diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index f27d26d6e..0ea4d7edf 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -17,6 +17,7 @@ from app.core.task_enqueue import ( safe_enqueue_generation_task, ) from app.dependencies import ( + get_db_session, get_generated_video_repository, get_generation_task_repository, ) @@ -25,7 +26,11 @@ from app.schemas.generation_task import ( PreviewGenerationTaskResponse, ) from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.orm import Session +from packages.adapters.sqlalchemy_impl.template_repository import ( + SQLAlchemyTemplateRepository, +) from packages.application import ( CreateGenerationTaskCommand, CreateGenerationTaskUseCase, @@ -39,6 +44,13 @@ router = APIRouter() PREVIEW_RESOLUTION = "854x480" +# 模板 mode → 视频比例映射 +_TEMPLATE_MODE_TO_RATIO = { + "pip": "9:16", + "standard": "16:9", + "square": "1:1", +} + def _calc_preview_resolution(video_ratio: str = "") -> str: """根据视频比例计算预览分辨率(短边 480,长边按比例)。 @@ -55,6 +67,36 @@ def _calc_preview_resolution(video_ratio: str = "") -> str: return ratio_map.get(video_ratio.strip(), PREVIEW_RESOLUTION) +def _infer_video_ratio_from_template(template_id: str, db: Session) -> str: + """从模板 mode 推断视频比例,前端未传 video_ratio 时使用。 + + Returns: + 视频比例字符串(如 "9:16"),查询失败返回空字符串。 + """ + if not template_id: + return "" + try: + repo = SQLAlchemyTemplateRepository(db) + template = repo.get_by_id(template_id) + if template: + mode = getattr(template, "mode", "") or "" + ratio = _TEMPLATE_MODE_TO_RATIO.get(mode.strip(), "") + if ratio: + logger.info( + "[预览生成] 从模板 mode=%s 推断 video_ratio=%s", + mode, + ratio, + ) + return ratio + except Exception: + logger.warning( + "[预览生成] 查询模板失败,跳过 video_ratio 推断: template_id=%s", + template_id, + exc_info=True, + ) + return "" + + def _mark_task_failed(repo, task, reason: str) -> None: """入队失败时将任务标记为 failed,避免产生僵尸 pending 数据。""" try: @@ -146,6 +188,7 @@ def create_preview_generation_task( request: CreatePreviewGenerationTaskRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), generation_task_repository=Depends(get_generation_task_repository), + db: Session = Depends(get_db_session), ) -> PreviewGenerationTaskResponse: """创建预览生成任务。 @@ -176,7 +219,7 @@ def create_preview_generation_task( except UserPendingLimitExceeded as e: raise HTTPException( status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待完成后再提交", + detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", ) from e except GlobalQueueFull as e: raise HTTPException( @@ -184,6 +227,11 @@ def create_preview_generation_task( detail="系统繁忙,请稍后再试", ) from e + # 确定视频比例:优先前端传入,否则从模板 mode 推断 + video_ratio = request.video_ratio or "" + if not video_ratio and request.template_id: + video_ratio = _infer_video_ratio_from_template(request.template_id, db) + use_case = CreateGenerationTaskUseCase(generation_task_repository) try: @@ -202,7 +250,7 @@ def create_preview_generation_task( asset_select_mode="", batch_id="", video_title=request.video_title, - resolution=_calc_preview_resolution(request.video_ratio), + resolution=_calc_preview_resolution(video_ratio), bgm_config=request.bgm_config or {}, auto_retry_enabled=False, auto_retry_max=0, @@ -214,7 +262,7 @@ def create_preview_generation_task( raise HTTPException(status_code=400, detail=str(e)) from e except Exception as e: logger.error("[预览生成] 创建失败: %s", e, exc_info=True) - raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后重试") from e + raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e # 入队执行;若入队失败则标记任务为 failed 避免僵尸数据 try: @@ -228,11 +276,11 @@ def create_preview_generation_task( logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id) _mark_task_failed(generation_task_repository, task, "任务入队失败") raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试") - except UserPendingLimitExceeded: + except UserPendingLimitExceeded as e: _mark_task_failed(generation_task_repository, task, "待处理任务超限") raise HTTPException( status_code=429, - detail="您的待处理任务过多,请等待完成后再提交", + detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", ) from None except GlobalQueueFull: _mark_task_failed(generation_task_repository, task, "系统队列已满") diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index de1ad7bec..11b584146 100755 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -726,6 +726,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert resp.task_id == "preview_task_001" assert resp.status == "pending" @@ -743,6 +744,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 429 @@ -759,6 +761,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 503 @@ -777,6 +780,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 400 @@ -795,6 +799,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 500 @@ -818,6 +823,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 500 @@ -841,6 +847,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 429 @@ -864,6 +871,7 @@ class TestCreatePreviewRoute: self._make_request(), authenticated_user=_make_user(), generation_task_repository=repo, + db=MagicMock(), ) assert exc_info.value.status_code == 503 @@ -1152,3 +1160,126 @@ class TestCalcPreviewResolution: from app.api.routes.generation_preview import _calc_preview_resolution assert _calc_preview_resolution("") == "854x480" + + +class TestInferVideoRatioFromTemplate: + """_infer_video_ratio_from_template 单元测试。""" + + def test_pip_mode_returns_9_16(self): + """模板 mode=pip → 返回 '9:16'""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_template = MagicMock() + mock_template.mode = "pip" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get_by_id.return_value = mock_template + result = _infer_video_ratio_from_template("tpl_001", mock_db) + assert result == "9:16" + + def test_standard_mode_returns_16_9(self): + """模板 mode=standard → 返回 '16:9'""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_template = MagicMock() + mock_template.mode = "standard" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get_by_id.return_value = mock_template + result = _infer_video_ratio_from_template("tpl_001", mock_db) + assert result == "16:9" + + def test_unknown_mode_returns_empty(self): + """模板 mode 未知 → 返回空字符串""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_template = MagicMock() + mock_template.mode = "unknown_mode" + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get_by_id.return_value = mock_template + result = _infer_video_ratio_from_template("tpl_001", mock_db) + assert result == "" + + def test_template_not_found_returns_empty(self): + """模板不存在 → 返回空字符串""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get_by_id.return_value = None + result = _infer_video_ratio_from_template("nonexistent", mock_db) + assert result == "" + + def test_empty_template_id_returns_empty(self): + """空 template_id → 直接返回空字符串""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_db = MagicMock() + result = _infer_video_ratio_from_template("", mock_db) + assert result == "" + + def test_db_exception_returns_empty(self): + """DB 异常 → 返回空字符串,不抛出""" + from app.api.routes.generation_preview import _infer_video_ratio_from_template + + mock_db = MagicMock() + + with patch("app.api.routes.generation_preview.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.side_effect = Exception("db connection error") + result = _infer_video_ratio_from_template("tpl_001", mock_db) + assert result == "" + + +class TestPreviewRouteAutoInfersVideoRatio: + """验证预览路由在前端未传 video_ratio 时自动从模板推断。""" + + def test_auto_infer_pip_resolution(self): + """前端传 video_ratio='',模板 mode=pip → resolution=480x854""" + 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_template = MagicMock() + mock_template.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.SQLAlchemyTemplateRepository") as MockRepo: + MockRepo.return_value.get_by_id.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}"