From f5f54b98b6ec9f3c2978d62e9fa7c3befae3254d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=94=A8=E6=88=B7CI=20Test?= Date: Thu, 9 Jul 2026 15:34:35 +0800 Subject: [PATCH 1/2] =?UTF-8?q?fix(generation):=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E7=94=9F=E6=88=90=E4=BB=BB=E5=8A=A1=E7=8A=B6=E6=80=81=E6=9C=BA?= =?UTF-8?q?=E7=BC=BA=E5=A4=B1=E5=AF=BC=E8=87=B4=E7=8A=B6=E6=80=81=E6=B0=B8?= =?UTF-8?q?=E4=B8=8D=E6=9B=B4=E6=96=B0=E7=9A=84=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 问题根因 generate_video Celery 任务完全没有状态更新逻辑: - 任务开始执行时,不更新 GenerationTask 状态为 running - 任务成功时,不更新状态为 completed - 任务失败时,不更新状态为 failed - GenerationTask 领域模型缺少 mark_processing/mark_completed/mark_failed 方法 结果:任务状态永远停留在 pending,前端查询时一直显示"生成中"或超时显示失败,即使视频实际上已经生成成功了。 ## 修复方案 ### 1. GenerationTask 领域模型增加完整状态机 - 新增 TERMINAL_STATUSES 终态集合 - 新增 _VALID_TRANSITIONS 合法状态转换表 - 新增 transition_to() 通用状态转换方法(含合法性校验) - 新增 mark_processing() — pending → running,设置 started_at - 新增 mark_completed(result_count) — running → completed,设置 completed_at/progress/result_count - 新增 mark_failed(error_message) — pending/running → failed,设置 error_message/completed_at - 新增 mark_cancelled() — pending/running → cancelled - 新增 mark_pending_from_failed() — failed → pending(用于重试) - 新增 is_terminal/is_completed/is_failed/is_running 属性 状态机: pending → running → completed ↘ failed → pending (重试) ↘ cancelled ### 2. generate_video Celery 任务增加状态更新 - 新增 _update_task_status() 辅助函数:独立 session + 异常隔离 - 任务开始时:调用 mark_processing() 更新为 running - 任务成功时:调用 mark_completed() 更新为 completed,传入视频数量 - 任务失败时:在 except 块中调用 mark_failed() 记录错误信息 - _create_video_record_and_dedup 返回值改为 int,方便统计成功数量 状态更新使用独立数据库 session,且所有状态更新操作都包裹在 try/except 中,确保不会因为状态更新失败导致整个任务异常。 ### 3. 新增状态机单元测试 新增 tests/unit/test_generation_task_status.py,覆盖 42 个测试用例: - 初始状态验证 - mark_processing 正常/异常路径 - mark_completed 正常/异常路径 - mark_failed 正常/异常路径 - mark_cancelled 正常/异常路径 - mark_pending_from_failed 重试路径 - transition_to 通用方法 - 完整流转路径(成功、失败、重试、取消) ## 影响范围 - packages/domain/generation_task.py — 领域模型状态机方法 - apps/worker/worker_app/tasks/generation.py — Celery 任务状态更新 - tests/unit/test_generation_task_status.py — 新增单元测试 ## 验证方式 - 42 个状态机单元测试全部通过 - 现有 8 个生成相关测试全部通过 - 现有 15 个生成 API 集成测试全部通过 - 现有 6 个编辑计划 worker 失败测试全部通过 - 现有 25 个生成视频管理测试全部通过 --- apps/worker/worker_app/tasks/generation.py | 97 ++++- packages/domain/generation_task.py | 162 ++++++++ tests/unit/test_generation_task_status.py | 455 +++++++++++++++++++++ 3 files changed, 701 insertions(+), 13 deletions(-) create mode 100644 tests/unit/test_generation_task_status.py diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 08ff3119b..4cdde7145 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -14,6 +14,7 @@ from urllib.parse import urlparse import oss2 from worker_app.celery_app import celery_app +from worker_app.db import SessionLocal OUTPUT_WIDTH = 1280 OUTPUT_HEIGHT = 720 @@ -28,6 +29,58 @@ PUBLIC_API_BASE_URL = os.getenv("PUBLIC_API_BASE_URL", "https://api.xiaoxiajianj logger = logging.getLogger(__name__) +# ── 状态更新辅助函数 ────────────────────────────────────────────────────────── + + +def _update_task_status(task_id: str, status_action: str, **kwargs) -> bool: + """更新 GenerationTask 状态(独立 session,异常不向外抛出)。 + + Args: + task_id: 任务 ID + status_action: 状态动作名,如 "mark_processing" / "mark_completed" / "mark_failed" + **kwargs: 传递给对应方法的参数 + + Returns: + True 表示更新成功,False 表示更新失败 + """ + try: + from packages.adapters.sqlalchemy_impl.generation_task_repository import ( + SQLAlchemyGenerationTaskRepository, + ) + + session = SessionLocal() + try: + repo = SQLAlchemyGenerationTaskRepository(session) + task = repo.get(task_id) + if task is None: + logger.warning("更新任务状态失败:任务不存在 task_id=%s", task_id) + return False + + action = getattr(task, status_action, None) + if action is None: + logger.warning("未知的状态动作: %s", status_action) + return False + + action(**kwargs) + repo.update(task) + logger.info("GenerationTask 状态更新成功: task_id=%s action=%s", task_id, status_action) + return True + finally: + session.close() + except Exception as e: + logger.error( + "更新 GenerationTask 状态异常: task_id=%s action=%s error=%s", + task_id, + status_action, + e, + exc_info=True, + ) + return False + + +# ── FFmpeg / OSS helpers ───────────────────────────────────────────────────── + + def _run_ffmpeg(command: list[str]) -> None: """执行 FFmpeg 命令""" subprocess.run(command, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) # nosec B603 @@ -155,8 +208,6 @@ def _download_library_assets( """ # 导入模型和会话 try: - from worker_app.db import SessionLocal - from packages.adapters.sqlalchemy_impl.models import AssetModel session = SessionLocal() @@ -231,6 +282,9 @@ def _process_with_editing_mode( ) +# ── Celery Task ────────────────────────────────────────────────────────────── + + @celery_app.task(bind=True, name="worker.generate_video", max_retries=2) def generate_video(self, task_id: str) -> dict: """ @@ -242,19 +296,21 @@ def generate_video(self, task_id: str) -> dict: Returns: 生成结果字典 """ - from worker_app.db import SessionLocal + from packages.domain import EditingMode - from packages.adapters.sqlalchemy_impl.generation_task_repository import ( - SQLAlchemyGenerationTaskRepository, - ) - from packages.domain import EditingMode, GeneratedVideo, GenerationTaskStatus + logger.info("开始生成视频任务: task_id=%s", task_id) # 从数据库加载任务信息 session = SessionLocal() try: + from packages.adapters.sqlalchemy_impl.generation_task_repository import ( + SQLAlchemyGenerationTaskRepository, + ) + task_repo = SQLAlchemyGenerationTaskRepository(session) gen_task = task_repo.get(task_id) if gen_task is None: + logger.error("生成任务不存在: task_id=%s", task_id) return {"status": "failed", "error": f"generation task {task_id} not found"} project_id = gen_task.project_id asset_library_id = gen_task.asset_library_id @@ -265,6 +321,9 @@ def generate_video(self, task_id: str) -> dict: finally: session.close() + # 标记任务为 running + _update_task_status(task_id, "mark_processing") + try: editing_mode = EditingMode(mode) except ValueError: @@ -315,7 +374,7 @@ def generate_video(self, task_id: str) -> dict: file_url = f"{GENERATED_FILES_URL_PREFIX}/{task_id}/{output_name}" # 创建 GeneratedVideo 记录 + 查重 - _create_video_record_and_dedup( + video_count = _create_video_record_and_dedup( task_id=task_id, project_id=project_id, batch_id=batch_id, @@ -326,6 +385,11 @@ def generate_video(self, task_id: str) -> dict: mode=editing_mode.value, ) + # 标记任务为 completed + _update_task_status(task_id, "mark_completed", result_count=video_count or 1) + + logger.info("视频生成完成: task_id=%s duration=%.2fs file_size=%d", task_id, duration, file_size) + return { "status": "completed", "task_id": task_id, @@ -337,7 +401,9 @@ def generate_video(self, task_id: str) -> dict: "mode": editing_mode.value, } except Exception as error: - logger.error(f"Video generation failed: {error}") + logger.error(f"Video generation failed: {error}", exc_info=True) + # 标记任务为 failed + _update_task_status(task_id, "mark_failed", error_message=str(error)) return { "status": "failed", "task_id": task_id, @@ -355,12 +421,15 @@ def _create_video_record_and_dedup( duration: float, video_path: str, mode: str, -) -> None: - """创建 GeneratedVideo 记录,计算指纹并执行查重(历史 + 批次)。""" +) -> int: + """创建 GeneratedVideo 记录,计算指纹并执行查重(历史 + 批次)。 + + Returns: + 创建的视频记录数量(1 表示成功,0 表示失败) + """ from uuid import uuid4 from video_processing.dedup import VideoDeduplicator - from worker_app.db import SessionLocal from packages.adapters.sqlalchemy_impl.generated_video_repository import ( SQLAlchemyGeneratedVideoRepository, @@ -395,7 +464,7 @@ def _create_video_record_and_dedup( except Exception as fp_err: logger.warning(f"Fingerprint computation failed for {video_id}: {fp_err}") session.commit() - return + return 1 generated_video.video_fingerprint = fingerprint.to_dict() @@ -420,8 +489,10 @@ def _create_video_record_and_dedup( video_repo.update(generated_video) session.commit() logger.info(f"GeneratedVideo record created: {video_id} (task={task_id}, dup={generated_video.is_duplicate})") + return 1 except Exception as e: logger.error(f"Failed to create video record / dedup for task {task_id}: {e}") session.rollback() + return 0 finally: session.close() diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index f9f0b2541..6308e28f8 100644 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -1,3 +1,11 @@ +"""GenerationTask 领域模型 — 视频生成任务. + +状态机: + pending → running → completed + ↘ failed → pending (重试) + ↘ cancelled +""" + from __future__ import annotations import sys @@ -17,11 +25,43 @@ from uuid import uuid4 class GenerationTaskStatus(StrEnum): + """生成任务状态枚举。""" + PENDING = "pending" + """待处理(任务已创建,等待执行)""" + RUNNING = "running" + """运行中(正在生成视频)""" + COMPLETED = "completed" + """已完成(视频生成成功)""" + FAILED = "failed" + """失败(生成失败)""" + CANCELLED = "cancelled" + """已取消(用户取消或系统取消)""" + + +# 终态集合 +TERMINAL_STATUSES = frozenset( + {GenerationTaskStatus.COMPLETED, GenerationTaskStatus.FAILED, GenerationTaskStatus.CANCELLED} +) + +# 合法状态转换 +_VALID_TRANSITIONS: dict[GenerationTaskStatus, set[GenerationTaskStatus]] = { + GenerationTaskStatus.PENDING: { + GenerationTaskStatus.RUNNING, + GenerationTaskStatus.FAILED, + GenerationTaskStatus.CANCELLED, + }, + GenerationTaskStatus.RUNNING: { + GenerationTaskStatus.COMPLETED, + GenerationTaskStatus.FAILED, + GenerationTaskStatus.CANCELLED, + }, + GenerationTaskStatus.FAILED: {GenerationTaskStatus.PENDING}, # 重试回到 pending +} @dataclass(slots=True) @@ -83,3 +123,125 @@ class GenerationTask: asset_select_mode=asset_select_mode, batch_id=batch_id, ) + + # ── 状态查询 ──────────────────────────────────────────────────────────── + + @property + def is_terminal(self) -> bool: + """是否处于终态(completed / failed / cancelled)。""" + return self.status in TERMINAL_STATUSES + + @property + def is_completed(self) -> bool: + """是否已完成。""" + return self.status == GenerationTaskStatus.COMPLETED + + @property + def is_failed(self) -> bool: + """是否失败。""" + return self.status == GenerationTaskStatus.FAILED + + @property + def is_running(self) -> bool: + """是否运行中。""" + return self.status == GenerationTaskStatus.RUNNING + + # ── 状态转换 ──────────────────────────────────────────────────────────── + + def transition_to(self, new_status: GenerationTaskStatus | str) -> None: + """执行状态转换。 + + Args: + new_status: 目标状态 + + Raises: + ValueError: 非法状态转换 + """ + if isinstance(new_status, str): + try: + new_status = GenerationTaskStatus(new_status) + except ValueError: + raise ValueError(f"无效状态: {new_status}") + + allowed = _VALID_TRANSITIONS.get(self.status, set()) + if new_status not in allowed: + raise ValueError( + f"非法状态转换: {self.status.value} → {new_status.value}," + f"允许: {{{', '.join(sorted(s.value for s in allowed))}}}" + ) + + self.status = new_status + + def mark_processing(self) -> None: + """标记为处理中(pending → running)。 + + 设置 started_at,清除 error_message。 + + Raises: + ValueError: 当前状态不允许转换到 running + """ + self.transition_to(GenerationTaskStatus.RUNNING) + self.started_at = datetime.now(timezone.utc) + self.error_message = "" + + def mark_completed(self, result_count: int = 1) -> None: + """标记为已完成(running → completed)。 + + 设置 completed_at、progress=100.0、result_count,清除 error_message。 + + Args: + result_count: 生成的视频数量,默认为 1 + + Raises: + ValueError: 当前状态不允许转换到 completed + """ + self.transition_to(GenerationTaskStatus.COMPLETED) + self.completed_at = datetime.now(timezone.utc) + self.progress = 100.0 + self.result_count = result_count + self.error_message = "" + + def mark_failed(self, error_message: str) -> None: + """标记为失败(pending / running → failed)。 + + 设置 error_message、completed_at。 + + Args: + error_message: 错误信息 + + Raises: + ValueError: 当前状态不允许转换到 failed + """ + self.transition_to(GenerationTaskStatus.FAILED) + self.error_message = error_message + self.completed_at = datetime.now(timezone.utc) + + def mark_cancelled(self) -> None: + """标记为已取消(pending / running → cancelled)。 + + 设置 completed_at。 + + Raises: + ValueError: 当前状态不允许转换到 cancelled + """ + self.transition_to(GenerationTaskStatus.CANCELLED) + self.completed_at = datetime.now(timezone.utc) + + def mark_pending_from_failed(self) -> None: + """从失败状态重置为待处理(用于重试)。 + + 清除 error_message、started_at、completed_at、progress。 + + Raises: + ValueError: 当前状态不是 failed + """ + if self.status != GenerationTaskStatus.FAILED: + raise ValueError( + f"只有 failed 状态的任务可以重置为 pending,当前状态: {self.status.value}" + ) + self.transition_to(GenerationTaskStatus.PENDING) + self.error_message = "" + self.started_at = None + self.completed_at = None + self.progress = 0.0 + self.result_count = 0 diff --git a/tests/unit/test_generation_task_status.py b/tests/unit/test_generation_task_status.py new file mode 100644 index 000000000..dd46815e6 --- /dev/null +++ b/tests/unit/test_generation_task_status.py @@ -0,0 +1,455 @@ +"""GenerationTask 领域模型状态机单元测试. + +覆盖: +- 初始状态为 pending +- mark_processing: pending → running +- mark_completed: running → completed +- mark_failed: pending/running → failed +- mark_cancelled: pending/running → cancelled +- mark_pending_from_failed: failed → pending(重试) +- 非法状态转换抛出 ValueError +- is_terminal / is_completed / is_failed / is_running 属性 +- 状态转换时的时间戳设置 +""" + +from __future__ import annotations + +import pytest + +from packages.domain.generation_task import ( + TERMINAL_STATUSES, + GenerationTask, + GenerationTaskStatus, +) + + +def _make_task(**overrides) -> GenerationTask: + """创建一个测试用的 GenerationTask。""" + defaults = dict( + id="task-test-001", + project_id="proj-1", + asset_library_id="lib-1", + ) + defaults.update(overrides) + return GenerationTask(**defaults) + + +# ── 初始状态 ────────────────────────────────────────────────────────────────── + + +class TestInitialState: + """测试初始状态。""" + + def test_default_status_is_pending(self) -> None: + """新创建的任务默认状态为 pending。""" + task = _make_task() + assert task.status == GenerationTaskStatus.PENDING + assert task.progress == 0.0 + assert task.result_count == 0 + assert task.error_message == "" + assert task.started_at is None + assert task.completed_at is None + + def test_create_factory_returns_pending(self) -> None: + """GenerationTask.create() 返回的任务状态为 pending。""" + task = GenerationTask.create( + project_id="proj-1", + asset_library_id="lib-1", + created_by_user_id="user-1", + ) + assert task.status == GenerationTaskStatus.PENDING + assert task.progress == 0.0 + assert task.result_count == 0 + + def test_is_not_terminal_initially(self) -> None: + """初始状态不是终态。""" + task = _make_task() + assert not task.is_terminal + assert not task.is_completed + assert not task.is_failed + assert not task.is_running + + def test_terminal_statuses_constant(self) -> None: + """终态集合包含 completed / failed / cancelled。""" + assert GenerationTaskStatus.COMPLETED in TERMINAL_STATUSES + assert GenerationTaskStatus.FAILED in TERMINAL_STATUSES + assert GenerationTaskStatus.CANCELLED in TERMINAL_STATUSES + assert GenerationTaskStatus.PENDING not in TERMINAL_STATUSES + assert GenerationTaskStatus.RUNNING not in TERMINAL_STATUSES + + +# ── mark_processing ────────────────────────────────────────────────────────── + + +class TestMarkProcessing: + """测试 pending → running 转换。""" + + def test_pending_to_running_success(self) -> None: + """pending 状态的任务可以标记为 running。""" + task = _make_task() + task.mark_processing() + assert task.status == GenerationTaskStatus.RUNNING + assert task.is_running + assert task.started_at is not None + assert task.error_message == "" + + def test_started_at_is_set(self) -> None: + """mark_processing 设置 started_at 时间戳。""" + task = _make_task() + assert task.started_at is None + task.mark_processing() + assert task.started_at is not None + + def test_error_message_cleared(self) -> None: + """mark_processing 清除 error_message(如果有的话)。""" + task = _make_task() + # 注意:pending 状态通常没有 error_message,这里验证确保被清除 + task.error_message = "some old error" + # 直接设置状态绕过校验(模拟异常场景) + task.status = GenerationTaskStatus.PENDING + task.mark_processing() + assert task.error_message == "" + + def test_running_to_running_raises(self) -> None: + """running 状态不能再次 mark_processing。""" + task = _make_task() + task.mark_processing() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_processing() + + def test_completed_to_running_raises(self) -> None: + """completed 状态不能回到 running。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_processing() + + def test_failed_to_running_raises(self) -> None: + """failed 状态不能直接到 running(应先重置为 pending)。""" + task = _make_task() + task.mark_processing() + task.mark_failed("some error") + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_processing() + + +# ── mark_completed ─────────────────────────────────────────────────────────── + + +class TestMarkCompleted: + """测试 running → completed 转换。""" + + def test_running_to_completed_success(self) -> None: + """running 状态的任务可以标记为 completed。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + assert task.status == GenerationTaskStatus.COMPLETED + assert task.is_completed + assert task.is_terminal + assert task.completed_at is not None + + def test_progress_set_to_100(self) -> None: + """mark_completed 设置 progress 为 100.0。""" + task = _make_task() + task.mark_processing() + task.progress = 50.0 # 模拟中间进度 + task.mark_completed() + assert task.progress == 100.0 + + def test_default_result_count_is_1(self) -> None: + """默认 result_count 为 1。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + assert task.result_count == 1 + + def test_custom_result_count(self) -> None: + """可以指定 result_count。""" + task = _make_task() + task.mark_processing() + task.mark_completed(result_count=5) + assert task.result_count == 5 + + def test_error_message_cleared(self) -> None: + """mark_completed 清除 error_message。""" + task = _make_task() + task.mark_processing() + task.error_message = "temporary error" + task.mark_completed() + assert task.error_message == "" + + def test_completed_at_is_set(self) -> None: + """mark_completed 设置 completed_at。""" + task = _make_task() + task.mark_processing() + assert task.completed_at is None + task.mark_completed() + assert task.completed_at is not None + + def test_pending_to_completed_raises(self) -> None: + """pending 状态不能直接到 completed。""" + task = _make_task() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_completed() + + def test_completed_to_completed_raises(self) -> None: + """completed 状态不能再次 mark_completed。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_completed() + + def test_failed_to_completed_raises(self) -> None: + """failed 状态不能直接到 completed。""" + task = _make_task() + task.mark_processing() + task.mark_failed("error") + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_completed() + + +# ── mark_failed ────────────────────────────────────────────────────────────── + + +class TestMarkFailed: + """测试 pending/running → failed 转换。""" + + def test_pending_to_failed_success(self) -> None: + """pending 状态可以直接标记为 failed。""" + task = _make_task() + task.mark_failed("资源不足") + assert task.status == GenerationTaskStatus.FAILED + assert task.is_failed + assert task.is_terminal + assert task.error_message == "资源不足" + assert task.completed_at is not None + + def test_running_to_failed_success(self) -> None: + """running 状态可以标记为 failed。""" + task = _make_task() + task.mark_processing() + task.mark_failed("生成失败:FFmpeg 错误") + assert task.status == GenerationTaskStatus.FAILED + assert task.is_failed + assert task.is_terminal + assert task.error_message == "生成失败:FFmpeg 错误" + assert task.completed_at is not None + + def test_completed_to_failed_raises(self) -> None: + """completed 状态不能标记为 failed。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_failed("late error") + + def test_failed_to_failed_raises(self) -> None: + """failed 状态不能再次 mark_failed。""" + task = _make_task() + task.mark_failed("first error") + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_failed("second error") + + def test_error_message_preserved(self) -> None: + """错误信息被正确保存。""" + task = _make_task() + error_msg = "FFmpeg returned non-zero exit status 1" + task.mark_failed(error_msg) + assert task.error_message == error_msg + + +# ── mark_cancelled ─────────────────────────────────────────────────────────── + + +class TestMarkCancelled: + """测试 pending/running → cancelled 转换。""" + + def test_pending_to_cancelled_success(self) -> None: + """pending 状态可以取消。""" + task = _make_task() + task.mark_cancelled() + assert task.status == GenerationTaskStatus.CANCELLED + assert task.is_terminal + assert task.completed_at is not None + + def test_running_to_cancelled_success(self) -> None: + """running 状态可以取消。""" + task = _make_task() + task.mark_processing() + task.mark_cancelled() + assert task.status == GenerationTaskStatus.CANCELLED + assert task.is_terminal + + def test_completed_to_cancelled_raises(self) -> None: + """completed 状态不能取消。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_cancelled() + + def test_failed_to_cancelled_raises(self) -> None: + """failed 状态不能取消。""" + task = _make_task() + task.mark_failed("some error") + with pytest.raises(ValueError, match="非法状态转换"): + task.mark_cancelled() + + +# ── mark_pending_from_failed (重试) ───────────────────────────────────────── + + +class TestMarkPendingFromFailed: + """测试 failed → pending(重试)转换。""" + + def test_failed_to_pending_success(self) -> None: + """failed 状态可以重置为 pending(用于重试)。""" + task = _make_task() + task.mark_processing() + task.mark_failed("临时错误") + task.mark_pending_from_failed() + assert task.status == GenerationTaskStatus.PENDING + assert not task.is_terminal + assert task.error_message == "" + assert task.started_at is None + assert task.completed_at is None + assert task.progress == 0.0 + assert task.result_count == 0 + + def test_pending_to_pending_raises(self) -> None: + """pending 状态不能调用 mark_pending_from_failed。""" + task = _make_task() + with pytest.raises(ValueError, match="只有 failed 状态"): + task.mark_pending_from_failed() + + def test_running_to_pending_raises(self) -> None: + """running 状态不能调用 mark_pending_from_failed。""" + task = _make_task() + task.mark_processing() + with pytest.raises(ValueError, match="只有 failed 状态"): + task.mark_pending_from_failed() + + def test_completed_to_pending_raises(self) -> None: + """completed 状态不能调用 mark_pending_from_failed。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError, match="只有 failed 状态"): + task.mark_pending_from_failed() + + +# ── transition_to 通用方法 ─────────────────────────────────────────────────── + + +class TestTransitionTo: + """测试通用的 transition_to 方法。""" + + def test_string_status_conversion(self) -> None: + """可以传入字符串形式的状态。""" + task = _make_task() + task.transition_to("running") + assert task.status == GenerationTaskStatus.RUNNING + + def test_invalid_string_raises(self) -> None: + """无效的状态字符串抛出 ValueError。""" + task = _make_task() + with pytest.raises(ValueError, match="无效状态"): + task.transition_to("invalid_status") + + def test_enum_status(self) -> None: + """可以传入枚举形式的状态。""" + task = _make_task() + task.transition_to(GenerationTaskStatus.RUNNING) + assert task.status == GenerationTaskStatus.RUNNING + + def test_error_message_includes_allowed_statuses(self) -> None: + """错误信息包含允许的状态列表。""" + task = _make_task() + task.mark_processing() + task.mark_completed() + with pytest.raises(ValueError) as exc_info: + task.transition_to(GenerationTaskStatus.RUNNING) + assert "completed" in str(exc_info.value) + assert "running" in str(exc_info.value) + + +# ── 完整流转路径 ───────────────────────────────────────────────────────────── + + +class TestFullFlow: + """测试完整的状态流转路径。""" + + def test_happy_path(self) -> None: + """正常路径:pending → running → completed。""" + task = _make_task() + assert task.status == GenerationTaskStatus.PENDING + assert not task.is_terminal + + task.mark_processing() + assert task.status == GenerationTaskStatus.RUNNING + assert task.started_at is not None + assert not task.is_terminal + + task.mark_completed(result_count=3) + assert task.status == GenerationTaskStatus.COMPLETED + assert task.is_completed + assert task.is_terminal + assert task.completed_at is not None + assert task.result_count == 3 + assert task.progress == 100.0 + + def test_failure_path_from_running(self) -> None: + """失败路径:pending → running → failed。""" + task = _make_task() + task.mark_processing() + assert task.is_running + + task.mark_failed("网络超时") + assert task.is_failed + assert task.is_terminal + assert task.error_message == "网络超时" + assert task.completed_at is not None + + def test_failure_path_from_pending(self) -> None: + """失败路径:pending → failed(启动前校验失败等)。""" + task = _make_task() + task.mark_failed("参数校验失败") + assert task.is_failed + assert task.is_terminal + + def test_retry_path(self) -> None: + """重试路径:pending → running → failed → pending → running → completed。""" + task = _make_task() + + # 第一次尝试失败 + task.mark_processing() + task.mark_failed("临时错误") + assert task.is_failed + + # 重试 + task.mark_pending_from_failed() + assert task.status == GenerationTaskStatus.PENDING + assert task.error_message == "" + + # 第二次成功 + task.mark_processing() + task.mark_completed() + assert task.is_completed + + def test_cancel_from_pending(self) -> None: + """取消路径:pending → cancelled。""" + task = _make_task() + task.mark_cancelled() + assert task.status == GenerationTaskStatus.CANCELLED + assert task.is_terminal + + def test_cancel_from_running(self) -> None: + """取消路径:pending → running → cancelled。""" + task = _make_task() + task.mark_processing() + task.mark_cancelled() + assert task.status == GenerationTaskStatus.CANCELLED + assert task.is_terminal -- 2.54.0 From 4437a08d4fd058c4734ffee843d6cb55572780e0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=94=A8=E6=88=B7CI=20Test?= Date: Thu, 9 Jul 2026 15:37:02 +0800 Subject: [PATCH 2/2] style(domain): format generation_task.py with black --- packages/domain/generation_task.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index 6308e28f8..ca6b62a77 100644 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -236,9 +236,7 @@ class GenerationTask: ValueError: 当前状态不是 failed """ if self.status != GenerationTaskStatus.FAILED: - raise ValueError( - f"只有 failed 状态的任务可以重置为 pending,当前状态: {self.status.value}" - ) + raise ValueError(f"只有 failed 状态的任务可以重置为 pending,当前状态: {self.status.value}") self.transition_to(GenerationTaskStatus.PENDING) self.error_message = "" self.started_at = None -- 2.54.0