diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index c40fdf343..64e5b16ca 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -23,6 +23,7 @@ from app.dependencies import ( get_generation_task_repository, ) from app.schemas.generation_task import ( + BatchPreviewGenerationTaskResponse, CreatePreviewGenerationTaskRequest, PreviewGenerationTaskResponse, ) @@ -193,11 +194,19 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG if started_at and completed_at: generate_duration = (completed_at - started_at).total_seconds() + title_cfg = getattr(task, "title_config", None) + title_cfg = title_cfg if isinstance(title_cfg, dict) else {} + extra_meta = getattr(task, "extra_meta", None) + extra_meta = extra_meta if isinstance(extra_meta, dict) else {} + voice_library_id = getattr(task, "voice_library_id", "") or "" + if not isinstance(voice_library_id, str): + voice_library_id = str(voice_library_id) if voice_library_id else "" 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)), + variant_index=int(extra_meta.get("variant_index", 0) or 0), resolution=getattr(task, "resolution", "") or "", video_url=video_url, duration=duration, @@ -206,6 +215,8 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG transition_count=transition_count, material_usage=material_usage, error_message=task.error_message or "", + title_text=str(title_cfg.get("text", "") or ""), + voice_library_id=voice_library_id, created_at=task.created_at, started_at=started_at, finished_at=completed_at, @@ -213,45 +224,95 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG ) -@router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201) +def _resolve_preview_edit_plan_id( + *, + request: CreatePreviewGenerationTaskRequest, + task, + db: Session, + user_id: str, +) -> str: + """确定任务关联的编辑计划ID:优先前端传入,否则按 template_id+user 兜底查找。""" + if task.source_edit_plan_id: + return task.source_edit_plan_id + if not request.template_id: + return "" + try: + from packages.adapters.sqlalchemy_impl.edit_plan_repository import ( + SQLAlchemyEditPlanRepository, + ) + + _plan_repo = SQLAlchemyEditPlanRepository(db) + _plans = _plan_repo.list_by_template(request.template_id, limit=20) + for _p in _plans: + if (_p.created_by_user_id or "") == user_id: + logger.info( + "[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s", + task.id, + _p.id, + ) + return _p.id + except Exception: + logger.warning( + "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s", + task.id, + exc_info=True, + ) + return "" + + +def _variant_value(values: list[str], index: int, fallback: str = "") -> str: + """从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。""" + if not values: + return fallback + if len(values) == 1: + return values[0] + return values[index] if index < len(values) else fallback + + +@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, 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), db: Session = Depends(get_db_session), asset_repo=Depends(get_asset_repository), -) -> PreviewGenerationTaskResponse: - """创建预览生成任务。 +) -> BatchPreviewGenerationTaskResponse: + """创建预览生成任务(支持批量)。 - 预览渲染品质与正式生成一致(1080p, CRF 23, medium preset),确认生成时可直接复用预览产物。 - - Args: - request: 预览任务创建请求(template_id + asset_ids 等) + preview_count=1 时行为与旧版完全一致(创建 1 个任务); + preview_count=N 时一次创建 N 个独立变体任务: + - 每个变体克隆独立编辑计划(独立 clips、独立随机素材起点),N 个预览内容互不相同 + - 每个变体拥有独立 task_id / 状态 / 预览视频 URL,前端按 task_id 分别轮询 + - 标题样式(font/color/position 等)全局共用;标题文字/配音/封面可按变体独立 + (titles[] / voice_library_ids[] / cover_urls[],长度1=共用,长度N=独立) Returns: - 201 + 预览任务详情 + 201 + 变体任务数组 {items: [...], total: N} """ user_id = authenticated_user.user.id + count = max(1, request.preview_count) logger.info( "[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d", user_id, request.template_id, len(request.asset_ids), - request.preview_count, + count, ) - # 预检查队列限流 + # 预检查队列限流(按变体总数计) 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) + if user_pending + count > USER_PENDING_LIMIT: + raise UserPendingLimitExceeded( + user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT + ) + if global_pending + count > GLOBAL_PENDING_LIMIT: + raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT) except UserPendingLimitExceeded as e: raise HTTPException( status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", + detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待后再提交", ) from e except GlobalQueueFull as e: raise HTTPException( @@ -273,14 +334,11 @@ def create_preview_generation_task( w, h = int(parts[0]), int(parts[1]) base = 1920 if w < h: - # 竖屏 output_width = round(base * w / h) output_height = base else: - # 横屏 output_width = base output_height = round(base * h / w) - # 对齐到偶数 output_width = output_width - output_width % 2 output_height = output_height - output_height % 2 except (ValueError, ZeroDivisionError): @@ -289,42 +347,71 @@ def create_preview_generation_task( logger.info( "[预览生成] 分辨率: video_ratio=%s → %s (%dx%d)", - video_ratio, resolution, output_width, output_height, + video_ratio, + resolution, + output_width, + output_height, ) - # 从模板读取 editing_mode / mode 作为 strategy_id(渲染 pipeline 的 mode 参数) strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id) - - title_config = request.title_config or {} + base_title_config = request.title_config or {} use_case = CreateGenerationTaskUseCase(generation_task_repository) + # ── 预创建第一个任务,仅用于解析源编辑计划(不落库为最终任务)── + # 先创建一个临时任务拿到 task 对象上下文,实际 N 个任务在循环中统一创建; + # 为保持与旧版一致的源 plan 解析逻辑,先创建任务0、解析源 plan, + # 再预克隆 N 个变体 plan,最后重建任务关联。 + # 简化实现:直接创建全部任务,plan 关联在创建后、入队前完成。 + + created_tasks: list = [] + variant_plan_ids: list[str] = [] # 每个变体最终关联的 plan_id(按变体顺序) + try: - task = use_case.execute( - CreateGenerationTaskCommand( - project_id="", - asset_library_id="", - strategy_id=strategy_id, - voice_library_id=request.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=request.source_edit_plan_id, - asset_select_mode="", - batch_id="", - video_title=request.video_title, - resolution=resolution, - bgm_config=request.bgm_config or {}, - auto_retry_enabled=False, - auto_retry_max=0, - is_preview=True, - title_config=title_config, - output_width=output_width, - output_height=output_height, + for variant_index in range(count): + # 变体独立标题文字:titles[] 覆盖 title_config.text + variant_title_text = _variant_value(request.titles, variant_index, "") + variant_title_config = dict(base_title_config) + if variant_title_text.strip(): + variant_title_config["text"] = variant_title_text.strip() + + # 变体独立配音 + variant_voice_library_id = _variant_value( + request.voice_library_ids, variant_index, request.voice_library_id ) - ) + + task = use_case.execute( + CreateGenerationTaskCommand( + project_id="", + asset_library_id="", + strategy_id=strategy_id, + voice_library_id=variant_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=request.source_edit_plan_id, + asset_select_mode="", + batch_id="", + video_title=request.video_title, + resolution=resolution, + bgm_config=request.bgm_config or {}, + auto_retry_enabled=False, + auto_retry_max=0, + is_preview=True, + title_config=variant_title_config, + output_width=output_width, + output_height=output_height, + ) + ) + task.extra_meta["variant_index"] = variant_index + + # 解析源编辑计划(前端传入或按模板兜底查找) + source_plan_id = _resolve_preview_edit_plan_id(request=request, task=task, db=db, user_id=user_id) + task.source_edit_plan_id = source_plan_id + generation_task_repository.update(task) + created_tasks.append(task) except ValueError as e: logger.warning("[预览生成] 创建失败: %s", e) raise HTTPException(status_code=400, detail=str(e)) from e @@ -332,93 +419,121 @@ def create_preview_generation_task( logger.error("[预览生成] 创建失败: %s", e, exc_info=True) raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e - # 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找 - if not task.source_edit_plan_id and request.template_id: - try: - from packages.adapters.sqlalchemy_impl.edit_plan_repository import ( - SQLAlchemyEditPlanRepository, - ) - - _plan_repo = SQLAlchemyEditPlanRepository(db) - _plans = _plan_repo.list_by_template(request.template_id, limit=20) - for _p in _plans: - if (_p.created_by_user_id or "") == user_id: - task.source_edit_plan_id = _p.id - generation_task_repository.update(task) - logger.info( - "[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s", - task.id, - _p.id, - ) - break - except Exception: - logger.warning( - "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s", - task.id, - exc_info=True, - ) - - # 每条预览都关联独立克隆 plan:多预览前端为 N 次并发调用,若共用同一 plan - # 则 N 条预览片段完全相同;克隆时片段起点按持久化历史区间重算(含受控复用), - # 保证各预览版本内容不同 - if task.source_edit_plan_id: + # ── 克隆独立变体 plan:N 个预览全部克隆(预览不污染源 plan)── + # 源 plan 不存在(无编辑历史)时各任务走自身随机选片流程,不克隆。 + source_plan_id = created_tasks[0].source_edit_plan_id if created_tasks else "" + if source_plan_id: try: from app.services.edit_plan_service import EditPlanService _plan_svc = EditPlanService(db) - _preview_plan = _plan_svc.clone_plan_for_variant( - task.source_edit_plan_id, - created_by_user_id=user_id, - name_suffix="预览变体", - ) - task.source_edit_plan_id = _preview_plan.id - generation_task_repository.update(task) - logger.info( - "[预览生成] 预览关联独立克隆 plan: task_id=%s clone_plan_id=%s", - task.id, - _preview_plan.id, - ) - except Exception as clone_err: - # 不退回共用原 plan(否则多条预览内容相同,违反去重诉求): - # 标记任务失败并中断,前端可重新发起预览 - logger.error( - "[预览生成] 克隆预览变体 plan 失败,任务标记失败: task_id=%s error=%s", - task.id, - clone_err, - exc_info=True, - ) - _mark_task_failed(generation_task_repository, task, "预览变体计划创建失败") + for variant_index in range(count): + last_err: Exception | None = None + variant_plan = None + for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动 + try: + variant_plan = _plan_svc.clone_plan_for_variant( + source_plan_id, + created_by_user_id=user_id, + name_suffix=f"预览变体{variant_index + 1}" if count > 1 else "预览变体", + ) + break + except Exception as clone_err: # noqa: PERF203 + last_err = clone_err + logger.warning( + "[预览生成] 克隆变体 plan 失败(尝试%d/2): variant=%d error=%s", + _attempt + 1, + variant_index, + clone_err, + exc_info=True, + ) + if variant_plan is None: + logger.error( + "[预览生成] 克隆预览变体 plan 重试仍失败: variant=%d source=%s", + variant_index, + source_plan_id, + exc_info=last_err, + ) + # 标记已创建任务失败 + for t in created_tasks: + _mark_task_failed(generation_task_repository, t, "预览变体计划创建失败") + raise HTTPException( + status_code=500, + detail="创建预览任务失败:无法生成独立剪辑计划,请重试", + ) from last_err + variant_plan_ids.append(variant_plan.id) + except HTTPException: + raise + except Exception as e: + logger.error("[预览生成] 克隆变体 plan 异常: %s", e, exc_info=True) + for t in created_tasks: + _mark_task_failed(generation_task_repository, t, "预览变体计划创建失败") raise HTTPException( status_code=500, detail="创建预览任务失败:无法生成独立剪辑计划,请重试", - ) from clone_err + ) from e - # 入队执行;若入队失败则标记任务为 failed 避免僵尸数据 - 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) - _mark_task_failed(generation_task_repository, task, "任务入队失败") - raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试") - except UserPendingLimitExceeded as e: - _mark_task_failed(generation_task_repository, task, "待处理任务超限") - raise HTTPException( - status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", - ) from None - except GlobalQueueFull: - _mark_task_failed(generation_task_repository, task, "系统队列已满") - raise HTTPException( - status_code=503, - detail="系统繁忙,请稍后再试", - ) from None + # 关联变体 plan 并回写标题配置 + for variant_index, task in enumerate(created_tasks): + if variant_plan_ids: + task.source_edit_plan_id = variant_plan_ids[variant_index] + generation_task_repository.update(task) + # 回写变体标题到 plan config(worker 渲染时从 plan 读取 title 配置) + if task.source_edit_plan_id and (task.title_config or {}).get("text", "").strip(): + try: + from app.api.routes.generation_tasks import _writeback_edit_plan_config - return _to_preview_response(task) + _writeback_edit_plan_config( + plan_id=task.source_edit_plan_id, + task_id=task.id, + title_config=task.title_config, + db=db, + ) + except Exception: + logger.warning( + "[预览生成] 回写标题配置失败(不影响主流程): task_id=%s", + task.id, + exc_info=True, + ) + + # ── 入队 ── + responses: list[PreviewGenerationTaskResponse] = [] + for variant_index, task in enumerate(created_tasks): + try: + enqueued = safe_enqueue_generation_task( + task, + generation_task_repository, + user_id=user_id, + log_prefix=f"[预览生成][变体{variant_index + 1}]", + log_task_status=True, + ) + if not enqueued: + logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id) + _mark_task_failed(generation_task_repository, task, "任务入队失败") + except UserPendingLimitExceeded: + _mark_task_failed(generation_task_repository, task, "待处理任务超限") + except GlobalQueueFull: + _mark_task_failed(generation_task_repository, task, "系统队列已满") + except Exception: + logger.exception("[预览生成] 入队异常: task_id=%s", task.id) + _mark_task_failed(generation_task_repository, task, "任务入队异常") + # enqueue 会原地更新 task 状态/进度,直接用 task 构造响应 + responses.append(_to_preview_response(task)) + + # 队列满/限流时若全部失败,返回明确错误码 + if all(r.status == "failed" for r in responses): + first_err = next((r.error_message for r in responses if r.error_message), "") + if "待处理任务" in first_err: + raise HTTPException(status_code=429, detail=first_err or "待处理任务超限") + if "队列" in first_err: + raise HTTPException(status_code=503, detail=first_err or "系统繁忙,请稍后再试") + + logger.info( + "[预览生成] 创建完成: %d 个变体任务, task_ids=%s", + len(responses), + [r.task_id for r in responses], + ) + return BatchPreviewGenerationTaskResponse(items=responses, total=len(responses)) @router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index d5fbad1f1..32391a34d 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -47,6 +47,15 @@ logger = logging.getLogger(__name__) router = APIRouter() +def _variant_value(values: list[str], index: int, fallback: str = "") -> str: + """从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。""" + if not values: + return fallback + if len(values) == 1: + return values[0] + return values[index] if index < len(values) else fallback + + def _to_generation_task_response(task) -> GenerationTaskResponse: return GenerationTaskResponse( id=task.id, @@ -472,12 +481,21 @@ def create_generation_task( if task_index > 0 and variant_plan_ids: effective_plan_id = variant_plan_ids[task_index - 1] + # 变体级独立配置:titles[]/voice_library_ids[]/cover_urls[] + # 长度1=所有变体共用,长度=count=每个变体独立,空数组=回退单值字段 + variant_title_text = _variant_value(request.titles, task_index, "") + variant_title_config = dict(request.title_config or {}) + if variant_title_text.strip(): + variant_title_config["text"] = variant_title_text.strip() + variant_voice_library_id = _variant_value(request.voice_library_ids, task_index, request.voice_library_id) + variant_cover_url = _variant_value(request.cover_urls, task_index, request.cover_url) + task = use_case.execute( CreateGenerationTaskCommand( project_id=project_id, asset_library_id=asset_library_id, strategy_id=effective_strategy_id, - voice_library_id=request.voice_library_id, + voice_library_id=variant_voice_library_id, template_id=request.template_id, asset_ids=resolved_asset_ids, title_ids=request.title_ids, @@ -495,10 +513,12 @@ def create_generation_task( source_task_id=request.source_task_id, output_width=request.output_width, output_height=request.output_height, - cover_url=request.cover_url, - title_config=request.title_config or {}, + cover_url=variant_cover_url, + title_config=variant_title_config, ) ) + # 变体序号写入 extra_meta(响应/排查时可辨识) + task.extra_meta["variant_index"] = task_index try: # 兜底关联编辑计划:前端未传 source_edit_plan_id 时, # 通过 template_id + user_id 在 DB 层直接查找最新的 plan。 @@ -533,13 +553,13 @@ def create_generation_task( # 回写 plan.config:必须在 enqueue 之前执行, # 确保 worker 读取 plan 时 config 中已包含 generation_task_id。 - # 只在首个任务时回写一次,避免批量生成时循环覆盖。 + # 批量场景下每个变体关联独立 plan,需各自回写自己的变体标题配置。 _effective_plan_id = task.source_edit_plan_id - if _effective_plan_id and len(created_tasks) == 0: + if _effective_plan_id: _writeback_edit_plan_config( plan_id=_effective_plan_id, task_id=task.id, - title_config=request.title_config, + title_config=variant_title_config, db=db, ) diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index bead763e3..15f55d895 100755 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -25,6 +25,12 @@ class CreateGenerationTaskRequest(BaseModel): asset_library_id: str = "" strategy_id: str = "" voice_library_id: str = "" + # ── 多变体独立配音(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立配音;空数组 = 回退 voice_library_id + voice_library_ids: list[str] = Field( + default_factory=list, + description="各变体独立配音素材库ID数组:长度1=共用,长度=count=独立。为空时回退 voice_library_id", + ) created_by_user_id: str = "" # ── 模板模式新增字段 ── template_id: str = "" @@ -75,6 +81,27 @@ class CreateGenerationTaskRequest(BaseModel): output_width: int = Field(default=1280, description="输出视频宽度") output_height: int = Field(default=720, description="输出视频高度") cover_url: str = Field(default="", description="封面图片 URL") + # ── 多变体独立封面(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立封面;空数组 = 回退 cover_url + cover_urls: list[str] = Field( + default_factory=list, + description="各变体独立封面URL数组:长度1=共用,长度=count=独立。为空时回退 cover_url", + ) + # ── 多变体独立标题文字(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立标题文字;空数组 = 使用 title_config.text + titles: list[str] = Field( + default_factory=list, + description="各变体独立标题文字数组:长度1=共用,长度=count=独立。为空时使用 title_config.text", + ) + + @model_validator(mode="after") + def _check_variant_arrays(self) -> "CreateGenerationTaskRequest": + """变体数组字段长度校验:空数组(回退单值)、长度 1(共用)、或长度 = count(独立)。""" + for name in ("voice_library_ids", "cover_urls", "titles"): + arr = getattr(self, name) + if arr and len(arr) != 1 and len(arr) != self.count: + raise ValueError(f"{name} 长度必须为 1(共用)或 {self.count}(与 count 一致),当前为 {len(arr)}") + return self @model_validator(mode="after") def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest": @@ -185,8 +212,33 @@ class CreatePreviewGenerationTaskRequest(BaseModel): ) title_config: dict = Field( default_factory=dict, - description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow", + description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow。N个变体时样式全局共用", ) + # ── 多变体独立配置(preview_count > 1)── + # 长度 1 = 所有变体共用;长度 = preview_count = 每个变体独立;空数组 = 回退单值字段 + titles: list[str] = Field( + default_factory=list, + description="各变体独立标题文字数组:长度1=共用,长度=preview_count=独立。为空时使用 title_config.text", + ) + voice_library_ids: list[str] = Field( + default_factory=list, + description="各变体独立配音素材库ID数组:长度1=共用,长度=preview_count=独立。为空时回退 voice_library_id", + ) + cover_urls: list[str] = Field( + default_factory=list, + description="各变体独立封面URL数组:长度1=共用,长度=preview_count=独立(预览阶段通常为空)", + ) + + @model_validator(mode="after") + def _check_variant_arrays(self) -> "CreatePreviewGenerationTaskRequest": + """变体数组字段长度校验:空数组(回退单值)、长度 1(共用)、或长度 = preview_count(独立)。""" + for name in ("titles", "voice_library_ids", "cover_urls"): + arr = getattr(self, name) + if arr and len(arr) != 1 and len(arr) != self.preview_count: + raise ValueError( + f"{name} 长度必须为 1(共用)或 {self.preview_count}(与 preview_count 一致),当前为 {len(arr)}" + ) + return self @model_validator(mode="after") def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest": @@ -202,7 +254,7 @@ class CreatePreviewGenerationTaskRequest(BaseModel): class PreviewGenerationTaskResponse(BaseModel): - """预览生成任务响应。 + """单个预览变体任务响应。 包含任务状态、进度、分辨率、生成结果 URL 等关键字段。 """ @@ -211,6 +263,7 @@ class PreviewGenerationTaskResponse(BaseModel): status: str progress: float is_preview: bool = True + variant_index: int = 0 resolution: str = "" video_url: str = "" duration: float = 0.0 @@ -219,7 +272,21 @@ class PreviewGenerationTaskResponse(BaseModel): transition_count: int = 0 material_usage: dict = Field(default_factory=dict) error_message: str = "" + title_text: str = "" + voice_library_id: str = "" created_at: datetime | None = None started_at: datetime | None = None finished_at: datetime | None = None generate_duration: float = 0.0 + + +class BatchPreviewGenerationTaskResponse(BaseModel): + """批量预览任务响应:preview_count=N 时返回 N 个独立变体任务。 + + - items: 变体任务数组,按 variant_index 顺序排列,每个含独立 task_id/状态/预览视频URL + - total: 变体总数(= preview_count) + - 前端按 items[i].task_id 分别轮询 GET /preview/{task_id} 获取进度与结果 + """ + + items: list[PreviewGenerationTaskResponse] + total: int diff --git a/tests/unit/test_1677_batch_variants.py b/tests/unit/test_1677_batch_variants.py new file mode 100644 index 000000000..bf67b5b62 --- /dev/null +++ b/tests/unit/test_1677_batch_variants.py @@ -0,0 +1,507 @@ +"""Issue #1677 多视频批量生成 — 变体独立配置与批量预览/批量生成测试。 + +覆盖: +- 批量预览:preview_count=N 一次创建 N 个独立任务,返回变体数组 +- 变体克隆链路:N 个预览/正式任务各自关联独立克隆 plan +- 变体独立配置:titles[]/voice_library_ids[]/cover_urls[] 按变体注入 +- 长度校验:数组长度必须为 1 或 N(共用或独立),非法长度报错 +- N=1 向后兼容:旧字段单值行为不变 +""" + +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest +from app.core.task_enqueue import GlobalQueueFull, UserPendingLimitExceeded +from app.schemas.generation_task import ( + BatchPreviewGenerationTaskResponse, + CreateGenerationTaskRequest, + CreatePreviewGenerationTaskRequest, +) + +from packages.domain import GenerationTask +from packages.domain.generation_task import GenerationTaskStatus + +# ════════════════════════════════════════════════════════════════════════════ +# 辅助构造 +# ════════════════════════════════════════════════════════════════════════════ + + +def _make_user(user_id="test_user_001"): + mock_user = MagicMock() + mock_user.id = user_id + auth = MagicMock() + auth.user = mock_user + return auth + + +def _make_task(task_id=None, status=GenerationTaskStatus.PENDING, source_plan_id=None): + task = GenerationTask.create( + project_id="", + asset_library_id="", + template_id="tpl_001", + asset_ids=["asset_1"], + ) + if task_id: + task.id = task_id + task.status = status + task.is_preview = True + task.source_edit_plan_id = source_plan_id or "" + task.voice_library_id = "" + task.title_config = {} + task.cover_url = "" + return task + + +def _make_preview_request(**kwargs): + defaults = { + "template_id": "tpl_001", + "asset_ids": ["asset_1", "asset_2"], + } + defaults.update(kwargs) + return CreatePreviewGenerationTaskRequest(**defaults) + + +def _repo_mock(): + repo = MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + repo.get.side_effect = lambda tid: None + return repo + + +# ════════════════════════════════════════════════════════════════════════════ +# Schema 校验:变体数组长度 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestVariantArrayValidation: + """变体数组字段长度校验。""" + + def test_preview_titles_length_matches_count(self): + """titles 长度 = preview_count 合法""" + req = _make_preview_request(preview_count=3, titles=["标题A", "标题B", "标题C"]) + assert len(req.titles) == 3 + + def test_preview_titles_single_shared(self): + """titles 长度 1 = 所有变体共用,合法""" + req = _make_preview_request(preview_count=3, titles=["共用标题"]) + assert req.titles == ["共用标题"] + + def test_preview_titles_wrong_length_raises(self): + """titles 长度 2 与 preview_count=3 不匹配 → 报错""" + with pytest.raises(ValueError, match="titles"): + _make_preview_request(preview_count=3, titles=["A", "B"]) + + def test_preview_voice_ids_wrong_length_raises(self): + """voice_library_ids 长度非法 → 报错""" + from pydantic import ValidationError + + with pytest.raises(ValidationError, match="voice_library_ids"): + _make_preview_request(preview_count=4, voice_library_ids=["v1", "v2"]) + + def test_preview_empty_arrays_ok(self): + """空数组(回退单值字段)合法""" + req = _make_preview_request(preview_count=3) + assert req.titles == [] + assert req.voice_library_ids == [] + assert req.cover_urls == [] + + def test_generation_titles_length_matches_count(self): + """正式生成 titles 长度 = count 合法""" + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + titles=["A", "B", "C"], + ) + assert len(req.titles) == 3 + + def test_generation_arrays_wrong_length_raises(self): + """正式生成 cover_urls 长度与 count 不匹配 → 报错""" + from pydantic import ValidationError + + with pytest.raises(ValidationError, match="cover_urls"): + CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + cover_urls=["c1", "c2"], + ) + + def test_generation_single_count_no_arrays(self): + """N=1 且不传数组:完全旧行为""" + req = CreateGenerationTaskRequest(template_id="tpl_1", asset_ids=["a1"]) + assert req.count == 1 + assert req.titles == [] + assert req.voice_library_ids == [] + assert req.cover_urls == [] + + +# ════════════════════════════════════════════════════════════════════════════ +# 批量预览路由 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestBatchPreviewRoute: + """POST /preview 批量变体。""" + + def test_preview_count_1_returns_single_item_array(self): + """N=1 返回 items 长度 1 的批量响应(结构统一)""" + from app.api.routes.generation_preview import create_preview_generation_task + + task = _make_task(task_id="task_1") + repo = _repo_mock() + 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( + _make_preview_request(preview_count=1), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert isinstance(resp, BatchPreviewGenerationTaskResponse) + assert resp.total == 1 + assert len(resp.items) == 1 + assert resp.items[0].task_id == "task_1" + assert resp.items[0].variant_index == 0 + + def test_preview_count_3_creates_three_independent_tasks(self): + """N=3 创建 3 个独立任务,返回 3 个变体,task_id 各不相同""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + resp = create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert resp.total == 3 + task_ids = [item.task_id for item in resp.items] + assert task_ids == ["task_0", "task_1", "task_2"] + assert len(set(task_ids)) == 3 + for i, item in enumerate(resp.items): + assert item.variant_index == i + + def test_preview_count_3_clones_three_variant_plans(self): + """有源 plan 时,N=3 克隆 3 个独立变体 plan(预览全部克隆,不用源 plan)""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] + repo = _repo_mock() + cloned_plan_ids = ["clone_1", "clone_2", "clone_3"] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: + clone_results = [MagicMock(id=pid) for pid in cloned_plan_ids] + MockPlanSvc.return_value.clone_plan_for_variant.side_effect = clone_results + create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + # 克隆被调用 3 次 + assert MockPlanSvc.return_value.clone_plan_for_variant.call_count == 3 + # 每个任务关联到不同的克隆 plan + for i, task in enumerate(tasks): + assert task.source_edit_plan_id == cloned_plan_ids[i] + + def test_preview_variant_titles_injected_per_variant(self): + """titles[] 按变体注入 title_config.text""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured_commands = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured_commands.append(cmd) + return tasks[len(captured_commands) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request( + preview_count=3, + title_config={"font": "黑体", "position": "bottom"}, + titles=["标题A", "标题B", "标题C"], + ), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert len(captured_commands) == 3 + assert captured_commands[0].title_config["text"] == "标题A" + assert captured_commands[1].title_config["text"] == "标题B" + assert captured_commands[2].title_config["text"] == "标题C" + # 样式全局共用 + assert all(c.title_config["font"] == "黑体" for c in captured_commands) + + def test_preview_shared_title_when_single_length(self): + """titles 长度 1 = 所有变体共用同一标题""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request(preview_count=3, titles=["共用标题"]), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert all(c.title_config["text"] == "共用标题" for c in captured) + + def test_preview_independent_voice_per_variant(self): + """voice_library_ids[] 按变体注入独立配音""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request( + preview_count=3, + voice_library_ids=["voice_a", "voice_b", "voice_c"], + ), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert [c.voice_library_id for c in captured] == ["voice_a", "voice_b", "voice_c"] + + def test_preview_voice_fallback_to_single_field(self): + """voice_library_ids 为空时回退 voice_library_id 单值字段(向后兼容)""" + from app.api.routes.generation_preview import create_preview_generation_task + + task = _make_task(task_id="task_1") + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return task + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request(voice_library_id="legacy_voice"), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert captured[0].voice_library_id == "legacy_voice" + + def test_preview_queue_limit_checks_total_count(self): + """限流预检查按变体总数计:用户 pending + N 超限 → 429""" + from app.api.routes.generation_preview import create_preview_generation_task + from fastapi import HTTPException + + repo = MagicMock() + repo.count_pending_by_user.return_value = 3 + repo.count_pending_total.return_value = 0 + with pytest.raises(HTTPException) as exc: + create_preview_generation_task( + _make_preview_request(preview_count=5), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert exc.value.status_code == 429 + + def test_preview_clone_failure_marks_all_failed(self): + """克隆变体 plan 失败 → 已创建任务全部标记 failed 并 500""" + from app.api.routes.generation_preview import create_preview_generation_task + from fastapi import HTTPException + + tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] + repo = _repo_mock() + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: + MockPlanSvc.return_value.clone_plan_for_variant.side_effect = RuntimeError("db down") + with pytest.raises(HTTPException) as exc: + create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert exc.value.status_code == 500 + # 所有已创建任务都被标记 failed + assert all(t.status == GenerationTaskStatus.FAILED for t in tasks) + + +# ════════════════════════════════════════════════════════════════════════════ +# 批量正式生成:变体配置注入 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestBatchGenerationVariantConfig: + """POST /tasks count=N 时变体独立配置。""" + + def _call_create_tasks(self, request, repo=None): + from app.api.routes.generation_tasks import create_generation_task + + repo = repo or MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + repo.update.return_value = None + + # 模板模式:asset_repository.find_by_id 返回 None(无 project 关联, + # 纯模板模式 project_id/library_id 都为空),避免 MagicMock 属性污染 + asset_repo = MagicMock() + asset_repo.find_by_id.return_value = None + + # db.query().filter()...first() 返回 None:不走兜底关联编辑计划 + db = MagicMock() + db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None + + return create_generation_task( + request, + authenticated_user=_make_user(), + generation_task_repository=repo, + project_repository=MagicMock(), + asset_library_repository=MagicMock(), + asset_repository=asset_repo, + db=db, + ) + + def test_count_3_variant_titles_voices_covers_injected(self): + """count=3:titles/voice_library_ids/cover_urls 按变体注入""" + from app.api.routes import generation_tasks as routes + + tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + t = tasks[len(captured) - 1] + t.title_config = cmd.title_config + t.voice_library_id = cmd.voice_library_id + t.cover_url = cmd.cover_url + return t + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + title_config={"font": "宋体"}, + titles=["成片标题1", "成片标题2", "成片标题3"], + voice_library_ids=["v1", "v2", "v3"], + cover_urls=["http://c1", "http://c2", "http://c3"], + ) + resp = self._call_create_tasks(req) + assert resp.total == 3 + assert [c.title_config["text"] for c in captured] == ["成片标题1", "成片标题2", "成片标题3"] + assert [c.voice_library_id for c in captured] == ["v1", "v2", "v3"] + assert [c.cover_url for c in captured] == ["http://c1", "http://c2", "http://c3"] + # 样式共用 + assert all(c.title_config["font"] == "宋体" for c in captured) + + def test_count_1_legacy_fields_unchanged(self): + """N=1 不传数组:旧字段 voice_library_id/cover_url/title_config 行为不变""" + from app.api.routes import generation_tasks as routes + + task = _make_task(task_id="gen_1") + task.is_preview = False + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return task + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=1, + voice_library_id="legacy_voice", + cover_url="http://legacy-cover", + title_config={"text": "旧标题", "font": "黑体"}, + ) + resp = self._call_create_tasks(req) + assert resp.total == 1 + assert captured[0].voice_library_id == "legacy_voice" + assert captured[0].cover_url == "http://legacy-cover" + assert captured[0].title_config["text"] == "旧标题" + + def test_count_3_shared_single_value_arrays(self): + """数组长度 1:3 个变体共用同一配音/封面""" + from app.api.routes import generation_tasks as routes + + tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + voice_library_ids=["shared_voice"], + cover_urls=["http://shared"], + ) + self._call_create_tasks(req) + assert all(c.voice_library_id == "shared_voice" for c in captured) + assert all(c.cover_url == "http://shared" for c in captured) + + +class TestVariantValueHelper: + """_variant_value 取值逻辑。""" + + def test_empty_returns_fallback(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value([], 0, fallback="fb") == "fb" + + def test_single_length_shared(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["only"], 5) == "only" + + def test_indexed_access(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["a", "b", "c"], 1) == "b" + + def test_index_out_of_range_fallback(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["a", "b"], 9, fallback="x") == "x" diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index 696e889aa..69a3444bc 100644 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -725,8 +725,12 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) - assert resp.task_id == "preview_task_001" - assert resp.status == "pending" + # 批量响应:N=1 时 items 长度为 1 + assert resp.total == 1 + assert len(resp.items) == 1 + assert resp.items[0].task_id == "preview_task_001" + assert resp.items[0].status == "pending" + assert resp.items[0].variant_index == 0 def test_user_pending_limit_exceeded(self): """用户待处理任务超限 → 429""" @@ -807,7 +811,13 @@ class TestCreatePreviewRoute: repo.count_pending_total.return_value = 0 task = _make_task() - from fastapi import HTTPException + + # 模拟 mark_failed 真实更新任务状态(_mark_task_failed 内部调用) + def _set_failed(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message + + task.mark_failed.side_effect = _set_failed with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task @@ -815,14 +825,15 @@ class TestCreatePreviewRoute: "app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=False, ): - with pytest.raises(HTTPException) as exc_info: - create_preview_generation_task( - self._make_request(), - authenticated_user=_make_user(), - generation_task_repository=repo, - db=MagicMock(), - ) - assert exc_info.value.status_code == 500 + resp = create_preview_generation_task( + self._make_request(), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + # 入队失败:任务被标记 failed(mark_failed 设置错误信息),响应正常返回 + assert resp.total == 1 + assert resp.items[0].status == "failed" def test_enqueue_raises_user_limit(self): """safe_enqueue 抛出 UserPendingLimitExceeded → 429""" @@ -833,6 +844,12 @@ class TestCreatePreviewRoute: task = _make_task() from fastapi import HTTPException + def _set_failed_limit(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message or "待处理任务超限" + + task.mark_failed.side_effect = _set_failed_limit + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch( @@ -846,6 +863,7 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) + # 全部变体入队失败且错误消息含"待处理任务" → 429 assert exc_info.value.status_code == 429 def test_enqueue_raises_global_queue_full(self): @@ -857,6 +875,12 @@ class TestCreatePreviewRoute: task = _make_task() from fastapi import HTTPException + def _set_failed_queue(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message or "系统队列已满" + + task.mark_failed.side_effect = _set_failed_queue + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch( @@ -870,6 +894,7 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) + # 全部变体入队失败且错误消息含"队列" → 503 assert exc_info.value.status_code == 503