fix: 预览分辨率自动从模板mode推断(Fix3补充) #1230
@@ -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, "系统队列已满")
|
||||
|
||||
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user