feat(viral-video): v1.5 三步分步流水线 analyze-images/generate-copy/confirm-copy #2117

Merged
auto-approve-bot merged 1 commits from fix/2115-three-stage-pipeline into develop 2026-10-01 11:57:39 +08:00
11 changed files with 837 additions and 102 deletions
+175 -9
View File
@@ -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": "任务失败",
+67 -23
View File
@@ -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
+24
View File
@@ -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<ViralVideoJob>("/viral-video/analyze-images", payload).then((r) => r.data)
}
/** 阶段2:用户填完营销参数后生成文案+分镜+合规审核,完成后状态=copy_generated,返回 copy_result */
export function generateViralCopy(id: string, payload: GenerateCopyRequest) {
return apiClient
.post<ViralVideoJob>(`/viral-video/${id}/generate-copy`, payload)
.then((r) => r.data)
}
/** 阶段3:用户确认/编辑文案后开始 TTS→渲染→上传,完成后状态=completed */
export function confirmViralCopy(id: string, payload: ConfirmCopyRequest = {}) {
return apiClient
.post<ViralVideoJob>(`/viral-video/${id}/confirm-copy`, payload)
.then((r) => r.data)
}
+43
View File
@@ -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<string, unknown>
}
/** 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
}
+247 -59
View File
@@ -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()
@@ -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="")
@@ -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()
@@ -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()
)
+69 -8
View File
@@ -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 "精选好物推荐"
+7 -2
View File
@@ -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
+149
View File
@@ -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