Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 72f28ff4eb | |||
| 33a1485e28 | |||
| 34188881b3 |
@@ -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),
|
||||
|
||||
@@ -297,4 +297,3 @@ test.describe("Core generation flow", () => {
|
||||
expect(Array.isArray(tasksData.items)).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
@@ -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}'"
|
||||
|
||||
Reference in New Issue
Block a user