diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index db4390203..746eac25f 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -28,6 +28,9 @@ from app.schemas.generation_task import ( from fastapi import APIRouter, Depends, HTTPException from sqlalchemy.orm import Session +from packages.adapters.sqlalchemy_impl.edit_template_repository import ( + SQLAlchemyEditTemplateRepository, +) from packages.adapters.sqlalchemy_impl.template_repository import ( SQLAlchemyTemplateRepository, ) @@ -99,6 +102,61 @@ def _infer_video_ratio_from_template( return "" +def _resolve_strategy_id_from_template( + template_id: str, db: Session, user_id: str = "" +) -> str: + """从模板读取 editing_mode / mode 作为 strategy_id。 + + 优先查新模板系统(EditTemplate.editing_mode),fallback 旧模板(Template.mode)。 + Worker 端使用 strategy_id 作为渲染 mode,为空则默认 one_take。 + """ + if not template_id: + return "" + + # 优先查新模板系统 + try: + new_repo = SQLAlchemyEditTemplateRepository(db) + new_template = new_repo.get(template_id) + if new_template and getattr(new_template, "editing_mode", ""): + mode = new_template.editing_mode.strip() + if mode: + logger.info( + "[预览生成] 从新模板 editing_mode=%s (template_id=%s)", + mode, + template_id, + ) + return mode + except Exception: + logger.debug( + "[预览生成] 新模板查询失败,尝试旧模板: template_id=%s", + template_id, + exc_info=True, + ) + + # fallback 旧模板系统 + try: + old_repo = SQLAlchemyTemplateRepository(db) + old_template = old_repo.get(template_id, user_id) + if old_template: + mode = getattr(old_template, "mode", "") or "" + mode = mode.strip() + if mode: + logger.info( + "[预览生成] 从旧模板 mode=%s (template_id=%s)", + mode, + template_id, + ) + return mode + except Exception: + logger.warning( + "[预览生成] 旧模板查询也失败,strategy_id 留空: template_id=%s", + template_id, + exc_info=True, + ) + + return "" + + def _mark_task_failed(repo, task, reason: str) -> None: """入队失败时将任务标记为 failed,避免产生僵尸 pending 数据。""" try: @@ -234,6 +292,9 @@ def create_preview_generation_task( if not video_ratio and request.template_id: video_ratio = _infer_video_ratio_from_template(request.template_id, db, user_id) + # 从模板读取 editing_mode / mode 作为 strategy_id(渲染 pipeline 的 mode 参数) + strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id) + use_case = CreateGenerationTaskUseCase(generation_task_repository) try: @@ -241,7 +302,7 @@ def create_preview_generation_task( CreateGenerationTaskCommand( project_id="", asset_library_id="", - strategy_id="", + strategy_id=strategy_id, voice_library_id="", template_id=request.template_id, asset_ids=list(request.asset_ids), diff --git a/apps/web/e2e/core-generation.spec.ts b/apps/web/e2e/core-generation.spec.ts index afc997d2a..e33339e5d 100755 --- a/apps/web/e2e/core-generation.spec.ts +++ b/apps/web/e2e/core-generation.spec.ts @@ -297,4 +297,3 @@ test.describe("Core generation flow", () => { expect(Array.isArray(tasksData.items)).toBe(true) }) }) - diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index b23c6499c..0f02fb8e3 100755 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -1264,8 +1264,183 @@ class TestPreviewRouteAutoInfersVideoRatio: bgm_config={}, ) - with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: - MockRepo.return_value.get.return_value = mock_template + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = None # 新模板系统无数据,fallback 旧系统 + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get.return_value = mock_template + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.return_value = task + with patch( + "app.api.routes.generation_preview.safe_enqueue_generation_task", + return_value=True, + ): + resp = create_preview_generation_task( + request, + authenticated_user=_make_user(), + generation_task_repository=repo, + db=mock_db, + ) + + # Verify the resolution passed to CreateGenerationTaskCommand + call_args = MockUC.return_value.execute.call_args + cmd = call_args[0][0] + assert cmd.resolution == "480x854", f"Expected 480x854, got {cmd.resolution}" + + +# ═══════════════════════════════════════════════════════════════════════════════ +# _resolve_strategy_id_from_template 单元测试 +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestResolveStrategyIdFromTemplate: + """_resolve_strategy_id_from_template 单元测试:预览生成从模板读取 editing_mode。""" + + def test_new_template_found_returns_editing_mode(self): + """新模板系统找到模板 → 返回 editing_mode。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_new_template = MagicMock() + mock_new_template.editing_mode = "pip" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = mock_new_template + result = _resolve_strategy_id_from_template("tpl_001", mock_db, "user_1") + + assert result == "pip" + + def test_new_template_not_found_fallback_to_old(self): + """新模板系统未找到 → fallback 旧模板系统返回 mode。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_old_template = MagicMock() + mock_old_template.mode = "standard" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = None + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.return_value.get.return_value = mock_old_template + result = _resolve_strategy_id_from_template("tpl_001", mock_db, "user_1") + + assert result == "standard" + + def test_both_not_found_returns_empty(self): + """新旧模板系统都未找到 → 返回空字符串。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = None + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.return_value.get.return_value = None + result = _resolve_strategy_id_from_template("tpl_nonexistent", mock_db, "user_1") + + assert result == "" + + def test_empty_template_id_returns_empty(self): + """空 template_id → 直接返回空字符串。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_db = MagicMock() + result = _resolve_strategy_id_from_template("", mock_db, "user_1") + assert result == "" + + def test_new_template_exception_fallback_to_old(self): + """新模板系统异常 → 降级到旧模板系统。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_old_template = MagicMock() + mock_old_template.mode = "voice_over" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.side_effect = RuntimeError("db error") + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.return_value.get.return_value = mock_old_template + result = _resolve_strategy_id_from_template("tpl_001", mock_db, "user_1") + + assert result == "voice_over" + + def test_both_exception_returns_empty(self): + """新旧模板系统都异常 → 返回空字符串。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.side_effect = RuntimeError("new db error") + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.side_effect = RuntimeError("old db error") + result = _resolve_strategy_id_from_template("tpl_001", mock_db, "user_1") + + assert result == "" + + def test_new_template_empty_editing_mode_fallback(self): + """新模板找到但 editing_mode 为空 → fallback 旧模板。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_new_template = MagicMock() + mock_new_template.editing_mode = "" + mock_old_template = MagicMock() + mock_old_template.mode = "one_take" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = mock_new_template + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.return_value.get.return_value = mock_old_template + result = _resolve_strategy_id_from_template("tpl_001", mock_db, "user_1") + + assert result == "one_take" + + def test_voice_pip_mode(self): + """模板 editing_mode=voice_pip → 返回 'voice_pip'。""" + from app.api.routes.generation_preview import _resolve_strategy_id_from_template + + mock_new_template = MagicMock() + mock_new_template.editing_mode = "voice_pip" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = mock_new_template + result = _resolve_strategy_id_from_template("tpl_voice", mock_db, "user_1") + + assert result == "voice_pip" + + +class TestPreviewRoutePassesStrategyId: + """验证预览路由正确传递 strategy_id 到 CreateGenerationTaskCommand。""" + + def test_strategy_id_from_new_template(self): + """预览路由从新模板读取 strategy_id=pip。""" + from app.api.routes.generation_preview import create_preview_generation_task + from app.schemas.generation_task import CreatePreviewGenerationTaskRequest + + repo = MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + mock_db = MagicMock() + + mock_new_template = MagicMock() + mock_new_template.editing_mode = "pip" + + task = _make_task() + + request = CreatePreviewGenerationTaskRequest( + template_id="tpl_pip", + asset_ids=["a1"], + title_ids=[], + voice_ids=[], + video_title="test", + duration=0.0, + video_ratio="", + bgm_config={}, + ) + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = mock_new_template with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch( @@ -1279,7 +1454,53 @@ class TestPreviewRouteAutoInfersVideoRatio: db=mock_db, ) - # Verify the resolution passed to CreateGenerationTaskCommand call_args = MockUC.return_value.execute.call_args cmd = call_args[0][0] - assert cmd.resolution == "480x854", f"Expected 480x854, got {cmd.resolution}" + assert cmd.strategy_id == "pip", f"Expected strategy_id='pip', got '{cmd.strategy_id}'" + + def test_strategy_id_fallback_to_old_template(self): + """新模板无数据时,从旧模板读取 strategy_id=standard。""" + from app.api.routes.generation_preview import create_preview_generation_task + from app.schemas.generation_task import CreatePreviewGenerationTaskRequest + + repo = MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + mock_db = MagicMock() + + mock_old_template = MagicMock() + mock_old_template.mode = "standard" + + task = _make_task() + + request = CreatePreviewGenerationTaskRequest( + template_id="tpl_standard", + asset_ids=["a1"], + title_ids=[], + voice_ids=[], + video_title="test", + duration=0.0, + video_ratio="", + bgm_config={}, + ) + + with patch("app.api.routes.generation_preview.SQLAlchemyEditTemplateRepository") as MockNewRepo: + MockNewRepo.return_value.get.return_value = None + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockOldRepo: + MockOldRepo.return_value.get.return_value = mock_old_template + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.return_value = task + with patch( + "app.api.routes.generation_preview.safe_enqueue_generation_task", + return_value=True, + ): + resp = create_preview_generation_task( + request, + authenticated_user=_make_user(), + generation_task_repository=repo, + db=mock_db, + ) + + call_args = MockUC.return_value.execute.call_args + cmd = call_args[0][0] + assert cmd.strategy_id == "standard", f"Expected strategy_id='standard', got '{cmd.strategy_id}'"