feat(#1197): Phase 1 预览生成接口(单版本预览)
实现预览生成功能的后端 Phase 1:单版本预览接口(创建 + 查询)。 变更内容: - 领域模型 GenerationTask 新增 is_preview: bool = False 字段 - SQLAlchemy GenerationTaskModel 新增 is_preview 列(带索引) - 新增 alembic migration 053_generation_task_is_preview - 仓储层 _to_domain/create/update 同步 is_preview 字段(getattr 兼容旧数据) - CreateGenerationTaskCommand/UseCase 新增 is_preview 参数 - Schema 新增 CreatePreviewGenerationTaskRequest 和 PreviewGenerationTaskResponse - 新增 /api/v1/generation/preview 路由(POST 创建 + GET 查询) - Worker generate_video 检测 is_preview=True 时强制 854x480 + 1M 低码率 - 新增 29 个单元测试,全部通过 验收: - ✅ 所有新增文件写完,所有修改点完成 - ✅ 单元测试 29 个,全部通过 - ✅ 现有 generation 相关 174 个测试全部通过 - ✅ 代码通过 ruff + black 检查 - ✅ 在 feat/preview-generation-1197 分支上
This commit is contained in:
+61
@@ -0,0 +1,61 @@
|
||||
"""#1197 - 预览生成:generation_tasks 表新增 is_preview 字段
|
||||
|
||||
Revision ID: 053
|
||||
Revises: 052
|
||||
Create Date: 2026-08-15
|
||||
|
||||
Changes:
|
||||
1. generation_tasks 表新增 is_preview 字段,标记是否为预览生成任务(低清 480p)
|
||||
2. 默认 False,与现有正式生成任务兼容
|
||||
3. 加索引以支持按预览/正式任务筛选
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision = "053_generation_task_is_preview"
|
||||
down_revision = "052_generation_task_bgm_config"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if context.get_context().dialect.name == "postgresql":
|
||||
# 检查列是否已存在(幂等)
|
||||
result = conn.execute(
|
||||
sa.text(
|
||||
"SELECT column_name FROM information_schema.columns "
|
||||
"WHERE table_name = 'generation_tasks' AND column_name = 'is_preview'"
|
||||
)
|
||||
)
|
||||
if result.scalar() is not None:
|
||||
return
|
||||
|
||||
op.add_column(
|
||||
"generation_tasks",
|
||||
sa.Column("is_preview", sa.Boolean, nullable=False, server_default=sa.text("false")),
|
||||
)
|
||||
# 加索引
|
||||
op.create_index(
|
||||
"ix_generation_tasks_is_preview",
|
||||
"generation_tasks",
|
||||
["is_preview"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if context.get_context().dialect.name == "postgresql":
|
||||
result = conn.execute(
|
||||
sa.text(
|
||||
"SELECT column_name FROM information_schema.columns "
|
||||
"WHERE table_name = 'generation_tasks' AND column_name = 'is_preview'"
|
||||
)
|
||||
)
|
||||
if result.scalar() is None:
|
||||
return
|
||||
|
||||
op.drop_index("ix_generation_tasks_is_preview", table_name="generation_tasks")
|
||||
op.drop_column("generation_tasks", "is_preview")
|
||||
@@ -8,6 +8,7 @@ from app.api.routes.classification_jobs import router as classification_jobs_rou
|
||||
from app.api.routes.duplication import router as duplication_router
|
||||
from app.api.routes.feature_flags import router as feature_flags_router
|
||||
from app.api.routes.generation_tasks import router as generation_tasks_router
|
||||
from app.api.routes.generation_preview import router as generation_preview_router
|
||||
from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.internal_render import router as internal_render_router
|
||||
@@ -87,6 +88,11 @@ api_router.include_router(
|
||||
prefix="/generation",
|
||||
tags=["Generation"],
|
||||
)
|
||||
api_router.include_router(
|
||||
generation_preview_router,
|
||||
prefix="/generation",
|
||||
tags=["Generation"],
|
||||
)
|
||||
api_router.include_router(
|
||||
titles_router,
|
||||
prefix="/titles",
|
||||
|
||||
+230
@@ -0,0 +1,230 @@
|
||||
"""预览生成路由 — Phase 1:单版本预览接口(创建 + 查询)。
|
||||
|
||||
路径前缀:/api/v1/generation/preview(与 /generation/tasks 同体系)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.task_enqueue import (
|
||||
GLOBAL_PENDING_LIMIT,
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
from app.dependencies import (
|
||||
get_generated_video_repository,
|
||||
get_generation_task_repository,
|
||||
)
|
||||
from app.schemas.generation_task import (
|
||||
CreatePreviewGenerationTaskRequest,
|
||||
PreviewGenerationTaskResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from packages.application import (
|
||||
CreateGenerationTaskCommand,
|
||||
CreateGenerationTaskUseCase,
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
PREVIEW_RESOLUTION = "854x480"
|
||||
|
||||
|
||||
def _to_preview_response(task, generated_videos: list | None = None) -> PreviewGenerationTaskResponse:
|
||||
"""将领域任务对象转换为预览响应 DTO。
|
||||
|
||||
Args:
|
||||
task: GenerationTask 领域对象
|
||||
generated_videos: 生成的视频列表(可选),取第一个作为 video_url
|
||||
|
||||
Returns:
|
||||
PreviewGenerationTaskResponse
|
||||
"""
|
||||
video_url = ""
|
||||
duration = 0.0
|
||||
file_size = 0
|
||||
if generated_videos:
|
||||
first_video = generated_videos[0]
|
||||
video_url = getattr(first_video, "file_url", "") or ""
|
||||
duration = float(getattr(first_video, "duration", 0.0) or 0.0)
|
||||
file_size = int(getattr(first_video, "file_size", 0) or 0)
|
||||
|
||||
# 从 extra_meta / metadata 中提取统计信息(如果有)
|
||||
extra_meta = getattr(task, "extra_meta", {}) or {}
|
||||
clip_count = int(extra_meta.get("clip_count", len(getattr(task, "asset_ids", [])) or 0))
|
||||
transition_count = int(extra_meta.get("transition_count", max(0, clip_count - 1)))
|
||||
material_usage = extra_meta.get("material_usage", {}) or {}
|
||||
|
||||
# 计算生成耗时
|
||||
generate_duration = 0.0
|
||||
started_at = getattr(task, "started_at", None)
|
||||
completed_at = getattr(task, "completed_at", None)
|
||||
if started_at and completed_at:
|
||||
generate_duration = (completed_at - started_at).total_seconds()
|
||||
|
||||
return PreviewGenerationTaskResponse(
|
||||
task_id=task.id,
|
||||
status=task.status.value if hasattr(task.status, "value") else str(task.status),
|
||||
progress=float(task.progress or 0.0),
|
||||
is_preview=bool(getattr(task, "is_preview", True)),
|
||||
resolution=getattr(task, "resolution", PREVIEW_RESOLUTION) or PREVIEW_RESOLUTION,
|
||||
video_url=video_url,
|
||||
duration=duration,
|
||||
file_size=file_size,
|
||||
clip_count=clip_count,
|
||||
transition_count=transition_count,
|
||||
material_usage=material_usage,
|
||||
error_message=task.error_message or "",
|
||||
created_at=task.created_at,
|
||||
started_at=started_at,
|
||||
finished_at=completed_at,
|
||||
generate_duration=generate_duration,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201)
|
||||
def create_preview_generation_task(
|
||||
request: CreatePreviewGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generation_task_repository=Depends(get_generation_task_repository),
|
||||
) -> PreviewGenerationTaskResponse:
|
||||
"""创建预览生成任务。
|
||||
|
||||
预览为完整时长的低清版(480p + 低码率),效果与正式生成一致,仅清晰度降低。
|
||||
|
||||
Args:
|
||||
request: 预览任务创建请求(template_id + asset_ids 等)
|
||||
|
||||
Returns:
|
||||
201 + 预览任务详情
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
logger.info(
|
||||
"[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d",
|
||||
user_id,
|
||||
request.template_id,
|
||||
len(request.asset_ids),
|
||||
)
|
||||
|
||||
# 预检查队列限流
|
||||
try:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending + 1 > USER_PENDING_LIMIT:
|
||||
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending + 1, limit=USER_PENDING_LIMIT)
|
||||
if global_pending + 1 > GLOBAL_PENDING_LIMIT:
|
||||
raise GlobalQueueFull(pending_count=global_pending + 1, limit=GLOBAL_PENDING_LIMIT)
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待完成后再提交",
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
) from e
|
||||
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
|
||||
try:
|
||||
task = use_case.execute(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
strategy_id="",
|
||||
voice_library_id="",
|
||||
template_id=request.template_id,
|
||||
asset_ids=list(request.asset_ids),
|
||||
title_ids=list(request.title_ids),
|
||||
voice_ids=list(request.voice_ids),
|
||||
created_by_user_id=user_id,
|
||||
source_edit_plan_id="",
|
||||
asset_select_mode="",
|
||||
batch_id="",
|
||||
video_title=request.video_title,
|
||||
resolution=PREVIEW_RESOLUTION,
|
||||
bgm_config=request.bgm_config or {},
|
||||
auto_retry_enabled=False,
|
||||
auto_retry_max=0,
|
||||
is_preview=True,
|
||||
)
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.warning("[预览生成] 创建失败: %s", e)
|
||||
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
|
||||
|
||||
# 入队执行
|
||||
try:
|
||||
if not safe_enqueue_generation_task(
|
||||
task,
|
||||
generation_task_repository,
|
||||
user_id=user_id,
|
||||
log_prefix="[预览生成]",
|
||||
log_task_status=True,
|
||||
):
|
||||
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
|
||||
raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试")
|
||||
except UserPendingLimitExceeded:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail="您的待处理任务过多,请等待完成后再提交",
|
||||
) from None
|
||||
except GlobalQueueFull:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
) from None
|
||||
|
||||
return _to_preview_response(task)
|
||||
|
||||
|
||||
@router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse)
|
||||
def get_preview_generation_task(
|
||||
task_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generation_task_repository=Depends(get_generation_task_repository),
|
||||
generated_video_repository=Depends(get_generated_video_repository),
|
||||
) -> PreviewGenerationTaskResponse:
|
||||
"""查询预览生成任务状态。
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID
|
||||
|
||||
Returns:
|
||||
预览任务详情(含状态、进度、结果 URL 等)
|
||||
"""
|
||||
use_case = GetGenerationTaskUseCase(generation_task_repository)
|
||||
task = use_case.execute(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail=f"预览任务 {task_id} 不存在")
|
||||
|
||||
# 权限校验:任务必须属于当前用户
|
||||
task_user_id = getattr(task, "created_by_user_id", "") or ""
|
||||
if task_user_id and task_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权访问该任务")
|
||||
|
||||
# 校验是否为预览任务
|
||||
if not getattr(task, "is_preview", False):
|
||||
raise HTTPException(status_code=404, detail=f"预览任务 {task_id} 不存在")
|
||||
|
||||
# 查询生成的视频(取第一个)
|
||||
generated_videos = []
|
||||
status_val = task.status.value if hasattr(task.status, "value") else str(task.status)
|
||||
if status_val == "completed":
|
||||
list_use_case = ListGeneratedVideosByTaskUseCase(generated_video_repository)
|
||||
generated_videos = list_use_case.execute(task_id)
|
||||
|
||||
return _to_preview_response(task, generated_videos=generated_videos)
|
||||
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
@@ -122,3 +123,62 @@ class ListGenerationTasksResponse(BaseModel):
|
||||
"""用户级生成任务列表响应(跨 project)。"""
|
||||
|
||||
items: list[GenerationTaskResponse]
|
||||
|
||||
|
||||
# ── 预览生成(Phase 1) ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class CreatePreviewGenerationTaskRequest(BaseModel):
|
||||
"""创建预览生成任务请求。
|
||||
|
||||
仅支持模板模式:template_id + asset_ids 等素材 ID 列表。
|
||||
预览为完整时长低清版(480p + 低码率)。
|
||||
"""
|
||||
|
||||
template_id: str
|
||||
asset_ids: list[str] = Field(default_factory=list)
|
||||
title_ids: list[str] = Field(default_factory=list)
|
||||
voice_ids: list[str] = Field(default_factory=list)
|
||||
video_title: str = Field(default="", description="生成视频的标题/名称,为空则使用默认命名")
|
||||
duration: float = Field(default=0.0, ge=0, description="期望视频时长(秒),0 表示由模板决定")
|
||||
video_ratio: str = Field(default="", description="视频比例,如 16:9 / 9:16,为空使用模板默认")
|
||||
bgm_config: dict = Field(
|
||||
default_factory=dict,
|
||||
description="自定义BGM配置,覆盖模板BGM设置。支持 enabled/source/asset_id/preset_id/audio_url/volume 等字段",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest":
|
||||
if not self.template_id.strip():
|
||||
raise ValueError("template_id 不能为空")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_asset_ids(self) -> "CreatePreviewGenerationTaskRequest":
|
||||
if not self.asset_ids and not self.title_ids and not self.voice_ids:
|
||||
raise ValueError("asset_ids/title_ids/voice_ids 至少需要提供一个")
|
||||
return self
|
||||
|
||||
|
||||
class PreviewGenerationTaskResponse(BaseModel):
|
||||
"""预览生成任务响应。
|
||||
|
||||
包含任务状态、进度、分辨率、生成结果 URL 等关键字段。
|
||||
"""
|
||||
|
||||
task_id: str
|
||||
status: str
|
||||
progress: float
|
||||
is_preview: bool = True
|
||||
resolution: str = ""
|
||||
video_url: str = ""
|
||||
duration: float = 0.0
|
||||
file_size: int = 0
|
||||
clip_count: int = 0
|
||||
transition_count: int = 0
|
||||
material_usage: dict = Field(default_factory=dict)
|
||||
error_message: str = ""
|
||||
created_at: datetime | None = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
generate_duration: float = 0.0
|
||||
|
||||
@@ -943,6 +943,7 @@ def _load_task_info(task_id: str) -> dict | None:
|
||||
"video_title": getattr(gen_task, "video_title", "") or "",
|
||||
"resolution": getattr(gen_task, "resolution", "") or "",
|
||||
"bgm_config": dict(getattr(gen_task, "bgm_config", {}) or {}),
|
||||
"is_preview": bool(getattr(gen_task, "is_preview", False)),
|
||||
}
|
||||
finally:
|
||||
session.close()
|
||||
@@ -1003,11 +1004,15 @@ def _render_video(
|
||||
output_name: str,
|
||||
resolution: str = "",
|
||||
bgm_config: dict | None = None,
|
||||
is_preview: bool = False,
|
||||
) -> tuple[Path, float]:
|
||||
"""渲染视频(含配音混音)。
|
||||
|
||||
使用 RenderAdapter 统一渲染入口,复用 BGM/ASR/分辨率/缩略图逻辑。
|
||||
|
||||
Args:
|
||||
is_preview: 是否为预览生成,若是则强制 480p + 低码率
|
||||
|
||||
Returns:
|
||||
(output_path, render_duration)
|
||||
"""
|
||||
@@ -1050,9 +1055,20 @@ def _render_video(
|
||||
|
||||
# 确保输出分辨率配置存在
|
||||
# 优先级:用户指定 > 模板配置 > 默认 1280x720
|
||||
# 预览模式:强制 854x480 + 低码率
|
||||
plan_cfg = virtual_plan.config or {}
|
||||
export_cfg = plan_cfg.get("export", {}) or {}
|
||||
if resolution:
|
||||
if is_preview:
|
||||
# 预览模式强制 480p + 低码率
|
||||
export_cfg["resolution"] = "854x480"
|
||||
export_cfg["bitrate"] = "1M"
|
||||
logger.info(
|
||||
"[task_id=%s] [渲染] 预览模式:强制分辨率=%s, 码率=%s",
|
||||
task_id,
|
||||
"854x480",
|
||||
"1M",
|
||||
)
|
||||
elif resolution:
|
||||
# 用户在 API 调用时指定的分辨率优先级最高
|
||||
export_cfg["resolution"] = resolution
|
||||
elif not export_cfg.get("resolution"):
|
||||
@@ -1301,6 +1317,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
output_name=output_name,
|
||||
resolution=task_info.get("resolution", ""),
|
||||
bgm_config=task_info.get("bgm_config", {}),
|
||||
is_preview=task_info.get("is_preview", False),
|
||||
)
|
||||
|
||||
if gen_task:
|
||||
|
||||
@@ -36,6 +36,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
video_title=getattr(model, "video_title", "") or "",
|
||||
resolution=getattr(model, "resolution", "") or "",
|
||||
bgm_config=dict(getattr(model, "bgm_config", {}) or {}),
|
||||
is_preview=bool(getattr(model, "is_preview", False)),
|
||||
logs=model.logs or "[]",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
@@ -74,6 +75,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
video_title=task.video_title or "",
|
||||
resolution=task.resolution or "",
|
||||
bgm_config=task.bgm_config or {},
|
||||
is_preview=task.is_preview or False,
|
||||
logs=task.logs,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
@@ -238,6 +240,8 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
model.resolution = task.resolution or ""
|
||||
if hasattr(model, "bgm_config"):
|
||||
model.bgm_config = task.bgm_config or {}
|
||||
if hasattr(model, "is_preview"):
|
||||
model.is_preview = task.is_preview or False
|
||||
model.logs = task.logs
|
||||
self.session.commit()
|
||||
return task
|
||||
|
||||
@@ -292,6 +292,7 @@ class GenerationTaskModel(Base):
|
||||
batch_id = Column(String(36), nullable=False, default="", index=True)
|
||||
video_title = Column(String(255), nullable=False, default="")
|
||||
resolution = Column(String(20), nullable=False, default="")
|
||||
is_preview = Column(Boolean, nullable=False, default=False, index=True)
|
||||
bgm_config = Column(JSON, nullable=False, default=dict)
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
logs = Column(Text, nullable=False, default="[]", server_default="[]")
|
||||
|
||||
@@ -26,6 +26,7 @@ class CreateGenerationTaskCommand:
|
||||
bgm_config: dict = field(default_factory=dict)
|
||||
auto_retry_enabled: bool = False
|
||||
auto_retry_max: int = 0
|
||||
is_preview: bool = False
|
||||
|
||||
|
||||
class CreateGenerationTaskUseCase:
|
||||
@@ -56,6 +57,7 @@ class CreateGenerationTaskUseCase:
|
||||
bgm_config=command.bgm_config,
|
||||
auto_retry_enabled=command.auto_retry_enabled,
|
||||
auto_retry_max=command.auto_retry_max,
|
||||
is_preview=command.is_preview,
|
||||
)
|
||||
return self.generation_task_repository.create(task)
|
||||
|
||||
|
||||
@@ -115,6 +115,7 @@ class GenerationTask:
|
||||
video_title: str = ""
|
||||
resolution: str = ""
|
||||
bgm_config: dict = field(default_factory=dict)
|
||||
is_preview: bool = False
|
||||
logs: str = "[]"
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
@@ -140,6 +141,7 @@ class GenerationTask:
|
||||
bgm_config: dict | None = None,
|
||||
auto_retry_enabled: bool = False,
|
||||
auto_retry_max: int = 0,
|
||||
is_preview: bool = False,
|
||||
) -> "GenerationTask":
|
||||
if not project_id.strip() and not template_id.strip():
|
||||
raise ValueError("project_id 或 template_id 至少需要提供一个")
|
||||
@@ -164,6 +166,7 @@ class GenerationTask:
|
||||
bgm_config=dict(bgm_config) if bgm_config else {},
|
||||
auto_retry_enabled=auto_retry_enabled,
|
||||
auto_retry_max=auto_retry_max,
|
||||
is_preview=is_preview,
|
||||
)
|
||||
|
||||
# ── 状态查询 ────────────────────────────────────────────────────────────
|
||||
|
||||
Executable
+561
@@ -0,0 +1,561 @@
|
||||
"""预览生成(Phase 1)单元测试.
|
||||
|
||||
覆盖:
|
||||
- 领域模型 is_preview 字段
|
||||
- CreateGenerationTaskCommand is_preview 字段
|
||||
- UseCase 传递 is_preview
|
||||
- Schema 校验(CreatePreviewGenerationTaskRequest / PreviewGenerationTaskResponse)
|
||||
- 仓储层 _to_domain 兼容旧数据(getattr + 默认值)
|
||||
- 仓储层 create/update 保留 is_preview 字段
|
||||
- 状态流转验证
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
# 确保 app 模块可导入
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.application.generation_tasks import (
|
||||
CreateGenerationTaskCommand,
|
||||
CreateGenerationTaskUseCase,
|
||||
)
|
||||
from packages.domain.generation_task import GenerationTask, GenerationTaskStatus
|
||||
|
||||
# ── 领域模型测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerationTaskIsPreviewField:
|
||||
"""GenerationTask 领域模型 is_preview 字段测试"""
|
||||
|
||||
def test_is_preview_default_false(self):
|
||||
"""默认 is_preview=False"""
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
)
|
||||
assert task.is_preview is False
|
||||
|
||||
def test_is_preview_true_when_specified(self):
|
||||
"""显式指定 is_preview=True"""
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=True,
|
||||
)
|
||||
assert task.is_preview is True
|
||||
|
||||
def test_is_preview_false_when_explicit_false(self):
|
||||
"""显式指定 is_preview=False"""
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=False,
|
||||
)
|
||||
assert task.is_preview is False
|
||||
|
||||
def test_is_preview_with_template_mode(self):
|
||||
"""模板模式下 is_preview 正常工作"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1", "a2"],
|
||||
is_preview=True,
|
||||
)
|
||||
assert task.is_preview is True
|
||||
assert task.template_id == "tmpl1"
|
||||
assert task.asset_ids == ["a1", "a2"]
|
||||
|
||||
def test_is_preview_preserved_in_dataclass(self):
|
||||
"""is_preview 是 dataclass 字段,可以被赋值"""
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
)
|
||||
task.is_preview = True
|
||||
assert task.is_preview is True
|
||||
|
||||
|
||||
# ── UseCase 测试 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateGenerationTaskUseCaseIsPreview:
|
||||
"""CreateGenerationTaskUseCase 中 is_preview 传递测试"""
|
||||
|
||||
def test_command_default_is_preview_false(self):
|
||||
"""CreateGenerationTaskCommand 默认 is_preview=False"""
|
||||
cmd = CreateGenerationTaskCommand(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
)
|
||||
assert cmd.is_preview is False
|
||||
|
||||
def test_command_is_preview_true(self):
|
||||
"""CreateGenerationTaskCommand 设置 is_preview=True"""
|
||||
cmd = CreateGenerationTaskCommand(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=True,
|
||||
)
|
||||
assert cmd.is_preview is True
|
||||
|
||||
def test_use_case_passes_is_preview_to_task(self):
|
||||
"""UseCase 将 is_preview 传递给 GenerationTask"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=True,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.is_preview is True
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_use_case_default_is_preview_false(self):
|
||||
"""UseCase 默认 is_preview=False"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.is_preview is False
|
||||
|
||||
|
||||
# ── Schema 测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreatePreviewGenerationTaskRequest:
|
||||
"""CreatePreviewGenerationTaskRequest schema 校验测试"""
|
||||
|
||||
def test_valid_request_with_asset_ids(self):
|
||||
"""有效请求:template_id + asset_ids"""
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
req = CreatePreviewGenerationTaskRequest(
|
||||
template_id="tmpl_123",
|
||||
asset_ids=["asset_1", "asset_2"],
|
||||
)
|
||||
assert req.template_id == "tmpl_123"
|
||||
assert req.asset_ids == ["asset_1", "asset_2"]
|
||||
assert req.title_ids == []
|
||||
assert req.voice_ids == []
|
||||
assert req.video_title == ""
|
||||
assert req.duration == 0.0
|
||||
assert req.bgm_config == {}
|
||||
|
||||
def test_valid_request_with_title_ids_only(self):
|
||||
"""有效请求:template_id + title_ids(替代 asset_ids)"""
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
req = CreatePreviewGenerationTaskRequest(
|
||||
template_id="tmpl_123",
|
||||
title_ids=["title_1"],
|
||||
)
|
||||
assert req.template_id == "tmpl_123"
|
||||
assert req.title_ids == ["title_1"]
|
||||
|
||||
def test_valid_request_with_voice_ids_only(self):
|
||||
"""有效请求:template_id + voice_ids"""
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
req = CreatePreviewGenerationTaskRequest(
|
||||
template_id="tmpl_123",
|
||||
voice_ids=["voice_1"],
|
||||
)
|
||||
assert req.voice_ids == ["voice_1"]
|
||||
|
||||
def test_missing_template_id_raises(self):
|
||||
"""缺少 template_id 报错"""
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
with pytest.raises(ValidationError, match="template_id"):
|
||||
CreatePreviewGenerationTaskRequest(
|
||||
template_id="",
|
||||
asset_ids=["asset_1"],
|
||||
)
|
||||
|
||||
def test_missing_template_id_not_provided_raises(self):
|
||||
"""完全不提供 template_id 报错"""
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreatePreviewGenerationTaskRequest(
|
||||
asset_ids=["asset_1"],
|
||||
)
|
||||
|
||||
def test_empty_asset_ids_raises(self):
|
||||
"""asset_ids/title_ids/voice_ids 全空报错"""
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
with pytest.raises(ValidationError, match="asset_ids"):
|
||||
CreatePreviewGenerationTaskRequest(
|
||||
template_id="tmpl_123",
|
||||
asset_ids=[],
|
||||
title_ids=[],
|
||||
voice_ids=[],
|
||||
)
|
||||
|
||||
def test_request_with_all_fields(self):
|
||||
"""所有字段都设置的请求"""
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
|
||||
req = CreatePreviewGenerationTaskRequest(
|
||||
template_id="tmpl_123",
|
||||
asset_ids=["a1", "a2"],
|
||||
title_ids=["t1"],
|
||||
voice_ids=["v1"],
|
||||
video_title="测试预览视频",
|
||||
duration=30.0,
|
||||
video_ratio="9:16",
|
||||
bgm_config={"enabled": True, "volume": 0.5},
|
||||
)
|
||||
assert req.video_title == "测试预览视频"
|
||||
assert req.duration == 30.0
|
||||
assert req.video_ratio == "9:16"
|
||||
assert req.bgm_config["enabled"] is True
|
||||
assert req.bgm_config["volume"] == 0.5
|
||||
|
||||
|
||||
class TestPreviewGenerationTaskResponse:
|
||||
"""PreviewGenerationTaskResponse schema 测试"""
|
||||
|
||||
def test_pending_state_response(self):
|
||||
"""pending 状态的响应"""
|
||||
from app.schemas.generation_task import PreviewGenerationTaskResponse
|
||||
|
||||
resp = PreviewGenerationTaskResponse(
|
||||
task_id="task_123",
|
||||
status="pending",
|
||||
progress=0.0,
|
||||
)
|
||||
assert resp.task_id == "task_123"
|
||||
assert resp.status == "pending"
|
||||
assert resp.progress == 0.0
|
||||
assert resp.is_preview is True
|
||||
assert resp.resolution == ""
|
||||
assert resp.video_url == ""
|
||||
assert resp.duration == 0.0
|
||||
assert resp.file_size == 0
|
||||
assert resp.clip_count == 0
|
||||
assert resp.error_message == ""
|
||||
|
||||
def test_completed_state_response(self):
|
||||
"""completed 状态的响应"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from app.schemas.generation_task import PreviewGenerationTaskResponse
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
resp = PreviewGenerationTaskResponse(
|
||||
task_id="task_123",
|
||||
status="completed",
|
||||
progress=100.0,
|
||||
is_preview=True,
|
||||
resolution="854x480",
|
||||
video_url="https://example.com/preview.mp4",
|
||||
duration=30.5,
|
||||
file_size=5_000_000,
|
||||
clip_count=5,
|
||||
transition_count=4,
|
||||
material_usage={"videos": 5, "images": 2},
|
||||
created_at=now,
|
||||
started_at=now,
|
||||
finished_at=now,
|
||||
generate_duration=12.5,
|
||||
)
|
||||
assert resp.status == "completed"
|
||||
assert resp.progress == 100.0
|
||||
assert resp.resolution == "854x480"
|
||||
assert resp.video_url == "https://example.com/preview.mp4"
|
||||
assert resp.duration == 30.5
|
||||
assert resp.file_size == 5_000_000
|
||||
assert resp.clip_count == 5
|
||||
assert resp.transition_count == 4
|
||||
assert resp.generate_duration == 12.5
|
||||
|
||||
def test_failed_state_response(self):
|
||||
"""failed 状态的响应"""
|
||||
from app.schemas.generation_task import PreviewGenerationTaskResponse
|
||||
|
||||
resp = PreviewGenerationTaskResponse(
|
||||
task_id="task_123",
|
||||
status="failed",
|
||||
progress=30.0,
|
||||
error_message="渲染失败:素材格式不支持",
|
||||
)
|
||||
assert resp.status == "failed"
|
||||
assert resp.error_message == "渲染失败:素材格式不支持"
|
||||
assert resp.video_url == ""
|
||||
|
||||
|
||||
# ── 仓储层兼容测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRepositoryIsPreviewCompatibility:
|
||||
"""仓储层 is_preview 向后兼容测试"""
|
||||
|
||||
def test_to_domain_with_is_preview_true(self):
|
||||
"""新数据 is_preview=True 时正确映射"""
|
||||
mock_model = MagicMock()
|
||||
mock_model.id = "task_123"
|
||||
mock_model.project_id = "proj1"
|
||||
mock_model.strategy_id = ""
|
||||
mock_model.asset_library_id = "lib1"
|
||||
mock_model.voice_library_id = ""
|
||||
mock_model.template_id = "tmpl1"
|
||||
mock_model.asset_ids = ["a1"]
|
||||
mock_model.title_ids = []
|
||||
mock_model.voice_ids = []
|
||||
mock_model.status = "pending"
|
||||
mock_model.progress = 0.0
|
||||
mock_model.result_count = 0
|
||||
mock_model.error_message = ""
|
||||
mock_model.error_info = {}
|
||||
mock_model.retry_count = 0
|
||||
mock_model.auto_retry_enabled = False
|
||||
mock_model.auto_retry_max = 0
|
||||
mock_model.started_at = None
|
||||
mock_model.completed_at = None
|
||||
mock_model.created_by_user_id = "user1"
|
||||
mock_model.source_edit_plan_id = None
|
||||
mock_model.asset_select_mode = ""
|
||||
mock_model.batch_id = ""
|
||||
mock_model.video_title = ""
|
||||
mock_model.resolution = "854x480"
|
||||
mock_model.bgm_config = {}
|
||||
mock_model.is_preview = True
|
||||
mock_model.logs = "[]"
|
||||
mock_model.created_at = datetime.now(timezone.utc)
|
||||
mock_model.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import _to_domain
|
||||
|
||||
result = _to_domain(mock_model)
|
||||
assert result.is_preview is True
|
||||
assert result.resolution == "854x480"
|
||||
|
||||
def test_to_domain_with_is_preview_false(self):
|
||||
"""is_preview=False 时正确映射"""
|
||||
mock_model = MagicMock()
|
||||
mock_model.id = "task_123"
|
||||
mock_model.project_id = "proj1"
|
||||
mock_model.strategy_id = ""
|
||||
mock_model.asset_library_id = "lib1"
|
||||
mock_model.voice_library_id = ""
|
||||
mock_model.template_id = ""
|
||||
mock_model.asset_ids = []
|
||||
mock_model.title_ids = []
|
||||
mock_model.voice_ids = []
|
||||
mock_model.status = "pending"
|
||||
mock_model.progress = 0.0
|
||||
mock_model.result_count = 0
|
||||
mock_model.error_message = ""
|
||||
mock_model.error_info = {}
|
||||
mock_model.retry_count = 0
|
||||
mock_model.auto_retry_enabled = False
|
||||
mock_model.auto_retry_max = 0
|
||||
mock_model.started_at = None
|
||||
mock_model.completed_at = None
|
||||
mock_model.created_by_user_id = "user1"
|
||||
mock_model.source_edit_plan_id = None
|
||||
mock_model.asset_select_mode = ""
|
||||
mock_model.batch_id = ""
|
||||
mock_model.video_title = ""
|
||||
mock_model.resolution = ""
|
||||
mock_model.bgm_config = {}
|
||||
mock_model.is_preview = False
|
||||
mock_model.logs = "[]"
|
||||
mock_model.created_at = datetime.now(timezone.utc)
|
||||
mock_model.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import _to_domain
|
||||
|
||||
result = _to_domain(mock_model)
|
||||
assert result.is_preview is False
|
||||
|
||||
def test_repository_create_includes_is_preview(self):
|
||||
"""repository.create() 包含 is_preview 字段"""
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base, GenerationTaskModel
|
||||
|
||||
# 使用内存数据库
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
SessionLocal = sessionmaker(bind=engine)
|
||||
session = SessionLocal()
|
||||
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=True,
|
||||
resolution="854x480",
|
||||
)
|
||||
repo.create(task)
|
||||
|
||||
# 直接查 model 验证
|
||||
model = session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first()
|
||||
assert model is not None
|
||||
assert model.is_preview is True
|
||||
assert model.resolution == "854x480"
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
def test_repository_update_preserves_is_preview(self):
|
||||
"""repository.update() 保留 is_preview 字段"""
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base, GenerationTaskModel
|
||||
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
SessionLocal = sessionmaker(bind=engine)
|
||||
session = SessionLocal()
|
||||
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
is_preview=True,
|
||||
)
|
||||
repo.create(task)
|
||||
|
||||
# 更新任务状态
|
||||
task.mark_processing()
|
||||
repo.update(task)
|
||||
|
||||
model = session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first()
|
||||
assert model is not None
|
||||
assert model.is_preview is True # is_preview 应该保持不变
|
||||
assert model.status == "running"
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
# ── 状态流转测试 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPreviewTaskStatusFlow:
|
||||
"""预览任务状态流转测试"""
|
||||
|
||||
def test_pending_to_running(self):
|
||||
"""pending → running"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1"],
|
||||
is_preview=True,
|
||||
)
|
||||
assert task.status == GenerationTaskStatus.PENDING
|
||||
assert task.is_preview is True
|
||||
|
||||
task.mark_processing()
|
||||
assert task.status == GenerationTaskStatus.RUNNING
|
||||
assert task.started_at is not None
|
||||
|
||||
def test_running_to_completed(self):
|
||||
"""running → completed"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1"],
|
||||
is_preview=True,
|
||||
)
|
||||
task.mark_processing()
|
||||
task.mark_completed()
|
||||
|
||||
assert task.status == GenerationTaskStatus.COMPLETED
|
||||
assert task.progress == 100.0
|
||||
assert task.completed_at is not None
|
||||
assert task.is_preview is True
|
||||
|
||||
def test_running_to_failed(self):
|
||||
"""running → failed"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1"],
|
||||
is_preview=True,
|
||||
)
|
||||
task.mark_processing()
|
||||
task.mark_failed("渲染错误")
|
||||
|
||||
assert task.status == GenerationTaskStatus.FAILED
|
||||
assert task.error_message == "渲染错误"
|
||||
assert task.is_preview is True
|
||||
|
||||
def test_pending_to_cancelled(self):
|
||||
"""pending → cancelled"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1"],
|
||||
is_preview=True,
|
||||
)
|
||||
task.mark_cancelled()
|
||||
|
||||
assert task.status == GenerationTaskStatus.CANCELLED
|
||||
assert task.is_preview is True
|
||||
|
||||
def test_preview_resolution_is_480p(self):
|
||||
"""预览任务分辨率为 854x480"""
|
||||
task = GenerationTask.create(
|
||||
project_id="",
|
||||
asset_library_id="",
|
||||
template_id="tmpl1",
|
||||
asset_ids=["a1"],
|
||||
resolution="854x480",
|
||||
is_preview=True,
|
||||
)
|
||||
assert task.resolution == "854x480"
|
||||
assert task.is_preview is True
|
||||
|
||||
def test_non_preview_task_default_false(self):
|
||||
"""非预览任务 is_preview 默认 False"""
|
||||
task = GenerationTask.create(
|
||||
project_id="proj1",
|
||||
asset_library_id="lib1",
|
||||
)
|
||||
assert task.is_preview is False
|
||||
Reference in New Issue
Block a user