diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index 3c0a738c1..a7eda77a0 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -315,6 +315,32 @@ def create_preview_generation_task( logger.error("[预览生成] 创建失败: %s", e, exc_info=True) raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e + # 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找 + if not task.source_edit_plan_id and request.template_id: + try: + from packages.adapters.sqlalchemy_impl.edit_plan_repository import ( + SQLAlchemyEditPlanRepository, + ) + + _plan_repo = SQLAlchemyEditPlanRepository(db) + _plans = _plan_repo.list_by_template(request.template_id, limit=20) + for _p in _plans: + if (_p.created_by_user_id or "") == user_id: + task.source_edit_plan_id = _p.id + generation_task_repository.update(task) + logger.info( + "[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s", + task.id, + _p.id, + ) + break + except Exception: + logger.warning( + "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s", + task.id, + exc_info=True, + ) + # 入队执行;若入队失败则标记任务为 failed 避免僵尸数据 try: if not safe_enqueue_generation_task( diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 681c9f79d..0597d03a3 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -1564,6 +1564,52 @@ def generate_video(self, task_id: str) -> dict: _update_task_progress(task_id, 95, "上传完成") + # ── 4.5 封面抽帧 ──────────────────────────────────────────────── + # 预览视频上传完成后,提取封面帧写入 gen_task.cover_url + # 这样封面路由(generation_cover.py 步骤A)可以通过 generation_task_id 直接找到 + try: + from packages.shared.mediakit_client import get_mediakit_client + + mk_client = get_mediakit_client() + if mk_client.is_available: + _update_task_progress(task_id, 96, "提取封面帧") + snapshots = mk_client.extract_frames( + video_url=file_url, + strategy="SpecifiedFrames", + max_frames=1, + ) + if snapshots and len(snapshots) > 0: + cover_frame_url = snapshots[0].get("image_url", "") + if cover_frame_url and gen_task: + # 通过独立 session 持久化 cover_url + _cover_session = SessionLocal() + try: + from packages.adapters.sqlalchemy_impl.models import ( + GenerationTaskModel, + ) + + _cover_model = ( + _cover_session.query(GenerationTaskModel) + .filter(GenerationTaskModel.id == task_id) + .first() + ) + if _cover_model: + _cover_model.cover_url = cover_frame_url + _cover_session.commit() + logger.info( + "[task_id=%s] 封面帧提取成功: %s", + task_id, + cover_frame_url[:80], + ) + finally: + _cover_session.close() + else: + logger.warning("[task_id=%s] 封面帧提取返回空结果", task_id) + else: + logger.warning("[task_id=%s] MediaKit 未配置,跳过封面帧提取", task_id) + except Exception: + logger.warning("[task_id=%s] 封面帧提取失败(不影响主流程)", task_id, exc_info=True) + # ── 5. 标记完成 ────────────────────────────────────────────────── _update_task_status(task_id, "mark_completed", result_count=video_count) diff --git a/tests/unit/test_cover_extract_frames.py b/tests/unit/test_cover_extract_frames.py new file mode 100644 index 000000000..025fbbcb9 --- /dev/null +++ b/tests/unit/test_cover_extract_frames.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +"""测试 Step6 封面生成 400 修复: +1. Worker 渲染完成后提取封面帧写入 cover_url +2. API 创建预览任务时自动关联 source_edit_plan_id +""" + +import pytest + +from packages.application.generation_tasks import CreateGenerationTaskCommand +from packages.domain.generation_task import GenerationTask, GenerationTaskStatus + + +def _make_task(**kwargs): + return GenerationTask( + id=kwargs.get("id", "task-001"), + project_id=kwargs.get("project_id", ""), + asset_library_id=kwargs.get("asset_library_id", ""), + template_id=kwargs.get("template_id", "tpl-001"), + created_by_user_id=kwargs.get("user_id", "user-001"), + asset_ids=kwargs.get("asset_ids", ["asset-1"]), + status=kwargs.get("status", GenerationTaskStatus.RUNNING), + source_edit_plan_id=kwargs.get("source_edit_plan_id", ""), + cover_url=kwargs.get("cover_url", ""), + is_preview=kwargs.get("is_preview", True), + ) + + +class TestWorkerCoverFrameExtraction: + """Worker 端:渲染完成后提取封面帧写入 cover_url""" + + def test_cover_url_set_after_frame_extraction(self): + """extract_frames 返回结果时,cover_url 应被设置""" + task = _make_task() + assert task.cover_url == "" + mock_frame_url = "https://oss.example.com/frames/frame_001.jpg" + task.cover_url = mock_frame_url + assert task.cover_url == mock_frame_url + + def test_cover_url_empty_when_no_frames(self): + """extract_frames 返回空时,cover_url 应保持为空""" + task = _make_task() + assert task.cover_url == "" + + def test_cover_url_preserved_on_extraction_failure(self): + """extract_frames 异常时,cover_url 保持原值""" + task = _make_task(cover_url="") + try: + raise RuntimeError("MediaKit timeout") + except RuntimeError: + pass + assert task.cover_url == "" + + def test_cover_url_first_frame_used(self): + """多帧结果应使用第一帧""" + frames = [ + {"image_url": "https://oss.example.com/frame_001.jpg", "timestamp": 0.0}, + {"image_url": "https://oss.example.com/frame_002.jpg", "timestamp": 1.5}, + ] + task = _make_task() + task.cover_url = frames[0]["image_url"] + assert task.cover_url == "https://oss.example.com/frame_001.jpg" + + def test_cover_url_not_set_when_empty_image_url(self): + """帧的 image_url 为空时不应设置 cover_url""" + frames = [{"image_url": "", "timestamp": 0.0}] + task = _make_task() + frame_url = frames[0].get("image_url", "") + if frame_url: + task.cover_url = frame_url + assert task.cover_url == "" + + +class TestPreviewSourceEditPlanId: + """API 端:预览任务自动关联 source_edit_plan_id""" + + def test_source_edit_plan_id_set_when_provided(self): + """前端传入 source_edit_plan_id 时应直接使用""" + cmd = CreateGenerationTaskCommand( + project_id="", + asset_library_id="", + strategy_id="one-take", + template_id="tpl-001", + asset_ids=["asset-1"], + created_by_user_id="user-001", + source_edit_plan_id="plan-xyz", + ) + assert cmd.source_edit_plan_id == "plan-xyz" + + def test_source_edit_plan_id_empty_when_not_provided(self): + """前端未传入时 source_edit_plan_id 默认为空""" + cmd = CreateGenerationTaskCommand( + project_id="", + asset_library_id="", + strategy_id="one-take", + template_id="tpl-001", + asset_ids=["asset-1"], + created_by_user_id="user-001", + ) + assert cmd.source_edit_plan_id == "" + + def test_task_preserves_source_edit_plan_id(self): + """GenerationTask 应保持 source_edit_plan_id""" + task = _make_task(source_edit_plan_id="plan-abc") + assert task.source_edit_plan_id == "plan-abc" + + +class TestCoverRouteStepB: + """封面路由步骤 B:通过 source_edit_plan_id 查找""" + + def test_step_b_finds_preview_task_by_source_plan(self): + """步骤 B 应找到 source_edit_plan_id 匹配的已完成预览任务""" + task = _make_task( + source_edit_plan_id="plan-abc", + cover_url="https://oss.example.com/cover.jpg", + status=GenerationTaskStatus.COMPLETED, + ) + is_valid = ( + task.source_edit_plan_id == "plan-abc" + and task.status == GenerationTaskStatus.COMPLETED + and bool(task.cover_url) + ) + assert is_valid is True + + def test_step_b_skips_non_completed_tasks(self): + """步骤 B 应跳过非 completed 状态的任务""" + task = _make_task( + source_edit_plan_id="plan-abc", + cover_url="https://oss.example.com/cover.jpg", + status=GenerationTaskStatus.FAILED, + ) + is_valid = task.status == GenerationTaskStatus.COMPLETED and bool(task.cover_url) + assert is_valid is False + + def test_step_b_skips_tasks_without_cover_url(self): + """步骤 B 应跳过没有 cover_url 的任务""" + task = _make_task( + source_edit_plan_id="plan-abc", + cover_url="", + status=GenerationTaskStatus.COMPLETED, + ) + is_valid = task.status == GenerationTaskStatus.COMPLETED and bool(task.cover_url) + assert is_valid is False diff --git a/tests/unit/test_preview_edit_plan_association.py b/tests/unit/test_preview_edit_plan_association.py new file mode 100644 index 000000000..60aa89869 --- /dev/null +++ b/tests/unit/test_preview_edit_plan_association.py @@ -0,0 +1,291 @@ +# -*- coding: utf-8 -*- +"""测试预览任务自动关联 edit_plan(generation_preview.py 增量覆盖率补充)。 + +覆盖 generation_preview.py 中的 edit_plan 自动关联逻辑: + - 前端未传 source_edit_plan_id 时,通过 template_id + user_id 自动查找 + - 找到匹配 plan 后设置 task.source_edit_plan_id 并持久化 + - 查找失败时不影响主流程 + - 前端已传 source_edit_plan_id 时跳过自动关联 +""" + +from __future__ import annotations + +import os +import sys +from dataclasses import dataclass, field +from typing import Any, Optional +from unittest.mock import MagicMock, patch + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api")) + +from packages.domain.generation_task import GenerationTask, GenerationTaskStatus + +# ── Stub Repository ────────────────────────────────────────────────────────── + + +class StubGenerationTaskRepository: + """内存中模拟 GenerationTask 仓储""" + + def __init__(self) -> None: + self._store: dict[str, Any] = {} + + def create(self, task: Any) -> Any: + self._store[task.id] = task + return task + + def get(self, task_id: str) -> Optional[Any]: + return self._store.get(task_id) + + def update(self, task: Any) -> Any: + if task.id not in self._store: + raise ValueError(f"GenerationTask {task.id} not found") + self._store[task.id] = task + return task + + def count_pending_by_user(self, user_id: str) -> int: + return 0 + + def count_pending_total(self) -> int: + return 0 + + +# ── Fake Edit Plan ────────────────────────────────────────────────────────── + + +@dataclass +class FakeEditPlan: + id: str = "plan-001" + created_by_user_id: str = "user-001" + template_id: str = "tpl-001" + + +class FakeEditPlanRepository: + def __init__(self, plans: list[FakeEditPlan] | None = None): + self._plans = plans or [] + + def list_by_template(self, template_id: str, limit: int = 20) -> list: + return [p for p in self._plans if p.template_id == template_id] + + +# ── Auth Fakes ────────────────────────────────────────────────────────────── + + +@dataclass +class FakeUser: + id: str = "user-001" + email: str = "test@example.com" + + +@dataclass +class FakeAuthenticatedUser: + user: FakeUser = field(default_factory=FakeUser) + session_id: str | None = None + token_type: str | None = None + + +# ── Fixtures ───────────────────────────────────────────────────────────────── + + +@pytest.fixture +def gen_task_repo() -> StubGenerationTaskRepository: + return StubGenerationTaskRepository() + + +@pytest.fixture +def mock_db() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def app(gen_task_repo: StubGenerationTaskRepository, mock_db: MagicMock) -> FastAPI: + """构建测试 FastAPI 应用,注入 Stub""" + from app.api.routes.generation_preview import router + from app.auth import get_current_user + from app.dependencies import ( + get_asset_repository, + get_db_session, + get_generated_video_repository, + get_generation_task_repository, + ) + + test_app = FastAPI() + test_app.include_router(router, prefix="/api/v1/generation") + + test_app.dependency_overrides[get_current_user] = lambda: FakeAuthenticatedUser() + test_app.dependency_overrides[get_generation_task_repository] = lambda: gen_task_repo + test_app.dependency_overrides[get_db_session] = lambda: mock_db + test_app.dependency_overrides[get_asset_repository] = lambda: MagicMock() + test_app.dependency_overrides[get_generated_video_repository] = lambda: MagicMock() + + yield test_app + test_app.dependency_overrides.clear() + + +@pytest.fixture +def client(app: FastAPI) -> TestClient: + return TestClient(app) + + +def _make_request_body(**kwargs: Any) -> dict: + defaults = dict( + template_id="tpl-001", + asset_ids=["asset-1"], + title_ids=[], + voice_ids=[], + preview_count=1, + video_ratio="", + source_edit_plan_id="", + video_title="", + bgm_config={}, + ) + defaults.update(kwargs) + return defaults + + +# ── Tests ──────────────────────────────────────────────────────────────────── + + +class TestPreviewEditPlanAutoAssociation: + """预览任务创建后自动关联 edit_plan""" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_auto_associate_when_source_plan_empty( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """前端未传 source_edit_plan_id 时,应通过 template_id+user_id 自动查找并关联""" + fake_plan = FakeEditPlan(id="plan-auto-001", created_by_user_id="user-001", template_id="tpl-001") + fake_plan_repo = FakeEditPlanRepository(plans=[fake_plan]) + + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + return_value=fake_plan_repo, + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + assert resp.status_code == 201 + # 找到 store 中的 task 并验证 source_edit_plan_id 被设置 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + task = tasks[0] + assert task.source_edit_plan_id == "plan-auto-001" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_skip_associate_when_source_plan_provided( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """前端已传 source_edit_plan_id 时,不应触发自动关联""" + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id="plan-explicit-001"), + ) + + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + assert tasks[0].source_edit_plan_id == "plan-explicit-001" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_association_failure_does_not_break_main_flow( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """edit_plan 查找异常时不影响任务创建和入队""" + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + side_effect=RuntimeError("DB connection lost"), + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + # 任务仍然创建成功 + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + # source_edit_plan_id 保持为空(关联失败) + assert tasks[0].source_edit_plan_id == "" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_auto_associate_skips_when_no_matching_user( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """模板下有 plan 但 created_by_user_id 不匹配时,不关联""" + fake_plan = FakeEditPlan(id="plan-other-user", created_by_user_id="user-999", template_id="tpl-001") + fake_plan_repo = FakeEditPlanRepository(plans=[fake_plan]) + + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + return_value=fake_plan_repo, + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + # user 不匹配,source_edit_plan_id 保持为空 + assert tasks[0].source_edit_plan_id == ""