feat(viral-video): v1.5 三步分步流水线 analyze-images/generate-copy/confirm-copy #2117
@@ -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": "任务失败",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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 "精选好物推荐"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user