diff --git a/apps/api/app/api/routes/viral_video.py b/apps/api/app/api/routes/viral_video.py index 2ebae48ee..e6c16053f 100644 --- a/apps/api/app/api/routes/viral_video.py +++ b/apps/api/app/api/routes/viral_video.py @@ -1,14 +1,21 @@ """爆款视频 API 路由。 -端点: - POST /api/v1/viral-video/generate 创建爆款视频任务 - GET /api/v1/viral-video/{job_id} 查询任务状态 - GET /api/v1/viral-video/history 历史记录 - POST /api/v1/viral-video/{job_id}/retry 重试失败任务 - POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案 - POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析 - GET /api/v1/viral-video/style-templates 获取风格模板列表 - WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送(订阅 Redis pub/sub) +v1.5 三步分步流水线端点(前端新交互): + POST /api/v1/viral-video/analyze-images 阶段1:创建任务 + 仅做图片/视频分析,暂停在 image_analyzed + POST /api/v1/viral-video/{job_id}/generate-copy 阶段2:用户填完参数后跑意图+文案+分镜+审核,暂停在 copy_generated + POST /api/v1/viral-video/{job_id}/confirm-copy 阶段3:用户确认/编辑文案后跑渲染,直到完成 + +旧端点(兼容保留,旧前端/一键生成模式): + POST /api/v1/viral-video/generate 一键入队,前半段跑到 wait_user_confirm + POST /api/v1/viral-video/{job_id}/confirm-intent 旧的意图确认后继续渲染 + +通用: + GET /api/v1/viral-video/{job_id} 查询任务状态(含 image_analysis/storyboard/generated_copy_text) + GET /api/v1/viral-video/history 历史记录 + POST /api/v1/viral-video/{job_id}/retry 重试失败任务 + POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析 + GET /api/v1/viral-video/style-templates 风格模板列表 + WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送 """ from __future__ import annotations @@ -19,10 +26,13 @@ from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app from app.dependencies import get_db_session from app.schemas.viral_video import ( + AnalyzeImagesRequest, AnalyzeStyleRequest, AnalyzeStyleResponse, + ConfirmCopyRequest, ConfirmIntentRequest, CreateViralVideoRequest, + GenerateCopyRequest, StyleTemplateListResponse, StyleTemplateResponse, ViralVideoHistoryResponse, @@ -45,6 +55,31 @@ router = APIRouter() # ── Helpers ────────────────────────────────────────────────────────────── +def _build_copy_result(job) -> dict | None: + """将后端原始字段拼装为前端期望的 CopyResult 结构(final_copy/suggested_copy/title/scenes)。""" + copy_text = getattr(job, "generated_copy_text", "") or "" + sb = getattr(job, "storyboard", None) or [] + intent = getattr(job, "intent_result", None) or {} + if not copy_text and not sb: + return None + scenes = [] + for seg in sb: + if isinstance(seg, dict): + scenes.append( + { + "shot": seg.get("description", ""), + "narration": seg.get("text", ""), + "duration": seg.get("duration"), + } + ) + return { + "title": (intent.get("suggested_title") if isinstance(intent, dict) else None) or "", + "final_copy": copy_text, + "suggested_copy": copy_text, + "scenes": scenes, + } + + def _to_response(job) -> ViralVideoJobResponse: return ViralVideoJobResponse( id=job.id, @@ -65,6 +100,10 @@ def _to_response(job) -> ViralVideoJobResponse: style_guide=job.style_guide, style_template_id=job.style_template_id, status=job.status, + image_analysis=getattr(job, "image_analysis", None), + storyboard=getattr(job, "storyboard", None), + generated_copy_text=getattr(job, "generated_copy_text", "") or "", + copy_result=_build_copy_result(job), intent_result=job.intent_result, result_video_url=job.result_video_url, credits_cost=job.credits_cost, @@ -133,6 +172,127 @@ def create_viral_video( return _to_response(job) +@router.post("/analyze-images", response_model=ViralVideoJobResponse) +def analyze_images( + request: AnalyzeImagesRequest, + authenticated_user: AuthenticatedUser = Depends(get_current_user), + session: Session = Depends(get_db_session), +) -> ViralVideoJobResponse: + """v1.5 阶段1:创建任务并仅做图片/视频 VLM 分析,跑完后状态=image_analyzed。 + + 前端拿到 image_analysis(商品名/品牌/特征/颜色/材质等结构化结果)展示给用户; + 用户填完营销参数后再调 /{id}/generate-copy 进入阶段2。 + """ + from packages.domain.viral_video import ViralVideoJob + + repo = _get_job_repo(session) + job = ViralVideoJob( + user_id=authenticated_user.user.id, + images=list(request.images), + reference_video_url=request.reference_video_url or "", + style_template_id=request.style_template_id or "", + style_strength=request.style_strength or "medium", + ) + repo.save(job) + + try: + celery_app.send_task("worker.run_viral_video_analyze", args=[job.id]) + logger.info("[爆款视频][阶段1] analyze-images 入队: job_id=%s", job.id) + except Exception as e: + logger.error("[爆款视频][阶段1] analyze-images 入队失败: %s", e, exc_info=True) + job.mark_failed(f"任务入队失败: {e}") + repo.update(job) + + return _to_response(job) + + +@router.post("/{job_id}/generate-copy", response_model=ViralVideoJobResponse) +def generate_copy( + job_id: str, + request: GenerateCopyRequest, + authenticated_user: AuthenticatedUser = Depends(get_current_user), + session: Session = Depends(get_db_session), +) -> ViralVideoJobResponse: + """v1.5 阶段2:用户填完营销参数后,跑 意图解析 → 文案融合 → 分镜 → 合规审核。 + + 跑完后状态=copy_generated,响应包含 generated_copy_text + storyboard, + 前端展示文案供用户编辑;确认/编辑后调 /{id}/confirm-copy 进入阶段3。 + """ + repo = _get_job_repo(session) + job = repo.get(job_id) + if job is None: + raise HTTPException(status_code=404, detail="任务不存在") + if job.user_id != authenticated_user.user.id: + raise HTTPException(status_code=403, detail="无权操作此任务") + if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED): + raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案") + + # 允许失败任务重试:重置 + if job.status == ViralVideoStatus.FAILED: + job.retry_count += 1 + job.error_msg = "" + + # 把用户填的营销参数写到 job 上 + job.industry = request.industry or job.industry + job.target_customer = request.target_customer or job.target_customer + job.persona_id = request.persona_id or job.persona_id + job.viral_structure = request.viral_structure or job.viral_structure + job.marketing_purpose = request.marketing_purpose or job.marketing_purpose + job.bgm_preference = request.bgm_preference or job.bgm_preference + if request.duration: + job.duration = request.duration + job.user_copy_text = request.user_copy_text if request.user_copy_text else job.user_copy_text + job.fusion_level = request.fusion_level or job.fusion_level + job.reference_audio_path = request.reference_audio_path or job.reference_audio_path + job.reference_video_url = request.reference_video_url or job.reference_video_url + job.style_strength = request.style_strength or job.style_strength + job.style_template_id = request.style_template_id or job.style_template_id + + job.resume_from_image_analyzed() + repo.update(job) + + try: + celery_app.send_task("worker.run_viral_video_generate_copy", args=[job.id]) + logger.info("[爆款视频][阶段2] generate-copy 入队: job_id=%s", job.id) + except Exception as e: + logger.error("[爆款视频][阶段2] generate-copy 入队失败: %s", e, exc_info=True) + job.mark_failed(f"任务入队失败: {e}") + repo.update(job) + + return _to_response(job) + + +@router.post("/{job_id}/confirm-copy", response_model=ViralVideoJobResponse) +def confirm_copy( + job_id: str, + request: ConfirmCopyRequest, + authenticated_user: AuthenticatedUser = Depends(get_current_user), + session: Session = Depends(get_db_session), +) -> ViralVideoJobResponse: + """v1.5 阶段3:用户确认/编辑文案后开始 TTS+BGM(skip)+渲染+上传。""" + repo = _get_job_repo(session) + job = repo.get(job_id) + if job is None: + raise HTTPException(status_code=404, detail="任务不存在") + if job.user_id != authenticated_user.user.id: + raise HTTPException(status_code=403, detail="无权操作此任务") + if job.status != ViralVideoStatus.COPY_GENERATED: + raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)") + + job.resume_from_copy_generated(edited_copy=request.edited_copy or None) + repo.update(job) + + try: + celery_app.send_task("worker.run_viral_video_render", args=[job.id]) + logger.info("[爆款视频][阶段3] confirm-copy 入队: job_id=%s", job.id) + except Exception as e: + logger.error("[爆款视频][阶段3] confirm-copy 入队失败: %s", e, exc_info=True) + job.mark_failed(f"任务入队失败: {e}") + repo.update(job) + + return _to_response(job) + + @router.get("/history", response_model=ViralVideoHistoryResponse) def list_viral_video_history( limit: int = 50, @@ -505,6 +665,8 @@ def _job_status(job) -> str: _STATUS_STAGE = { "pending": "", "running": "", + "image_analyzed": "image_analysis", + "copy_generated": "review", "wait_user_confirm": "intent_parsing", "completed": "uploading", "failed": "", @@ -514,6 +676,8 @@ _STATUS_STAGE = { _STATUS_PROGRESS = { "pending": 0.0, "running": 5.0, + "image_analyzed": 15.0, + "copy_generated": 70.0, "wait_user_confirm": 35.0, "completed": 100.0, "failed": 0.0, @@ -523,6 +687,8 @@ _STATUS_PROGRESS = { _STATUS_MESSAGE = { "pending": "任务已创建,等待执行", "running": "任务执行中", + "image_analyzed": "图片分析完成,等待填写营销参数", + "copy_generated": "文案与分镜已生成,等待确认文案", "wait_user_confirm": "等待用户确认意图文案", "completed": "视频生成完成", "failed": "任务失败", diff --git a/apps/api/app/schemas/viral_video.py b/apps/api/app/schemas/viral_video.py index 359aeae18..2789796f6 100755 --- a/apps/api/app/schemas/viral_video.py +++ b/apps/api/app/schemas/viral_video.py @@ -6,7 +6,7 @@ from datetime import datetime from pydantic import BaseModel, Field, field_validator -# ── 枚举常量 ───────────────────────────────────────────────────────────── +# -- 枚举常量 -- VALID_FUSION_LEVELS = ("ai_full", "full_ai", "ai_polish", "user_primary") VALID_STYLE_STRENGTHS = ("light", "medium", "strict") @@ -25,11 +25,11 @@ VALID_STAGES = ( ) -# ── Request Schemas ──────────────────────────────────────────────────────── +# -- Request Schemas -- class CreateViralVideoRequest(BaseModel): - """创建爆款视频任务请求。""" + """创建爆款视频任务请求(旧接口:一键跑完前半段到 WAIT_USER_CONFIRM,保留兼容)。""" images: list[str] = Field(..., min_length=1, max_length=20, description="产品图片 URL 列表") industry: str = Field(default="", description="行业") @@ -42,7 +42,6 @@ class CreateViralVideoRequest(BaseModel): user_copy_text: str = Field(default="", description="用户原始文案(我说你写)") fusion_level: str = Field(default="ai_polish", description="文案融合级别: ai_full/ai_polish/user_primary") reference_audio_path: str = Field(default="", description="参考音频路径") - # v1.3 新增 reference_video_url: str = Field(default="", description="参考爆款视频 URL") style_strength: str = Field(default="medium", description="风格强度: light/medium/strict") style_template_id: str = Field(default="", description="风格模板 ID") @@ -50,25 +49,73 @@ class CreateViralVideoRequest(BaseModel): @field_validator("fusion_level") @classmethod def _validate_fusion_level(cls, v: str) -> str: - # 兼容前端历史写法 full_ai(等价 ai_full) if v == "full_ai": return "ai_full" if v not in VALID_FUSION_LEVELS: - raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一") + raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}") return v @field_validator("style_strength") @classmethod def _validate_style_strength(cls, v: str) -> str: if v not in VALID_STYLE_STRENGTHS: - raise ValueError(f"style_strength 必须是 {VALID_STYLE_STRENGTHS} 之一") + raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}") return v -class ConfirmIntentRequest(BaseModel): - """确认意图请求(confirm-intent)。""" +class AnalyzeImagesRequest(BaseModel): + """v1.5 阶段1:创建任务并仅做图片/视频分析。images 必填,其他参数可选(阶段2再传)。""" - confirmed_copy: str = Field(default="", description="用户确认/修改后的文案,为空表示使用 AI 生成的文案") + images: list[str] = Field(..., min_length=1, max_length=20) + reference_video_url: str = Field(default="", description="参考爆款视频 URL(可选,有则同步做风格分析)") + style_template_id: str = Field(default="", description="风格模板 ID(可选)") + style_strength: str = Field(default="medium") + + +class GenerateCopyRequest(BaseModel): + """v1.5 阶段2:用户填完参数后跑意图+文案+分镜+审核,暂停在 COPY_GENERATED。""" + + industry: str = Field(default="") + target_customer: str = Field(default="") + persona_id: str = Field(default="") + viral_structure: str = Field(default="") + marketing_purpose: str = Field(default="") + bgm_preference: str = Field(default="") + duration: int = Field(default=30, ge=5, le=180) + user_copy_text: str = Field(default="") + fusion_level: str = Field(default="ai_polish") + reference_audio_path: str = Field(default="") + reference_video_url: str = Field(default="") + style_strength: str = Field(default="medium") + style_template_id: str = Field(default="") + + @field_validator("fusion_level") + @classmethod + def _v_fl(cls, v: str) -> str: + if v == "full_ai": + return "ai_full" + if v not in VALID_FUSION_LEVELS: + raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}") + return v + + @field_validator("style_strength") + @classmethod + def _v_ss(cls, v: str) -> str: + if v not in VALID_STYLE_STRENGTHS: + raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}") + return v + + +class ConfirmCopyRequest(BaseModel): + """v1.5 阶段3:用户确认/编辑文案后开始渲染。""" + + edited_copy: str = Field(default="", description="用户编辑后的最终文案;为空则使用 AI 生成文案") + + +class ConfirmIntentRequest(BaseModel): + """确认意图请求(旧 confirm-intent,兼容)。""" + + confirmed_copy: str = Field(default="", description="用户确认/修改后的文案") adjustments: str = Field(default="", description="用户对 AI 文案的调整意见") @@ -79,11 +126,11 @@ class AnalyzeStyleRequest(BaseModel): style_template_id: str = Field(default="", description="风格模板 ID(可选覆盖)") -# ── Response Schemas ─────────────────────────────────────────────────────── +# -- Response Schemas -- class ViralVideoJobResponse(BaseModel): - """爆款视频任务响应。""" + """爆款视频任务响应。v1.5 新增 image_analysis/storyboard/generated_copy_text 字段。""" id: str user_id: str @@ -103,6 +150,13 @@ class ViralVideoJobResponse(BaseModel): style_guide: dict | None = None style_template_id: str = "" status: str + # v1.4 VLM 结果 + image_analysis: dict | None = None + # v1.5 三步分步产物(原始字段,保留给后端/老调用方) + storyboard: list | None = None + generated_copy_text: str = "" + # v1.5 前端 CopyResult 结构(final_copy/suggested_copy/title/scenes) + copy_result: dict | None = None intent_result: dict | None = None result_video_url: str = "" credits_cost: int = 0 @@ -115,15 +169,11 @@ class ViralVideoJobResponse(BaseModel): class ViralVideoHistoryResponse(BaseModel): - """历史记录列表响应。""" - items: list[ViralVideoJobResponse] total: int class StyleTemplateResponse(BaseModel): - """风格模板响应。""" - id: str name: str description: str = "" @@ -132,25 +182,19 @@ class StyleTemplateResponse(BaseModel): class StyleTemplateListResponse(BaseModel): - """风格模板列表响应。""" - items: list[StyleTemplateResponse] class AnalyzeStyleResponse(BaseModel): - """风格分析结果响应。""" - job_id: str status: str style_guide: dict | None = None -# ── WebSocket 事件 Schema ────────────────────────────────────────────────── +# -- WebSocket 事件 Schema -- class WSProgressEvent(BaseModel): - """WebSocket 进度推送事件。""" - type: str = "viral_video:progress" job_id: str stage: str diff --git a/apps/web/src/api/viral-video/index.ts b/apps/web/src/api/viral-video/index.ts index 691be3394..35c49321a 100644 --- a/apps/web/src/api/viral-video/index.ts +++ b/apps/web/src/api/viral-video/index.ts @@ -6,6 +6,9 @@ import type { ViralVideoJob, ImageAnalysisResult, CopyResult, + AnalyzeImagesRequest, + GenerateCopyRequest, + ConfirmCopyRequest, } from "./types" /** 创建爆款视频任务 */ @@ -105,3 +108,24 @@ export function mockGenerateCopy(params: { }, 2200) }) } + +/** ── 三步拆分 v1.5 真实后端 API(PR #2117 合入后启用,前端可替换 mock 调用) ── */ + +/** 阶段1:上传图片后仅做 VLM 图片分析 + 可选参考视频风格分析,完成后状态=image_analyzed */ +export function analyzeViralImages(payload: AnalyzeImagesRequest) { + return apiClient.post("/viral-video/analyze-images", payload).then((r) => r.data) +} + +/** 阶段2:用户填完营销参数后生成文案+分镜+合规审核,完成后状态=copy_generated,返回 copy_result */ +export function generateViralCopy(id: string, payload: GenerateCopyRequest) { + return apiClient + .post(`/viral-video/${id}/generate-copy`, payload) + .then((r) => r.data) +} + +/** 阶段3:用户确认/编辑文案后开始 TTS→渲染→上传,完成后状态=completed */ +export function confirmViralCopy(id: string, payload: ConfirmCopyRequest = {}) { + return apiClient + .post(`/viral-video/${id}/confirm-copy`, payload) + .then((r) => r.data) +} diff --git a/apps/web/src/api/viral-video/types.ts b/apps/web/src/api/viral-video/types.ts index 281dd00d1..0fd27db7b 100644 --- a/apps/web/src/api/viral-video/types.ts +++ b/apps/web/src/api/viral-video/types.ts @@ -175,3 +175,46 @@ export interface HistoryResponse { page: number page_size: number } + +/** v1.5 阶段1请求:仅做图片/视频分析(POST /viral-video/analyze-images) */ +export interface AnalyzeImagesRequest { + images: string[] + reference_video_url?: string + style_template_id?: string + style_strength?: StyleStrength +} + +/** v1.5 阶段2请求:填完营销参数后生成文案+分镜(POST /viral-video/{id}/generate-copy) */ +export interface GenerateCopyRequest { + industry?: string + target_customer?: string + persona_id?: string + viral_structure?: string + marketing_purpose?: string + bgm_preference?: string + duration?: number + user_copy_text?: string + fusion_level?: FusionLevel + reference_audio_path?: string + reference_video_url?: string + style_strength?: StyleStrength + style_template_id?: string + style_guide?: string | Record +} + +/** v1.5 阶段3请求:用户确认/编辑文案后开始渲染(POST /viral-video/{id}/confirm-copy) */ +export interface ConfirmCopyRequest { + /** 用户编辑后的最终文案;为空则使用 AI 生成文案 */ + edited_copy?: string +} + +/** 分镜片段结构(后端 storyboard 字段的元素形态,保留供调试/进阶使用;主流程请使用 copy_result.scenes) */ +export interface StoryboardSegment { + order: number + type: string + description: string + text: string + duration: number + ken_burns?: string + transition?: string +} diff --git a/apps/worker/worker_app/tasks/viral_video.py b/apps/worker/worker_app/tasks/viral_video.py index d4b7bb937..20cea372a 100644 --- a/apps/worker/worker_app/tasks/viral_video.py +++ b/apps/worker/worker_app/tasks/viral_video.py @@ -835,65 +835,9 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict: if job.status != ViralVideoStatus.RUNNING: return {"ok": False, "error": f"unexpected status: {job.status}"} - # P0-3: 从 job 读取 image_analysis(run_pipeline 阶段已持久化) - image_analysis = job.image_analysis or {"products": []} - - _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...") - - # ── Step 3: 文案融合 ── - copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis) - _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成") - - # ── Step 4: 分镜脚本 ── - _emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...") - storyboard = _step_storyboard(job, copy_text, image_analysis) - _emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)}) - - # ── Step 5: 合规审核 ── - _emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...") - review_result = _step_review(job, copy_text, storyboard) - if not review_result.get("passed", True): - _emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...") - copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis) - review_result = _step_review(job, copy_text, storyboard) - _emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成") - - # ── Step 6: CosyVoice 配音(返回 Path | None) ── - _emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...") - tts_path = _step_tts(job, copy_text) - _emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None}) - - # ── Step 7: BGM 选择(P1:暂返回 None,跳过) ── - _emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)") - bgm = _step_bgm_select(job) - - # ── Step 8: 渲染(逐分镜 Seedance → concat → 混 TTS) ── - _emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...") - video_path = _step_render(job, storyboard, tts_path, bgm) - _emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成") - - # ── Step 9: OSS 上传 + 扣点 ── - # 注:爆款视频由 Seedance 2.5 直接生成口型,不需要 MuseTalk 事后对口型(MuseTalk 是 AI 数字人路线用的)。 - _emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...") - video_url = _step_upload(job, video_path) - - job.credits_cost = CREDITS_VIRAL_VIDEO_COST - # TODO: 调用 credits.deduct() 实际扣点(#1895 总开关为 false 时不扣,保留 TODO) - - job.mark_completed(video_url) - _save_job(repo, job, session) - _emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url}) - _emit_progress( - job_id, - ViralVideoStage.UPLOADING, - 100.0, - "视频生成完成", - {"video_url": video_url}, - event_type="viral_video:completed", - ) - - logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url) - return {"ok": True, "job_id": job_id, "video_url": video_url} + # 旧 confirm-intent 路径:v1.4 及之前 storyboard/copy_text 未持久化,交给 _run_render_pipeline 兜底重算; + # 新 v1.5 路径(copy_generated -> run_viral_video_render)直接走新 task,不会进入这里。 + return _run_render_pipeline(job_id, session, repo, job) except Retry: raise @@ -951,3 +895,247 @@ def run_video_style_analysis(self: Task, job_id: str) -> dict: finally: if session: session.close() + + +# ── v1.5 三步分步流水线 Celery 任务 ────────────────────────────────────── + + +def _mark_failed_and_notify(job_id: str, session, repo, job, err_msg: str, stage: str = "") -> None: + """统一的失败处理:标记 FAILED + 发 failed WS 事件。""" + try: + if session is None: + session = SessionLocal() + repo = SQLAlchemyViralVideoJobRepository(session) + job = repo.get(job_id) + if job is not None and not job.is_terminal: + job.mark_failed(err_msg) + _save_job(repo, job, session) + except Exception as inner: + logger.warning("[爆款视频] 标记失败状态时出错: %s", inner) + _emit_progress( + job_id, + stage, + 0, + f"任务失败: {err_msg}", + {"error": err_msg}, + event_type="viral_video:failed", + ) + + +@shared_task(bind=True, max_retries=1, name="worker.run_viral_video_analyze") +def run_viral_video_analyze(self: Task, job_id: str) -> dict: + """v1.5 阶段1:仅跑图片 VLM 分析(+ 可选视频风格分析),完成后状态=image_analyzed。""" + session = None + try: + session, repo, job = _get_repo_and_job(job_id) + if job is None: + logger.error("[爆款视频][阶段1] 任务不存在: %s", job_id) + return {"ok": False, "error": "job not found"} + + job.mark_running() + _save_job(repo, job, session) + _emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 5.0, "开始图片分析") + + _emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 10.0, "正在分析产品图片...") + image_analysis = _step_image_analysis(job) + job.image_analysis = image_analysis + _save_job(repo, job, session) + _emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 60.0, "图片分析完成", {"result": image_analysis}) + + style_guide = None + if job.reference_video_url or job.style_template_id: + _emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 70.0, "正在分析参考视频风格...") + style_guide = _step_video_analysis(job) + job.style_guide = style_guide + _save_job(repo, job, session) + _emit_progress( + job_id, + ViralVideoStage.VIDEO_ANALYSIS, + 90.0, + "风格分析完成", + {"style_analyzed": True, "style_guide": style_guide}, + ) + + job.mark_image_analyzed() + _save_job(repo, job, session) + _emit_progress( + job_id, + ViralVideoStage.IMAGE_ANALYSIS, + 100.0, + "图片分析完成,请填写营销参数以生成文案", + {"image_analysis": image_analysis, "status": "image_analyzed"}, + event_type="viral_video:image_analyzed", + ) + logger.info("[爆款视频][阶段1] 图片分析完成 job_id=%s", job_id) + return {"ok": True, "job_id": job_id, "status": "image_analyzed", "image_analysis": image_analysis} + + except Retry: + raise + except Exception as e: + logger.error("[爆款视频][阶段1] 异常: %s", e, exc_info=True) + _mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.IMAGE_ANALYSIS) + return {"ok": False, "job_id": job_id, "error": str(e)} + finally: + if session: + session.close() + + +@shared_task(bind=True, max_retries=1, name="worker.run_viral_video_generate_copy") +def run_viral_video_generate_copy(self: Task, job_id: str) -> dict: + """v1.5 阶段2:跑 意图解析 → 文案融合 → 分镜 → 合规审核,完成后状态=copy_generated。 + + 入参要求:调用方已 resume_from_image_analyzed() 把状态切到 RUNNING,并把用户填的营销参数写到 job 上。 + """ + session = None + try: + session, repo, job = _get_repo_and_job(job_id) + if job is None: + return {"ok": False, "error": "job not found"} + + if job.status != ViralVideoStatus.RUNNING: + return {"ok": False, "error": f"unexpected status: {job.status}"} + + image_analysis = job.image_analysis or {"products": []} + + # Step 2: 意图解析 + _emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 20.0, "正在解析文案意图...") + intent_result = _step_intent_parsing(job, image_analysis) + job.intent_result = intent_result + _save_job(repo, job, session) + _emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 35.0, "意图解析完成") + + # Step 3: 文案融合 + _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...") + copy_text = _step_copy_fusion(job, intent_result, image_analysis) + _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成") + + # Step 4: 分镜 + _emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...") + storyboard = _step_storyboard(job, copy_text, image_analysis) + _emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)}) + + # Step 5: 合规审核 + _emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...") + review_result = _step_review(job, copy_text, storyboard) + if not review_result.get("passed", True): + _emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...") + copy_text = _step_copy_fusion(job, intent_result, image_analysis) + _step_review(job, copy_text, storyboard) + _emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成") + + job.mark_copy_generated(copy_text, storyboard) + _save_job(repo, job, session) + _emit_progress( + job_id, + ViralVideoStage.REVIEW, + 100.0, + "文案与分镜已生成,请确认或编辑文案", + { + "generated_copy_text": copy_text, + "storyboard": storyboard, + "status": "copy_generated", + }, + event_type="viral_video:copy_generated", + ) + logger.info( + "[爆款视频][阶段2] 文案+分镜生成完成 job_id=%s copy_len=%d segs=%d", job_id, len(copy_text), len(storyboard) + ) + return { + "ok": True, + "job_id": job_id, + "status": "copy_generated", + "generated_copy_text": copy_text, + "storyboard": storyboard, + } + + except Retry: + raise + except Exception as e: + logger.error("[爆款视频][阶段2] 异常: %s", e, exc_info=True) + _mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.COPY_FUSION) + return {"ok": False, "job_id": job_id, "error": str(e)} + finally: + if session: + session.close() + + +def _run_render_pipeline(job_id: str, session, repo, job) -> dict: + """v1.5 阶段3 / 旧 resume 共用:TTS → BGM → Render → Upload → Completed。 + + 入参要求:job.status == RUNNING,job.generated_copy_text 或 job.user_copy_text 非空,job.storyboard 已就绪。 + """ + image_analysis = job.image_analysis or {"products": []} + copy_text = job.effective_copy_text + storyboard = job.storyboard or [] + + # 兼容旧路径:老 resume_viral_video_pipeline 在 RUNNING 时可能还没 storyboard(v1.4 及之前 job.storyboard 没持久化), + # 这种情况下用 intent_result + image_analysis 现算 copy_text + storyboard。 + if not storyboard: + _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...") + copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis) + _emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成") + _emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...") + storyboard = _step_storyboard(job, copy_text, image_analysis) + _emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)}) + _emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...") + _step_review(job, copy_text, storyboard) + _emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成") + # 补持久化 + job.storyboard = storyboard + job.generated_copy_text = copy_text + _save_job(repo, job, session) + + # Step 6: TTS + _emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...") + tts_path = _step_tts(job, copy_text) + _emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None}) + + # Step 7: BGM (skip) + _emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)") + bgm = _step_bgm_select(job) + + # Step 8: Render + _emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...") + video_path = _step_render(job, storyboard, tts_path, bgm) + _emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成") + + # Step 9: Upload + _emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...") + video_url = _step_upload(job, video_path) + + job.credits_cost = CREDITS_VIRAL_VIDEO_COST + job.mark_completed(video_url) + _save_job(repo, job, session) + _emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url}) + _emit_progress( + job_id, + ViralVideoStage.UPLOADING, + 100.0, + "视频生成完成", + {"video_url": video_url}, + event_type="viral_video:completed", + ) + logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url) + return {"ok": True, "job_id": job_id, "video_url": video_url} + + +@shared_task(bind=True, max_retries=2, name="worker.run_viral_video_render") +def run_viral_video_render(self: Task, job_id: str) -> dict: + """v1.5 阶段3:用户确认/编辑文案后,跑 TTS+BGM+Render+Upload 直到完成。""" + session = None + try: + session, repo, job = _get_repo_and_job(job_id) + if job is None: + return {"ok": False, "error": "job not found"} + if job.status != ViralVideoStatus.RUNNING: + return {"ok": False, "error": f"unexpected status: {job.status}"} + return _run_render_pipeline(job_id, session, repo, job) + except Retry: + raise + except Exception as e: + logger.error("[爆款视频][阶段3] 异常: %s", e, exc_info=True) + _mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.RENDERING) + return {"ok": False, "job_id": job_id, "error": str(e)} + finally: + if session: + session.close() diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 5c0218cbb..06d423720 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -947,6 +947,8 @@ class ViralVideoJobModel(Base): status = Column(String(30), nullable=False, default="pending", index=True) intent_result = Column(JSON, nullable=True) image_analysis = Column(JSON, nullable=True) + storyboard = Column(JSON, nullable=True) + generated_copy_text = Column(Text, nullable=False, default="") result_video_url = Column(String(1000), nullable=False, default="") credits_cost = Column(Integer, nullable=False, default=0) error_msg = Column(Text, nullable=False, default="") diff --git a/packages/adapters/sqlalchemy_impl/session.py b/packages/adapters/sqlalchemy_impl/session.py index 02c0ce4c9..77642159f 100644 --- a/packages/adapters/sqlalchemy_impl/session.py +++ b/packages/adapters/sqlalchemy_impl/session.py @@ -81,6 +81,36 @@ def ensure_database_exists(database_url: str) -> None: admin_engine.dispose() +_VIRAL_VIDEO_BACKFILL_COLS = [ + ("storyboard", "JSON"), + ("generated_copy_text", "TEXT NOT NULL DEFAULT ''"), +] + + +def _ensure_viral_video_columns(connection) -> None: + """Idempotently add new columns to viral_video_jobs; create_all will not ALTER existing tables.""" + from sqlalchemy import inspect as _inspect + + try: + insp = _inspect(connection) + if not insp.has_table("viral_video_jobs"): + return + existing = {c["name"] for c in insp.get_columns("viral_video_jobs")} + except Exception: + return + import logging as _logging + + _log = _logging.getLogger(__name__) + for col, ddl in _VIRAL_VIDEO_BACKFILL_COLS: + if col in existing: + continue + try: + connection.execute(text(f"ALTER TABLE viral_video_jobs ADD COLUMN {col} {ddl}")) + _log.info("added column viral_video_jobs.%s", col) + except Exception as e: + _log.warning("add column %s failed: %s", col, e) + + def initialize_database(engine) -> None: """初始化数据库 schema。 @@ -100,4 +130,5 @@ def initialize_database(engine) -> None: text("SELECT pg_advisory_unlock(:lock_id)"), {"lock_id": SCHEMA_INIT_LOCK_ID}, ) + _ensure_viral_video_columns(connection) connection.commit() diff --git a/packages/adapters/sqlalchemy_impl/viral_video_repository.py b/packages/adapters/sqlalchemy_impl/viral_video_repository.py index 02e5b0d27..0e0ba53b8 100755 --- a/packages/adapters/sqlalchemy_impl/viral_video_repository.py +++ b/packages/adapters/sqlalchemy_impl/viral_video_repository.py @@ -35,6 +35,8 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob: status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING, intent_result=dict(model.intent_result) if model.intent_result else None, image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None, + storyboard=list(model.storyboard) if getattr(model, "storyboard", None) else None, + generated_copy_text=getattr(model, "generated_copy_text", "") or "", result_video_url=model.result_video_url or "", credits_cost=model.credits_cost or 0, error_msg=model.error_msg or "", @@ -74,6 +76,8 @@ class SQLAlchemyViralVideoJobRepository: status=job.status, intent_result=job.intent_result, image_analysis=job.image_analysis, + storyboard=job.storyboard, + generated_copy_text=job.generated_copy_text, result_video_url=job.result_video_url, credits_cost=job.credits_cost, error_msg=job.error_msg, @@ -94,6 +98,8 @@ class SQLAlchemyViralVideoJobRepository: model.status = job.status model.intent_result = job.intent_result model.image_analysis = job.image_analysis + model.storyboard = job.storyboard + model.generated_copy_text = job.generated_copy_text or "" model.result_video_url = job.result_video_url model.credits_cost = job.credits_cost model.error_msg = job.error_msg @@ -101,6 +107,20 @@ class SQLAlchemyViralVideoJobRepository: model.started_at = job.started_at model.completed_at = job.completed_at model.style_guide = job.style_guide + # v1.5 three-stage: persist user-editable params so resume uses latest values + model.user_copy_text = job.user_copy_text + model.industry = job.industry + model.target_customer = job.target_customer + model.persona_id = job.persona_id + model.viral_structure = job.viral_structure + model.marketing_purpose = job.marketing_purpose + model.bgm_preference = job.bgm_preference + model.duration = job.duration + model.fusion_level = job.fusion_level + model.reference_audio_path = job.reference_audio_path + model.reference_video_url = job.reference_video_url + model.style_strength = job.style_strength + model.style_template_id = job.style_template_id model.updated_at = datetime.now(timezone.utc) self.session.commit() @@ -126,7 +146,9 @@ class SQLAlchemyViralVideoJobRepository: self.session.query(ViralVideoJobModel) .filter( ViralVideoJobModel.user_id == user_id, - ViralVideoJobModel.status.in_(["pending", "running", "wait_user_confirm"]), + ViralVideoJobModel.status.in_( + ["pending", "running", "wait_user_confirm", "image_analyzed", "copy_generated"] + ), ) .count() ) diff --git a/packages/domain/viral_video.py b/packages/domain/viral_video.py index 148d77283..3ce7ed09a 100755 --- a/packages/domain/viral_video.py +++ b/packages/domain/viral_video.py @@ -1,10 +1,10 @@ """ViralVideoJob 领域模型 — 爆款视频任务. -状态机: - pending → running → completed - ↘ failed → pending (retry) - ↘ cancelled - running 中可暂停:running → wait_user_confirm → running (confirm-intent resume) +状态机(v1.5 三步分步): + pending -> running -> image_analyzed -> running -> copy_generated -> running -> completed + wait_user_confirm -> running -> completed (旧路径兼容) + 任意阶段 fail; 任意非终态 cancel. + failed -> pending (retry 重置后重跑)。 """ from __future__ import annotations @@ -30,6 +30,8 @@ class ViralVideoStatus(StrEnum): PENDING = "pending" RUNNING = "running" + IMAGE_ANALYZED = "image_analyzed" + COPY_GENERATED = "copy_generated" WAIT_USER_CONFIRM = "wait_user_confirm" COMPLETED = "completed" FAILED = "failed" @@ -120,6 +122,9 @@ class ViralVideoJob: style_template_id: str = "" # v1.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜) image_analysis: dict | None = None + # v1.5 三步分步流水线产物(持久化,供前端 GET 读取 + resume 消费) + storyboard: list | None = None + generated_copy_text: str = "" # 状态 id: str = field(default_factory=lambda: uuid4().hex) status: ViralVideoStatus = ViralVideoStatus.PENDING @@ -133,23 +138,74 @@ class ViralVideoJob: created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) - # ── 状态转换 ── + # -- 状态转换 -- def mark_running(self) -> None: - if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING): + if self.status not in ( + ViralVideoStatus.PENDING, + ViralVideoStatus.IMAGE_ANALYZED, + ViralVideoStatus.COPY_GENERATED, + ViralVideoStatus.WAIT_USER_CONFIRM, + ViralVideoStatus.RUNNING, + ): raise ValueError(f"Cannot transition from {self.status} to running") self.status = ViralVideoStatus.RUNNING - self.started_at = datetime.now(timezone.utc) + if self.started_at is None: + self.started_at = datetime.now(timezone.utc) + self.updated_at = datetime.now(timezone.utc) + + def mark_image_analyzed(self) -> None: + """阶段1完成:图片/视频分析完成,等待用户填参数或直接触发阶段2。""" + if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING): + raise ValueError(f"Cannot transition from {self.status} to image_analyzed") + self.status = ViralVideoStatus.IMAGE_ANALYZED + if self.started_at is None: + self.started_at = datetime.now(timezone.utc) + self.updated_at = datetime.now(timezone.utc) + + def mark_copy_generated(self, copy_text: str, storyboard: list) -> None: + """阶段2完成:文案+分镜+审核完成,等待用户编辑后确认。""" + if self.status not in ( + ViralVideoStatus.IMAGE_ANALYZED, + ViralVideoStatus.RUNNING, + ViralVideoStatus.PENDING, + ): + raise ValueError(f"Cannot transition from {self.status} to copy_generated") + self.status = ViralVideoStatus.COPY_GENERATED + self.generated_copy_text = copy_text or "" + self.storyboard = list(storyboard) if storyboard else [] self.updated_at = datetime.now(timezone.utc) def mark_wait_user_confirm(self, intent_result: dict) -> None: + """旧流水线兼容:意图解析完成等待用户确认(老接口)。""" if self.status != ViralVideoStatus.RUNNING: raise ValueError(f"Cannot transition from {self.status} to wait_user_confirm") self.status = ViralVideoStatus.WAIT_USER_CONFIRM self.intent_result = intent_result self.updated_at = datetime.now(timezone.utc) + def resume_from_image_analyzed(self, **kwargs) -> None: + """阶段1->阶段2:用户已填参数,开始跑文案/分镜。kwargs 覆盖参数字段。""" + if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING): + raise ValueError(f"Cannot resume from {self.status} to copy-gen") + for k, v in kwargs.items(): + if hasattr(self, k) and v not in (None, "", []): + setattr(self, k, v) + self.status = ViralVideoStatus.RUNNING + self.updated_at = datetime.now(timezone.utc) + + def resume_from_copy_generated(self, edited_copy: str | None = None) -> None: + """阶段2->阶段3:用户确认/编辑文案,开始跑渲染。""" + if self.status != ViralVideoStatus.COPY_GENERATED: + raise ValueError(f"Cannot resume from {self.status} to render") + if edited_copy: + self.user_copy_text = edited_copy + self.generated_copy_text = edited_copy + self.status = ViralVideoStatus.RUNNING + self.updated_at = datetime.now(timezone.utc) + def resume_from_confirm(self) -> None: + """旧流水线兼容:从 WAIT_USER_CONFIRM 恢复。""" if self.status != ViralVideoStatus.WAIT_USER_CONFIRM: raise ValueError(f"Cannot resume from {self.status}") self.status = ViralVideoStatus.RUNNING @@ -181,3 +237,8 @@ class ViralVideoJob: ViralVideoStatus.FAILED, ViralVideoStatus.CANCELLED, ) + + @property + def effective_copy_text(self) -> str: + """渲染阶段使用的最终文案:用户编辑 > 生成文案 > 用户原始 > 占位。""" + return self.user_copy_text or self.generated_copy_text or "精选好物推荐" diff --git a/tests/unit/test_viral_video_p0.py b/tests/unit/test_viral_video_p0.py index 76ccd31b0..5c560f684 100644 --- a/tests/unit/test_viral_video_p0.py +++ b/tests/unit/test_viral_video_p0.py @@ -224,10 +224,15 @@ class TestDoubaoClientVideoGen: class TestResumeReadsImageAnalysis: def test_resume_uses_persisted_image_analysis(self): - """resume_pipeline 应从 job.image_analysis 读(P0-3 持久化)。""" + """resume/render pipeline 应从 job.image_analysis 读(v1.5 _run_render_pipeline 共享渲染逻辑)。""" import inspect from apps.worker.worker_app.tasks import viral_video as vv - src = inspect.getsource(vv.resume_viral_video_pipeline) + # v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis + src = inspect.getsource(vv._run_render_pipeline) assert "job.image_analysis" in src + assert "image_analysis" in src + # resume 本身应该调用 _run_render_pipeline + resume_src = inspect.getsource(vv.resume_viral_video_pipeline) + assert "_run_render_pipeline" in resume_src diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index 02e986220..eaf231f2d 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -51,6 +51,10 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending "stage": "", "progress": 0.0, "intent_result": None, + "image_analysis": None, + "storyboard": None, + "generated_copy_text": "", + "credits_cost": 0, "updated_at": None, }.items(): setattr(job, k, kwargs.pop(k, v)) @@ -174,3 +178,148 @@ class TestAnalyzeStyle: mock_send.assert_called_once_with("worker.run_video_style_analysis", args=["job-sty"]) assert resp.job_id == "job-sty" assert resp.status == "analyzing" + + +# ── v1.5 three-stage endpoints ───────────────────────────────────────── + + +class TestAnalyzeImages: + def test_analyze_images_creates_job_and_dispatches(self): + """POST /analyze-images: 创建任务 + 入队 run_viral_video_analyze。""" + from unittest.mock import patch + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import AnalyzeImagesRequest + + user = _auth_user("u1") + session = MagicMock() + repo = MagicMock() + req = AnalyzeImagesRequest(images=["https://x.com/a.jpg"], reference_video_url="", style_template_id="") + + saved = {} + + def fake_save(job): + saved["job"] = job + return job + + repo.save.side_effect = fake_save + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.analyze_images(req, authenticated_user=user, session=session) + + job = saved["job"] + assert job.user_id == "u1" + assert job.images == ["https://x.com/a.jpg"] + mock_send.assert_called_once_with("worker.run_viral_video_analyze", args=[job.id]) + assert resp.status == "pending" + + +class TestGenerateCopy: + def test_generate_copy_updates_params_and_dispatches(self): + """POST /{id}/generate-copy: 在 image_analyzed 状态下写营销参数 + 入队 run_viral_video_generate_copy。""" + from unittest.mock import patch + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import GenerateCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-gc", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED) + repo = MagicMock() + repo.get.return_value = job + req = GenerateCopyRequest( + industry="美妆", + target_customer="年轻女性", + duration=25, + fusion_level="ai_full", + user_copy_text="试试这个", + ) + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.generate_copy("job-gc", req, authenticated_user=user, session=session) + + # 参数写入 + assert job.industry == "美妆" + assert job.target_customer == "年轻女性" + assert job.duration == 25 + assert job.fusion_level == "ai_full" + assert job.user_copy_text == "试试这个" + job.resume_from_image_analyzed.assert_called_once() + repo.update.assert_called() + mock_send.assert_called_once_with("worker.run_viral_video_generate_copy", args=["job-gc"]) + assert resp.id == "job-gc" + + def test_generate_copy_rejects_wrong_status(self): + """任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。""" + import pytest + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import GenerateCopyRequest + from fastapi import HTTPException + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + repo = MagicMock() + repo.get.return_value = job + + with (patch.object(vv_mod, "_get_job_repo", return_value=repo),): + with pytest.raises(HTTPException) as exc: + vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session) + assert exc.value.status_code == 409 + + +class TestConfirmCopy: + def test_confirm_copy_dispatches_render(self): + """POST /{id}/confirm-copy: copy_generated -> RUNNING + 入队 run_viral_video_render,编辑文案写入。""" + from unittest.mock import patch + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-cc", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest(edited_copy="我改了文案") + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.confirm_copy("job-cc", req, authenticated_user=user, session=session) + + job.resume_from_copy_generated.assert_called_once_with(edited_copy="我改了文案") + mock_send.assert_called_once_with("worker.run_viral_video_render", args=["job-cc"]) + assert resp.id == "job-cc" + + def test_confirm_copy_rejects_wrong_status(self): + import pytest + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + from fastapi import HTTPException + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-cc2", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED) + repo = MagicMock() + repo.get.return_value = job + + with patch.object(vv_mod, "_get_job_repo", return_value=repo): + with pytest.raises(HTTPException) as exc: + vv_mod.confirm_copy("job-cc2", ConfirmCopyRequest(), authenticated_user=user, session=session) + assert exc.value.status_code == 409