fix: 预览分辨率自动从模板mode推断(Fix3补充) #1230

Merged
xiaoxia merged 7 commits from fix/preview-resolution-auto-infer into develop 2026-08-04 01:08:23 +08:00
2 changed files with 184 additions and 5 deletions
+53 -5
View File
@@ -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, "系统队列已满")
+131
View File
@@ -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}"