Compare commits

..

1 Commits

Author SHA1 Message Date
saas前端工程师 3430c2fb6d fix(viral-video): 素材库音频选择弹窗修复+按钮禁用态可见性加固
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 21s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 58s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m24s
CI/CD Pipeline / Validate - Style (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 47s
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m53s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 3s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
AI Code Review / AI Code Review (pull_request) Successful in 6m38s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m47s
1. 修复「从素材库选择」配音弹窗列表只显示音符占位符不显示名称的问题:
   - 根因:.vv-thumb-ph(占位图标)在 voice 横向卡片中继承 width/height:100%,占满整行把 .vv-thumb-name 挤到视口外,用户只看到3个大号🎵图标看不到名称
   - 修复:为 .vv-asset-voice .vv-thumb-ph 设置固定 36px 圆形+浅紫底+紫色音符(与头像风格一致)
   - .vv-asset-voice 下的 img/video 也限 36px 圆形,避免缩略图过大
   - .vv-asset-voice .vv-thumb-name 改为 flex 布局+左对齐+透明底
   - 音频项名称后显示时长徽章(如 15s),与其他模块展示一致
   - 勾选标记 ✓ 从右上角移到垂直居中靠右(横向卡片更合理)

2. 空状态提示优化(voice kind):
   - 原先只显示「该素材库暂无音频素材」一行
   - 现在分两行:主提示「暂无配音素材」+ 副提示「请先在『配音/我的音色』中上传音频文件,或在素材库管理中添加」
   - 其他 kind(image/video)保持原样

3. 「选择内置音色」Modal 绑定代码复核:
   - 按钮 onClick={() => setVoicePickerOpen(true)} ✅
   - disabled={!canEditAssets} 在 STEP1 未上传时为 disabled(符合预期)
   - PresetVoicePickerModal 已在 JSX 末尾正确挂载,传 open/voices/selectedId/onClose/onConfirm ✅
   - vv-col-disabled 的 opacity 已通过 .vv-col-disabled .vv-btn-primary { opacity:1 !important } 覆盖,按钮在禁用列中仍用 vv-btn-disabled-gray 灰色样式可见

4. STEP2/STEP3 disabled 灰色按钮可见性加固:
   - vv-btn-disabled-gray 已设置 opacity:1 + 灰色#d1d5db背景+灰字#6b7280,不会被父级 opacity:0.55 冲淡
   - vv-col-disabled 下 .vv-btn-primary opacity:1 强制覆盖
   - disabled 主按钮清晰可见为灰色块状,和紫色激活态形成明显对比

tsc/vite build 全绿。
2026-10-01 11:39:34 +08:00
41 changed files with 2334 additions and 5659 deletions
-5
View File
@@ -212,14 +212,9 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
DOUBAO_API_KEY=your-doubao-api-key
DOUBAO_MODEL=doubao-seed-1-6-250615
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
DOUBAO_TIMEOUT=30
DOUBAO_MAX_RETRIES=2
# 视觉模型:pro 精度高,lite 速度快(viral-video 商品识别默认用 lite 提速)
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
DOUBAO_VISION_USE_LITE=true
# ==================== 积分/会员系统 (#1895) ====================
# 积分系统总开关:默认 false(暂停积分系统)。
@@ -1,51 +0,0 @@
"""viral video add copy_result + voice/video columns
Revision ID: 088_viral_video_copy_result
Revises: 087_viral_video_image_analysis
Create Date: 2026-10-01
v1.6 爆款视频字段补齐:
- copy_result JSON: 编导分镜脚本完整结构(overview/scene_and_lighting/shots/hard_constraints/negative_prompts/voiceover_script)
- voice_id/voice_source: TTS 音色参数
- video_ratio/video_model: Seedance 视频比例/模型
注意:线上启动也有幂等 ADD COLUMN 补列逻辑 (_ensure_viral_video_columns),本 migration 提供标准 Alembic 路径,
两套机制互不冲突(IF NOT EXISTS 等价行为)。
"""
import sqlalchemy as sa
from alembic import op
revision = "088_viral_video_copy_result"
down_revision = "087_viral_video_image_analysis"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 幂等添加列(通过单独执行 + 异常忽略兼容已由 backfill 补上的环境)
cols = [
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
conn = op.get_bind()
for name, ddl in cols:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
except Exception:
# 不支持 IF NOT EXISTS 的库(如老版本 SQLite)直接尝试 ADD COLUMN,失败则忽略
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
except Exception:
pass
def downgrade() -> None:
for name in ("copy_result", "video_model", "video_ratio", "voice_source", "voice_id"):
try:
op.drop_column("viral_video_jobs", name)
except Exception:
pass
-62
View File
@@ -1,62 +0,0 @@
"""viral video add storyboard + generated_copy_text (complement 088)
Revision ID: 089_viral_video_cols
Revises: 088_viral_video_copy_result
Create Date: 2026-10-01
#2129 兜底迁移:补齐 _VIRAL_VIDEO_BACKFILL_COLS 中所有列,覆盖
# watchtower 自动部署未跑历史 migration、且 AUTO_CREATE_SCHEMA=false 时
# _ensure_viral_video_columns 未执行的场景。
# 幂等 ADD COLUMN IF NOT EXISTS,已存在则跳过。
"""
import sqlalchemy as sa
from alembic import op
revision = "089_viral_video_cols"
down_revision = "088_viral_video_copy_result"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 扩展 alembic_version.version_num 字段长度(原来 VARCHAR(32) 装不下长 revision id)
conn = op.get_bind()
try:
conn.execute(sa.text("ALTER TABLE alembic_version ALTER COLUMN version_num TYPE VARCHAR(256)"))
except Exception:
pass
cols = [
("storyboard", "JSON"),
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
for name, ddl in cols:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
except Exception:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
except Exception:
pass
def downgrade() -> None:
for name in (
"copy_result",
"video_model",
"video_ratio",
"voice_source",
"voice_id",
"generated_copy_text",
"storyboard",
):
try:
op.drop_column("viral_video_jobs", name)
except Exception:
pass
@@ -1,35 +0,0 @@
"""viral video add phase_message column (#2134)
Revision ID: 090_viral_video_phase_msg
Revises: 089_viral_video_cols
Create Date: 2026-10-02
#2134 阶段细粒度提示:viral_video 表新增 phase_message 列(中文阶段提示文案)。
current_stage 列已在之前版本存在,本迁移只补 phase_message。
幂等 ADD COLUMN IF NOT EXISTS。
"""
import sqlalchemy as sa
from alembic import op
revision = "090_viral_video_phase_msg"
down_revision = "089_viral_video_cols"
branch_labels = None
depends_on = None
def upgrade() -> None:
# SQLite/PostgreSQL 兼容的幂等添加列
conn = op.get_bind()
inspector = sa.inspect(conn)
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
if "phase_message" not in cols:
op.add_column(
"viral_video_jobs",
sa.Column("phase_message", sa.String(length=500), nullable=False, server_default=""),
)
def downgrade() -> None:
op.drop_column("viral_video_jobs", "phase_message")
-49
View File
@@ -1,49 +0,0 @@
"""viral video add current_stage column (#2137 follow-up)
Revision ID: 091_viral_video_stage
Revises: 090_viral_video_phase_msg
Create Date: 2026-10-02
#2137 follow-up fix: 090 migration missed current_stage column on viral_video_jobs,
causing UndefinedColumn errors and 500s on all authenticated viral-video endpoints.
Idempotently add current_stage and double-check phase_message.
"""
import sqlalchemy as sa
from alembic import op
revision = "091_viral_video_stage"
down_revision = "090_viral_video_phase_msg"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
if "current_stage" not in cols:
op.add_column(
"viral_video_jobs",
sa.Column(
"current_stage",
sa.String(length=200),
nullable=False,
server_default="",
),
)
if "phase_message" not in cols:
op.add_column(
"viral_video_jobs",
sa.Column(
"phase_message",
sa.String(length=500),
nullable=False,
server_default="",
),
)
def downgrade() -> None:
op.drop_column("viral_video_jobs", "current_stage")
@@ -1,42 +0,0 @@
"""viral_video_jobs 增加 heartbeat_at 列(worker 心跳,用于僵尸任务超时回收)
Revision ID: 092_viral_video_heartbeat
Revises: 091_viral_video_stage
Create Date: 2026-10-02
"""
import sqlalchemy as sa
from alembic import op
revision = "092_viral_video_heartbeat"
down_revision = "091_viral_video_stage"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
if "heartbeat_at" not in cols:
op.add_column("viral_video_jobs", sa.Column("heartbeat_at", sa.DateTime(), nullable=True))
op.execute(
"UPDATE viral_video_jobs SET heartbeat_at = updated_at " "WHERE status = 'running' AND heartbeat_at IS NULL"
)
try:
op.create_index("ix_viral_video_jobs_heartbeat_at", "viral_video_jobs", ["heartbeat_at"])
except Exception:
pass
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
if "heartbeat_at" in cols:
try:
op.drop_index("ix_viral_video_jobs_heartbeat_at", table_name="viral_video_jobs")
except Exception:
pass
op.drop_column("viral_video_jobs", "heartbeat_at")
+2 -9
View File
@@ -9,11 +9,6 @@ from fastapi.responses import JSONResponse
router = APIRouter(tags=["Health"])
def _pg_url(url: str) -> str:
"""Convert SQLAlchemy URL (postgresql+psycopg://...) to libpq connection string."""
return url.replace("postgresql+psycopg://", "postgresql://", 1).replace("postgresql+psycopg2://", "postgresql://", 1)
@router.get("/health", status_code=status.HTTP_200_OK)
async def health_check():
return {
@@ -54,7 +49,7 @@ async def _check_database() -> dict:
"message": "Using in-memory database",
}
try:
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
@@ -129,7 +124,7 @@ async def _check_migrations() -> dict:
"message": "Using in-memory database, no migrations needed",
}
try:
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("""
SELECT COUNT(*) FROM information_schema.tables
@@ -142,5 +137,3 @@ async def _check_migrations() -> dict:
return {"status": "unhealthy", "message": f"Missing tables, found {count}/5"}
except Exception as error:
return {"status": "unhealthy", "message": f"Migration check failed: {error}"}
+16 -247
View File
@@ -1,21 +1,14 @@
"""爆款视频 API 路由。
v1.6 三步分步流水线端点(单次 Seedance 出片版):
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/copy_result 编导脚本)
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 进度推送
端点:
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)
"""
from __future__ import annotations
@@ -26,13 +19,10 @@ 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,
@@ -55,59 +45,6 @@ router = APIRouter()
# ── Helpers ──────────────────────────────────────────────────────────────
def _build_copy_result(job) -> dict | None:
"""v1.6: 返回编导分镜脚本 CopyResult 结构(给前端/Seedance 使用)。
- 若 job.copy_result 已持久化(v1.6 worker 生成),直接返回(补 final_copy 兜底)。
- 否则从老字段(generated_copy_text=口播, storyboard=分镜列表, intent_result)拼装兼容结构。
"""
cr = getattr(job, "copy_result", None)
if isinstance(cr, dict) and cr:
out = dict(cr)
# 向后兼容字段
voiceover = out.get("voiceover_script", "") or ""
out.setdefault("final_copy", voiceover)
out.setdefault("suggested_copy", voiceover)
out.setdefault("title", "")
return out
# 兼容 v1.5 老数据:storyboard 是老格式 [{order,type,description,text,duration,...}]
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
title = ""
if isinstance(intent, dict):
title = intent.get("suggested_title") or intent.get("intent", "") or ""
shots = []
for seg in sb:
if isinstance(seg, dict):
shots.append(
{
"time_range": "",
"shot_type_angle_movement": seg.get("ken_burns", ""),
"scene_and_dialogue": (seg.get("text") or "")
+ (" " + seg.get("description", "") if seg.get("description") else ""),
"action_details": "",
"audio_bgm": "",
"transition": seg.get("transition", "硬切"),
"reference_image_index": None,
}
)
ratio = getattr(job, "video_ratio", None) or "9:16"
return {
"overview": {"theme": title, "total_duration": getattr(job, "duration", 15), "aspect_ratio": ratio},
"scene_and_lighting": "",
"shots": shots,
"hard_constraints": ["无字幕", "无水印", "人物一致性"],
"negative_prompts": ["字幕", "水印", "错误文字", "五官崩坏"],
"voiceover_script": copy_text,
"final_copy": copy_text,
"suggested_copy": copy_text,
"title": title,
}
def _to_response(job) -> ViralVideoJobResponse:
return ViralVideoJobResponse(
id=job.id,
@@ -119,7 +56,7 @@ def _to_response(job) -> ViralVideoJobResponse:
viral_structure=job.viral_structure,
marketing_purpose=job.marketing_purpose,
bgm_preference=job.bgm_preference,
duration=job.duration or 15,
duration=job.duration,
user_copy_text=job.user_copy_text,
fusion_level=job.fusion_level,
reference_audio_path=job.reference_audio_path,
@@ -128,16 +65,6 @@ def _to_response(job) -> ViralVideoJobResponse:
style_guide=job.style_guide,
style_template_id=job.style_template_id,
status=job.status,
current_stage=getattr(job, "current_stage", "") or "",
phase_message=getattr(job, "phase_message", "") or "",
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),
voice_id=getattr(job, "voice_id", "") or "",
voice_source=getattr(job, "voice_source", "") or "",
video_ratio=getattr(job, "video_ratio", "9:16") or "9:16",
video_model=getattr(job, "video_model", "") or "",
intent_result=job.intent_result,
result_video_url=job.result_video_url,
credits_cost=job.credits_cost,
@@ -182,18 +109,13 @@ def create_viral_video(
viral_structure=request.viral_structure,
marketing_purpose=request.marketing_purpose,
bgm_preference=request.bgm_preference,
duration=request.duration or 15,
duration=request.duration,
user_copy_text=request.user_copy_text,
fusion_level=request.fusion_level,
reference_audio_path=request.reference_audio_path,
reference_video_url=request.reference_video_url,
style_strength=request.style_strength,
style_template_id=request.style_template_id,
voice_id=getattr(request, "voice_id", "") or "",
voice_source=getattr(request, "voice_source", "") or "",
video_ratio=getattr(request, "video_ratio", "9:16") or "9:16",
video_model=getattr(request, "video_model", "") or "",
copy_result=None,
)
# 持久化
@@ -211,139 +133,6 @@ 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",
voice_id=request.voice_id or "",
voice_source=request.voice_source or "",
video_ratio=request.video_ratio or "9:16",
video_model=request.video_model or "",
duration=request.duration or 15,
)
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.6 阶段2:用户填完营销参数后,跑 意图解析 → 编导分镜脚本生成 → 合规审核。
跑完后状态=copy_generated,响应 copy_result(含 overview/scene_and_lighting/shots/
hard_constraints/negative_prompts/voiceover_script),前端展示脚本与口播供用户编辑;
确认/编辑后调 /{id}/confirm-copy 进入阶段3(TTS + 单次 Seedance 出片)。
"""
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 = max(5, min(30, int(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
if request.style_guide is not None:
job.style_guide = request.style_guide
job.voice_id = request.voice_id or job.voice_id
job.voice_source = request.voice_source or job.voice_source
job.video_ratio = request.video_ratio or job.video_ratio or "9:16"
job.video_model = request.video_model or job.video_model or ""
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.6 阶段3:用户确认/编辑口播后开始 TTS + 单次 Seedance 生成 + 上传。"""
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,
@@ -400,42 +189,28 @@ def retry_viral_video_job(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
session: Session = Depends(get_db_session),
) -> ViralVideoJobResponse:
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。"""
from datetime import datetime, timezone
"""重试失败的爆款视频任务。"""
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="无权操作此任务")
# 判定是否为僵尸 running 任务:running 超过 10 分钟且心跳停止超过 2 分钟
now = datetime.now(timezone.utc)
is_stale_running = False
if job.status == ViralVideoStatus.RUNNING and job.started_at is not None:
hb = getattr(job, "heartbeat_at", None) or job.updated_at
if (now - job.started_at).total_seconds() > 10 * 60 and hb is not None and (now - hb).total_seconds() > 2 * 60:
is_stale_running = True
if job.status != ViralVideoStatus.FAILED and not is_stale_running:
raise HTTPException(status_code=409, detail="只有失败或超时的任务可以重试")
if job.status != ViralVideoStatus.FAILED:
raise HTTPException(status_code=409, detail="只有失败的任务可以重试")
# 重置状态
job.retry_count += 1
job.status = ViralVideoStatus.PENDING
job.error_msg = "" if not is_stale_running else "任务执行超时,已重置重试"
job.error_msg = ""
job.started_at = None
job.completed_at = None
job.current_stage = ""
job.phase_message = ""
job.heartbeat_at = None
repo.update(job)
# 重新入队
try:
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s", job.id, job.retry_count, is_stale_running)
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d", job.id, job.retry_count)
except Exception as e:
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
job.mark_failed(f"重试入队失败: {e}")
@@ -730,8 +505,6 @@ def _job_status(job) -> str:
_STATUS_STAGE = {
"pending": "",
"running": "",
"image_analyzed": "image_analysis",
"copy_generated": "review",
"wait_user_confirm": "intent_parsing",
"completed": "uploading",
"failed": "",
@@ -741,8 +514,6 @@ _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,
@@ -752,8 +523,6 @@ _STATUS_PROGRESS = {
_STATUS_MESSAGE = {
"pending": "任务已创建,等待执行",
"running": "任务执行中",
"image_analyzed": "图片分析完成,等待填写营销参数",
"copy_generated": "文案与分镜已生成,等待确认文案",
"wait_user_confirm": "等待用户确认意图文案",
"completed": "视频生成完成",
"failed": "任务失败",
+48 -151
View File
@@ -1,4 +1,4 @@
"""爆款视频 API schemas (v1.6 单次 Seedance 出片版)。"""
"""爆款视频 API schemas。"""
from __future__ import annotations
@@ -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")
@@ -14,174 +14,76 @@ VALID_STAGES = (
"image_analysis",
"video_analysis",
"intent_parsing",
"script_generation",
"copy_fusion",
"storyboard",
"review",
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
)
VALID_VIDEO_RATIOS = ("9:16", "16:9", "1:1", "4:3", "3:4", "21:9")
VALID_DURATIONS = (5, 10, 15, 20, 25, 30)
# -- 编导脚本结构(v1.6) --
class ShotScript(BaseModel):
"""逐镜头分镜。"""
time_range: str = Field(default="", description="时间区间,如 0-3秒")
shot_type_angle_movement: str = Field(default="", description="景别/角度/运镜,如『近景俯拍45度,缓慢推镜』")
scene_and_dialogue: str = Field(default="", description="场景描述+口播台词")
action_details: str = Field(default="", description="人物动作、表情、物品操作细节")
audio_bgm: str = Field(default="", description="环境音+BGM提示")
transition: str = Field(default="硬切", description="转场方式:硬切/淡入淡出/叠化")
reference_image_index: int | None = Field(
default=None, description="参考图片索引(0-based,对应上传的第几张产品图)"
)
class CopyResultOverview(BaseModel):
theme: str = ""
total_duration: int = 15
aspect_ratio: str = "9:16"
class CopyResult(BaseModel):
"""v1.6 编导分镜脚本结构(给前端 + Seedance 用)。"""
overview: CopyResultOverview = Field(default_factory=CopyResultOverview)
scene_and_lighting: str = ""
shots: list[ShotScript] = Field(default_factory=list)
hard_constraints: list[str] = Field(default_factory=list)
negative_prompts: list[str] = Field(default_factory=list)
voiceover_script: str = Field(
default="", description="纯口播对白,从各镜 scene_and_dialogue 的对白部分拼接,供 TTS 使用"
)
# 向后兼容:final_copy = voiceover_script
final_copy: str = ""
suggested_copy: str = ""
title: str = ""
# -- Request Schemas --
# ── Request Schemas ────────────────────────────────────────────────────────
class CreateViralVideoRequest(BaseModel):
"""旧接口:一键创建(保留兼容)。"""
"""创建爆款视频任务请求。"""
images: list[str] = Field(..., min_length=1, max_length=20)
industry: str = ""
target_customer: str = ""
persona_id: str = ""
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = Field(default=15, ge=5, le=30, description="视频时长(秒),5-30")
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
reference_video_url: str = ""
style_strength: str = "medium"
style_template_id: str = ""
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
images: list[str] = Field(..., min_length=1, max_length=20, description="产品图片 URL 列表")
industry: str = Field(default="", description="行业")
target_customer: str = Field(default="", description="目标客户描述")
persona_id: str = Field(default="", description="人设 ID")
viral_structure: str = Field(default="", description="爆款结构类型")
marketing_purpose: str = Field(default="", description="营销目的")
bgm_preference: str = Field(default="", description="BGM 偏好")
duration: int = Field(default=30, ge=5, le=180, description="视频时长(秒)")
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")
@field_validator("fusion_level")
@classmethod
def _v_fl(cls, v: str) -> str:
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 must be one of {VALID_FUSION_LEVELS}")
raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一")
return v
@field_validator("style_strength")
@classmethod
def _v_ss(cls, v: str) -> str:
def _validate_style_strength(cls, v: str) -> str:
if v not in VALID_STYLE_STRENGTHS:
raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}")
raise ValueError(f"style_strength 必须是 {VALID_STYLE_STRENGTHS} 之一")
return v
class AnalyzeImagesRequest(BaseModel):
"""v1.5+ 阶段1:创建任务 + 图片/视频分析。"""
images: list[str] = Field(..., min_length=1, max_length=30)
reference_video_url: str = ""
style_template_id: str = ""
style_strength: str = "medium"
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
duration: int = Field(default=15, ge=5, le=30)
class GenerateCopyRequest(BaseModel):
"""v1.5+ 阶段2:填完营销参数,生成编导脚本。"""
industry: str = ""
target_customer: str = ""
persona_id: str = ""
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = Field(default=15, ge=5, le=30)
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
reference_video_url: str = ""
style_strength: str = "medium"
style_template_id: str = ""
style_guide: dict | None = None
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
@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:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
class ConfirmIntentRequest(BaseModel):
"""旧 confirm-intent(兼容)。"""
"""确认意图请求(confirm-intent)。"""
confirmed_copy: str = ""
adjustments: str = ""
confirmed_copy: str = Field(default="", description="用户确认/修改后的文案,为空表示使用 AI 生成的文案")
adjustments: str = Field(default="", description="用户对 AI 文案的调整意见")
class AnalyzeStyleRequest(BaseModel):
"""触发参考视频风格分析请求。"""
reference_video_url: str = Field(..., description="参考视频 URL")
style_template_id: str = ""
style_template_id: str = Field(default="", description="风格模板 ID(可选覆盖)")
# -- Response Schemas --
# ── Response Schemas ───────────────────────────────────────────────────────
class ViralVideoJobResponse(BaseModel):
"""爆款视频任务响应(v1.6 包含 copy_result 编导脚本结构)。"""
"""爆款视频任务响应。"""
id: str
user_id: str
@@ -192,7 +94,7 @@ class ViralVideoJobResponse(BaseModel):
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = 15
duration: int = 30
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
@@ -201,21 +103,6 @@ class ViralVideoJobResponse(BaseModel):
style_guide: dict | None = None
style_template_id: str = ""
status: str
current_stage: str = (
"" # 细粒度阶段 snake_case(analyzing_images/parsing_intent/generating_script/reviewing/tts_synthesizing/rendering_video/uploading)
)
phase_message: str = "" # 中文阶段提示文案(前端轮询/SSE 直接展示)
image_analysis: dict | None = None
# v1.6 编导脚本(推荐前端使用)
copy_result: dict | None = None
# v1.5 兼容字段
storyboard: list | None = None
generated_copy_text: str = ""
# 音色/视频参数
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
intent_result: dict | None = None
result_video_url: str = ""
credits_cost: int = 0
@@ -228,11 +115,15 @@ class ViralVideoJobResponse(BaseModel):
class ViralVideoHistoryResponse(BaseModel):
"""历史记录列表响应。"""
items: list[ViralVideoJobResponse]
total: int
class StyleTemplateResponse(BaseModel):
"""风格模板响应。"""
id: str
name: str
description: str = ""
@@ -241,19 +132,25 @@ 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,9 +6,6 @@ import type {
ViralVideoJob,
ImageAnalysisResult,
CopyResult,
AnalyzeImagesRequest,
GenerateCopyRequest,
ConfirmCopyRequest,
} from "./types"
/** 创建爆款视频任务 */
@@ -108,24 +105,3 @@ 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)
}
+31 -152
View File
@@ -12,14 +12,6 @@ export const STYLE_STRENGTHS: { value: StyleStrength; label: string }[] = [
{ value: "strict", label: "像素级复刻" },
]
/** v1.6 前端时长下拉选项(5/10/15/20/25/30秒) */
export const VALID_DURATIONS = [5, 10, 15, 20, 25, 30] as const
export type VideoDuration = (typeof VALID_DURATIONS)[number]
/** v1.6 支持的画幅比例 */
export const VALID_RATIOS = ["9:16", "16:9", "1:1"] as const
export type VideoRatio = (typeof VALID_RATIOS)[number]
export type ViralVideoStatus =
| "pending"
| "running"
@@ -31,25 +23,39 @@ export type ViralVideoStatus =
| "cancelled"
/**
* v1.6 后端流水线阶段。单次 Seedance 出片版:
* image_analysis → video_analysis(可选) → intent_parsing → script_generation → review → tts → rendering → uploading
* 后端流水线阶段字符串。前端不展示逐阶段进度列表,仅保留类型
* 用于轮询时判断当前在哪个大阶段(分析中 vs 文案生成 vs 视频生成)。
*/
export type ViralVideoStage =
| "image_analysis"
| "video_analysis"
| "intent_parsing"
| "script_generation"
| "copy_fusion"
| "storyboard"
| "review"
| "tts"
| "bgm_select"
| "rendering"
| "musetalk"
| "uploading"
/** 图片+视频分析阶段:属于「分析图片」按钮的范围 */
const IMAGE_ANALYSIS_STAGES = new Set<ViralVideoStage>(["image_analysis", "video_analysis"])
/** 编导脚本阶段:属于「生成文案」按钮的范围 */
const COPY_STAGES = new Set<ViralVideoStage>(["intent_parsing", "script_generation", "review"])
/** 视频生成阶段:属于「开始生成视频」按钮的范围(v1.6: TTS+单次Seedance+上传) */
const VIDEO_STAGES = new Set<ViralVideoStage>(["tts", "rendering", "uploading"])
/** 文案相关阶段:属于「生成文案」按钮的范围 */
const COPY_STAGES = new Set<ViralVideoStage>([
"intent_parsing",
"copy_fusion",
"storyboard",
"review",
])
/** 视频相关阶段:属于「开始生成视频」按钮的范围 */
const VIDEO_STAGES = new Set<ViralVideoStage>([
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
])
export function isImageAnalysisStage(stage: ViralVideoStage | undefined): boolean {
return !!stage && IMAGE_ANALYSIS_STAGES.has(stage)
@@ -60,7 +66,7 @@ export function isCopyStage(stage: ViralVideoStage | undefined): boolean {
export function isVideoStage(stage: ViralVideoStage | undefined): boolean {
return !!stage && VIDEO_STAGES.has(stage)
}
/** 兼容旧调用:分析图片+生成文案 的所有前置阶段 */
/** 兼容旧调用:旧的 isAnalysisStage 视为「图片分析+文案」的所有前置阶段 */
export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
return isImageAnalysisStage(stage) || isCopyStage(stage)
}
@@ -68,20 +74,12 @@ export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
/** 单张图片 VLM 识别出的商品信息 */
export interface ImageProductAnalysis {
name?: string
category?: string
brand?: string
colors?: string[]
material_or_texture?: string
key_features?: string[]
visual_style?: string
scene?: string
target_audience_hint?: string
text_on_image?: string
/** 旧字段兼容 */
spec?: string
brand?: string
features?: string[] | string
label_text?: string
selling_points?: string
selling_points?: string[]
scene?: string
image_index?: number
}
@@ -89,50 +87,10 @@ export interface ImageAnalysisResult {
products?: ImageProductAnalysis[]
}
/** v1.6 编导分镜脚本 - 单镜头 */
export interface ShotScript {
/** 时间区间,如 "0-3秒" */
time_range?: string
/** 景别/角度/运镜,如 "近景俯拍45度,缓慢推镜" */
shot_type_angle_movement?: string
/** 场景描述+对白 */
scene_and_dialogue?: string
/** 人物动作/表情/物品操作细节 */
action_details?: string
/** 环境音+BGM提示 */
audio_bgm?: string
/** 转场方式(硬切/淡入淡出/叠化/结束) */
transition?: string
/** 参考图片索引(0-based,对应上传产品图数组) */
reference_image_index?: number | null
}
/** v1.6 编导分镜脚本 - 总览 */
export interface CopyResultOverview {
theme?: string
total_duration?: number
aspect_ratio?: string
}
/** v1.6 编导分镜脚本(核心输出结构,给 Seedance 做 prompt,给 TTS 取 voiceover_script) */
export interface CopyResult {
overview?: CopyResultOverview
/** 整体场景+光线描述 */
scene_and_lighting?: string
/** 逐镜头时间轴 */
shots?: ShotScript[]
/** 硬性约束(禁止字幕/水印/变形等) */
hard_constraints?: string[]
/** 负面提示词 */
negative_prompts?: string[]
/** 完整口播稿(纯文本,用于 TTS 合成) */
voiceover_script?: string
/** 向后兼容:= voiceover_script */
final_copy?: string
/** 向后兼容:= voiceover_script */
suggested_copy?: string
title?: string
/** v1.5 旧字段兼容(老数据降级时可能出现) */
scenes?: Array<{ shot: string; narration: string; duration?: number }>
}
@@ -146,18 +104,13 @@ export interface StyleTemplate {
}
export interface IntentResult {
intent?: string
key_messages?: string[]
tone?: string
target_emotion?: string
call_to_action?: string
product: string
selling_points: string[]
target_audience: string
tone: string
structure: string
duration: number
suggested_title?: string
/** v1.5 旧字段兼容 */
product?: string
selling_points?: string[]
target_audience?: string
structure?: string
duration?: number
suggested_copy?: string
}
@@ -170,9 +123,7 @@ export interface ViralVideoJob {
style_template_id?: string
style_guide?: string | Record<string, unknown>
user_copy_text?: string
/** v1.6: = copy_result.voiceover_script(从 copy_result 派生,向后兼容) */
final_copy_text?: string
generated_copy_text?: string
fusion_level?: FusionLevel
voice_id?: string
voice_mode?: "global" | "per_video"
@@ -180,17 +131,8 @@ export interface ViralVideoJob {
bgm_preference?: string
intent_result?: IntentResult
intent_text?: string
/** v1.6 编导分镜脚本(核心产物) */
copy_result?: CopyResult
/** 向后兼容:= copy_result.shots */
storyboard?: ShotScript[]
image_analysis?: ImageAnalysisResult
/** 视频比例:9:16 / 16:9 / 1:1,默认 9:16 */
video_ratio?: string
/** Seedance 模型 ID(空=后端默认) */
video_model?: string
/** 视频时长(秒,5-30,默认15) */
duration?: number
progress_stage?: ViralVideoStage
progress_percent?: number
progress_message?: string
@@ -220,7 +162,6 @@ export interface GenerateViralVideoRequest {
persona_id?: string
viral_structure?: string
marketing_purpose?: string
/** 视频时长(5-30秒,默认15) */
duration?: number
video_model?: string
video_ratio?: string
@@ -234,65 +175,3 @@ export interface HistoryResponse {
page: number
page_size: number
}
/** v1.6 阶段1请求:图片/视频分析(POST /viral-video/analyze-images) */
export interface AnalyzeImagesRequest {
images: string[]
reference_video_url?: string
style_template_id?: string
style_strength?: StyleStrength
/** TTS 音色 ID(STEP1 已选音色时传) */
voice_id?: string
/** 音色来源:preset | library | clone | upload */
voice_source?: "preset" | "library" | "clone" | "upload"
/** Seedance 视频比例:9:16 | 16:9 | 1:1 */
video_ratio?: string
/** Seedance 模型 ID(空则使用服务端默认) */
video_model?: string
/** 视频时长(秒,5-30,默认15) */
duration?: number
}
/** v1.6 阶段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
/** 视频时长(秒,5-30,默认15) */
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>
/** TTS 音色 ID(优先级高于 persona_id) */
voice_id?: string
/** 音色来源:preset | library | clone | upload */
voice_source?: "preset" | "library" | "clone" | "upload"
/** Seedance 视频比例(9:16/16:9/1:1 等) */
video_ratio?: string
/** Seedance 模型 ID(空则使用服务端默认) */
video_model?: string
}
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
export interface ConfirmCopyRequest {
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
edited_copy?: string
}
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
export interface StoryboardSegment {
order: number
type: string
description: string
text: string
duration: number
ken_burns?: string
transition?: string
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -34,63 +34,83 @@ interface Props {
/** 兜底 mock 音色(后端 /api/v1/tts/presets 返回字段不够时使用) */
const MOCK_VOICES: PresetVoice[] = [
// ⚠️ 兜底 mock,仅在 /voices/presets 接口不可达时使用;ID 必须与后端
// packages/domain/preset_voices.py PRESET_VOICES 的 voice_id 对齐(v3后缀)
{
id: "longxiaochun_v3",
id: "long-xiaochun",
name: "龙小淳",
gender: "female",
category: "女声",
desc: "知性积极女声,适合语音助手",
},
{
id: "longxiaoxia_v3",
id: "long-xiaoxia",
name: "龙小夏",
gender: "female",
category: "女声",
desc: "沉稳权威女声,适合新闻播报",
},
{
id: "longsanshu_v3",
id: "long-xiaoyan",
name: "龙小颜",
gender: "female",
category: "女声",
desc: "温柔甜美女声,适合情感口播",
},
{
id: "long-xiaotong",
name: "龙小彤",
gender: "female",
category: "女声",
desc: "活力少女音,适合短视频带货",
},
{
id: "long-sanshu",
name: "龙三叔",
gender: "male",
category: "男声",
desc: "沉稳质感男声,适合有声书",
},
{
id: "longyue_v3",
id: "long-xiaogang",
name: "龙小刚",
gender: "male",
category: "男声",
desc: "阳光活力男声,适合解说",
},
{
id: "long-xiaocheng",
name: "龙小诚",
gender: "male",
category: "男声",
desc: "磁性商务男声,适合品牌宣传",
},
{
id: "long-xiaozhi",
name: "龙小智",
gender: "child",
category: "童声",
desc: "可爱童声,适合亲子内容",
},
{
id: "long-yue",
name: "龙悦",
gender: "female",
category: "女声",
desc: "温暖磁性女声,适合广告配音",
category: "情绪",
desc: "温柔治愈女声,适合睡前/助眠",
},
{
id: "longshu_v3",
name: "龙书",
id: "long-xiaodong",
name: "龙晓东",
gender: "male",
category: "男声",
desc: "沉稳青年男声,适合教育讲解",
category: "方言",
desc: "东北方言男声,接地气",
},
{ id: "long-xiaoling", name: "龙小玲", gender: "female", category: "方言", desc: "粤语女声" },
{
id: "longyingjing_v3",
name: "龙应静",
id: "long-xiaoxiao-neural",
name: "晓晓",
gender: "female",
category: "女声",
desc: "低调冷静女声,适合纪录片解说",
},
{
id: "longshuo_v3",
name: "龙硕",
gender: "male",
category: "男声",
desc: "博才干练男声,适合科技类内容",
},
{
id: "longtian_v3",
name: "龙甜",
gender: "female",
category: "女声",
desc: "活泼女声,适合短视频配音",
desc: "温柔女声",
},
]
@@ -158,22 +178,23 @@ const PresetVoicePickerModal: React.FC<Props> = ({
// 合并真实数据和 mock:如果真实数据 gender/category 缺失,用 mock 兜底
const allVoices: PresetVoice[] = useMemo(() => {
// 真实 API 返回的 voice_id 以 API 为准(如 longxiaochun_v3),前端不做硬编码覆盖
const realList: PresetVoice[] = (voices || []).map((v) => {
// 按 id 精确匹配 mock 获取补充元信息(id 即 voice_id,唯一稳定键)
const mockMatch = MOCK_VOICES.find((m) => m.id === v.id)
// 按 name 模糊匹配 mock 获取补充信息
const mockMatch = MOCK_VOICES.find(
(m) => v.name?.includes(m.name.slice(1)) || m.name.includes(v.name?.slice(0, 2) || "___"),
)
return {
...v,
gender: v.gender || mockMatch?.gender,
category:
v.category ||
mockMatch?.category ||
(v.gender === "female" ? "女声" : v.gender === "male" ? "男声" : "其他"),
(v.gender === "female" ? "女声" : v.gender === "male" ? "男声" : undefined),
desc: v.desc || mockMatch?.desc,
sample_audio_url: v.sample_audio_url,
}
})
// 如果没有真实数据,使用兜底 mock(接口失败时)
// 如果没有真实数据,使用 mock
return realList.length > 0 ? realList : MOCK_VOICES
}, [voices])
-45
View File
@@ -1,45 +0,0 @@
import { describe, it, expect } from "vitest"
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
describe("api/errors", () => {
it("returns string error directly", () => {
expect(getErrorMessage("plain")).toBe("plain")
})
it("uses Error.message", () => {
expect(getErrorMessage(new Error("boom"))).toBe("boom")
})
it("returns fallback for empty/unknown", () => {
expect(getErrorMessage(null)).toBe("操作失败,请稍后重试")
expect(getErrorMessage(undefined, "f")).toBe("f")
})
it("reads axios-like response.data.detail", () => {
const err = { response: { data: { detail: "后端报错" } }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("后端报错")
})
it("reads axios-like response.data.message", () => {
const err = { response: { data: { message: "消息字段" } }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("消息字段")
})
it("HTTP 404 fallback", () => {
const err = { response: { status: 404, data: null }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("404")
})
it("HTTP 401 fallback", () => {
const err = { response: { status: 401, data: null }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("登录")
})
it("network error", () => {
const err = { request: {}, isAxiosError: true }
expect(getErrorMessage(err)).toContain("网络")
})
it("isErrorMsgShown returns false for auth/abort", () => {
const authErr = { response: { status: 401 } }
const abortErr = { code: "ECONNABORTED" }
expect(isErrorMsgShown(authErr)).toBe(false)
expect(isErrorMsgShown(abortErr)).toBe(false)
const e: any = new Error("x")
e.__msgShown = true
expect(isErrorMsgShown(e)).toBe(true)
expect(isErrorMsgShown(new Error("x"))).toBe(false)
})
})
-226
View File
@@ -1,226 +0,0 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import {
generateViralVideo,
getViralVideoJob,
confirmViralVideoIntent,
retryViralVideo,
getViralVideoHistory,
getViralStyleTemplates,
analyzeViralStyle,
mockImageAnalysis,
mockGenerateCopy,
analyzeViralImages,
generateViralCopy,
confirmViralCopy,
} from "@/api/viral-video"
import {
VALID_DURATIONS,
VALID_RATIOS,
isVideoStage,
isImageAnalysisStage,
isCopyStage,
isAnalysisStage,
} from "@/api/viral-video/types"
const mockGet = vi.fn()
const mockPost = vi.fn()
vi.mock("@/api/client", () => ({
default: {
get: (...args: unknown[]) => mockGet(...args),
post: (...args: unknown[]) => mockPost(...args),
},
}))
vi.mock("antd", () => ({ message: { error: vi.fn(), success: vi.fn() } }))
// 让 setTimeout 同步执行,避免测试等待 1.8s/2.2s
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
mockGet.mockResolvedValue({ data: {} })
mockPost.mockResolvedValue({ data: {} })
})
afterEach(() => {
vi.useRealTimers()
})
describe("viral-video constants & stage helpers", () => {
afterEach(() => {
vi.useRealTimers()
})
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
})
it("VALID_DURATIONS/VALID_RATIOS", () => {
expect(VALID_DURATIONS).toEqual([5, 10, 15, 20, 25, 30])
expect(VALID_RATIOS).toEqual(expect.arrayContaining(["9:16", "16:9", "1:1"]))
})
it("isVideoStage", () => {
expect(isVideoStage("tts")).toBe(true)
expect(isVideoStage("rendering")).toBe(true)
expect(isVideoStage("uploading")).toBe(true)
expect(isVideoStage("script_generation")).toBe(false)
expect(isVideoStage("completed")).toBe(false)
expect(isVideoStage(undefined)).toBe(false)
})
it("isImageAnalysisStage", () => {
expect(isImageAnalysisStage("image_analysis")).toBe(true)
expect(isImageAnalysisStage("video_analysis")).toBe(true)
expect(isImageAnalysisStage("script_generation")).toBe(false)
expect(isImageAnalysisStage(undefined)).toBe(false)
})
it("isCopyStage", () => {
expect(isCopyStage("intent_parsing")).toBe(true)
expect(isCopyStage("script_generation")).toBe(true)
expect(isCopyStage("review")).toBe(true)
expect(isCopyStage("tts")).toBe(false)
})
it("isAnalysisStage is union", () => {
expect(isAnalysisStage("image_analysis")).toBe(true)
expect(isAnalysisStage("script_generation")).toBe(true)
expect(isAnalysisStage("tts")).toBe(false)
expect(isAnalysisStage(undefined)).toBe(false)
})
})
describe("viral-video API wrappers", () => {
afterEach(() => {
vi.useRealTimers()
})
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
mockGet.mockResolvedValue({ data: {} })
mockPost.mockResolvedValue({ data: {} })
})
it("generateViralVideo", async () => {
mockPost.mockResolvedValue({ data: { id: "j1" } })
const r = generateViralVideo({ images: ["img1"] } as never)
vi.runAllTimersAsync()
expect(await r).toEqual({ id: "j1" })
expect(mockPost).toHaveBeenCalledWith("/viral-video/generate", { images: ["img1"] })
})
it("getViralVideoJob", async () => {
mockGet.mockResolvedValue({ data: { id: "j2" } })
const r = getViralVideoJob("j2")
vi.runAllTimersAsync()
expect(await r).toEqual({ id: "j2" })
expect(mockGet).toHaveBeenCalledWith("/viral-video/j2")
})
it("confirmViralVideoIntent", async () => {
mockPost.mockResolvedValue({ data: { id: "j3" } })
const r = confirmViralVideoIntent("j3", { confirmed_copy: "hi" })
vi.runAllTimersAsync()
await r
expect(mockPost).toHaveBeenCalledWith("/viral-video/j3/confirm-intent", {
confirmed_copy: "hi",
})
})
it("retryViralVideo", async () => {
mockPost.mockResolvedValue({ data: { id: "j4" } })
await retryViralVideo("j4")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j4/retry")
})
it("getViralVideoHistory", async () => {
mockGet.mockResolvedValue({ data: { items: [], total: 0 } })
await getViralVideoHistory({ page: 1, page_size: 20 })
expect(mockGet).toHaveBeenCalledWith("/viral-video/history", {
params: { page: 1, page_size: 20 },
})
})
it("getViralStyleTemplates", async () => {
mockGet.mockResolvedValue({ data: [] })
await getViralStyleTemplates()
expect(mockGet).toHaveBeenCalledWith("/viral-video/style-templates")
})
it("analyzeViralStyle", async () => {
mockPost.mockResolvedValue({ data: { id: "j5" } })
await analyzeViralStyle("j5")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j5/analyze-style")
})
it("analyzeViralImages", async () => {
mockPost.mockResolvedValue({ data: { id: "j6" } })
await analyzeViralImages({ images: ["a.png"] } as never)
expect(mockPost).toHaveBeenCalledWith("/viral-video/analyze-images", { images: ["a.png"] })
})
it("generateViralCopy", async () => {
mockPost.mockResolvedValue({ data: { id: "j7" } })
await generateViralCopy("j7", { duration: 15 } as never)
expect(mockPost).toHaveBeenCalledWith("/viral-video/j7/generate-copy", { duration: 15 })
})
it("confirmViralCopy", async () => {
mockPost.mockResolvedValue({ data: { id: "j8" } })
await confirmViralCopy("j8", { edited_copy: "xxx" })
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", { edited_copy: "xxx" })
mockPost.mockClear()
await confirmViralCopy("j8")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", {})
})
})
describe("viral-video client mocks", () => {
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
})
afterEach(() => {
vi.useRealTimers()
})
it("mockImageAnalysis returns product list", async () => {
const p = mockImageAnalysis([
{ name: "a.png" },
{ name: "b.jpg" },
{ name: "c.webp" },
{ name: "d.png" },
])
vi.advanceTimersByTime(2000)
const r = await p
expect(r.products).toHaveLength(3)
expect(r.products[0].image_index).toBe(0)
expect(r.products[0].brand).toBe("示例品牌")
expect(r.products[1].spec).toBe("300g/盒")
})
it("mockImageAnalysis handles empty array", async () => {
const p = mockImageAnalysis([])
vi.advanceTimersByTime(2000)
const r = await p
expect(r.products).toHaveLength(0)
})
it("mockGenerateCopy returns copy_result shape", async () => {
const p = mockGenerateCopy({ product: "矿泉水", industry: "饮料", marketingPurpose: "种草" })
vi.advanceTimersByTime(3000)
const r = await p
expect(r.title).toContain("种草")
expect(r.title).toContain("矿泉水")
expect(r.final_copy.length).toBeGreaterThan(50)
expect(r.suggested_copy).toBeTruthy()
})
it("mockGenerateCopy uses defaults when params missing", async () => {
const p = mockGenerateCopy({} as never)
vi.advanceTimersByTime(3000)
const r = await p
expect(r.title).toContain("品牌种草")
expect(r.final_copy).toContain("这款产品")
})
})
@@ -1,21 +0,0 @@
import { describe, it, expect } from "vitest"
import { getGenerationPhase } from "@/pages/generate/hooks/generate-video/phase"
describe("getGenerationPhase", () => {
it("returns 分析素材与配置 for p<20", () => {
expect(getGenerationPhase(0)).toEqual({ label: "分析素材与配置", icon: "🔍" })
expect(getGenerationPhase(19).label).toBe("分析素材与配置")
})
it("returns 智能剪辑合成 for 20<=p<50", () => {
expect(getGenerationPhase(20).label).toBe("智能剪辑合成")
expect(getGenerationPhase(49).label).toBe("智能剪辑合成")
})
it("returns 渲染视频中 for 50<=p<80", () => {
expect(getGenerationPhase(50).label).toBe("渲染视频中")
expect(getGenerationPhase(79).label).toBe("渲染视频中")
})
it("returns 即将完成 for p>=80", () => {
expect(getGenerationPhase(80)).toEqual({ label: "即将完成", icon: "✨" })
expect(getGenerationPhase(100).label).toBe("即将完成")
})
})
@@ -1,26 +0,0 @@
import { describe, it, expect, vi, afterEach } from "vitest"
import { formatDuration, formatFileSize, formatDate } from "@/pages/products/detailUtils"
describe("products/detailUtils", () => {
afterEach(() => {
vi.useRealTimers()
})
it("formatDuration", () => {
expect(formatDuration(0)).toBe("00:00")
expect(formatDuration(-1)).toBe("00:00")
expect(formatDuration(5)).toBe("00:05")
expect(formatDuration(65)).toBe("01:05")
expect(formatDuration(3600)).toBe("60:00")
})
it("formatFileSize MB/GB", () => {
expect(formatFileSize(0)).toBe("-")
expect(formatFileSize(-1)).toBe("-")
expect(formatFileSize(5.3)).toBe("5.3 MB")
expect(formatFileSize(2048)).toBe("2.00 GB")
})
it("formatDate returns zh-CN format", () => {
vi.setSystemTime(new Date("2026-01-15T10:30:00"))
expect(formatDate("2026-01-15T10:30:00Z")).toMatch(/2026/)
expect(formatDate("")).toBe("-")
})
})
@@ -1,54 +0,0 @@
import { describe, it, expect, beforeEach, vi, afterEach } from "vitest"
import { renderHook, act } from "@testing-library/react"
import { useViralVideoPolling } from "@/pages/viral-video/hooks/useViralVideoPolling"
const getViralVideoJobMock = vi.fn()
vi.mock("@/api/viral-video", () => ({
getViralVideoJob: (...args: unknown[]) => getViralVideoJobMock(...args),
}))
describe("useViralVideoPolling", () => {
beforeEach(() => {
vi.clearAllMocks()
vi.useFakeTimers()
})
afterEach(() => {
vi.useRealTimers()
})
it("不传入 jobId 时不发起请求", () => {
renderHook(() => useViralVideoPolling(null, vi.fn()))
expect(getViralVideoJobMock).not.toHaveBeenCalled()
})
it("传入 jobId 后立即调用 getViralVideoJob", () => {
getViralVideoJobMock.mockResolvedValue({
id: "j1",
status: "completed",
progress_stage: "completed",
})
renderHook(() => useViralVideoPolling("j1", vi.fn()))
expect(getViralVideoJobMock).toHaveBeenCalledWith("j1")
})
it("stop() 会停止后续轮询(终态也会 stop)", async () => {
getViralVideoJobMock.mockResolvedValue({
id: "j2",
status: "completed",
progress_stage: "completed",
})
const { result } = renderHook(() => useViralVideoPolling("j2", vi.fn(), { intervalMs: 50 }))
// 等第一次 promise 完成
await act(async () => {
await Promise.resolve()
await Promise.resolve()
})
// 终态后不会再调度新请求
const calls = getViralVideoJobMock.mock.calls.length
act(() => {
vi.advanceTimersByTime(2000)
})
expect(getViralVideoJobMock).toHaveBeenCalledTimes(calls)
expect(result.current.stop).toBeTypeOf("function")
})
})
@@ -1,35 +0,0 @@
import { describe, it, expect } from "vitest"
import {
genderLabel,
languageLabel,
genderClass,
formatTime,
formatFileSize,
} from "@/pages/voices/utils/format"
describe("voices utils/format", () => {
it("genderLabel returns label or falls back to value", () => {
expect(genderLabel("female")).toContain("女")
expect(genderLabel("male")).toContain("男")
expect(genderLabel("unknown" as never)).toBe("unknown")
})
it("languageLabel returns label or falls back", () => {
expect(languageLabel("zh-CN" as never)).toBeTruthy()
expect(languageLabel("xx-XX" as never)).toBe("xx-XX")
})
it("genderClass returns css class", () => {
expect(genderClass("female")).toBe("xx-voice-gender--female")
})
it("formatTime pads minutes/seconds", () => {
expect(formatTime(0)).toBe("00:00")
expect(formatTime(5)).toBe("00:05")
expect(formatTime(65)).toBe("01:05")
expect(formatTime(3600)).toBe("60:00")
})
it("formatFileSize human-readable", () => {
expect(formatFileSize(0)).toBe("0 B")
expect(formatFileSize(512)).toBe("512 B")
expect(formatFileSize(2048)).toBe("2.0 KB")
expect(formatFileSize(2 * 1024 * 1024)).toBe("2.0 MB")
})
})
+1 -2
View File
@@ -28,12 +28,11 @@ export default defineConfig({
"src/pages/editing-planner/EditingPlanner.tsx",
"src/pages/assets/AssetLibrary.tsx",
"src/pages/voice-materials/VoiceMaterialLibrary.tsx",
"src/pages/viral-video/ViralVideoPage.tsx",
],
// CI 覆盖率门禁(Phase 4 后提升,逐步逼近目标)
// 当前实际:行 ~62% / 分支 ~61% / 函数 ~25%
thresholds: {
lines: 49,
lines: 50,
branches: 50,
functions: 20,
},
File diff suppressed because it is too large Load Diff
-37
View File
@@ -104,43 +104,6 @@ Staging 当前可以保持 no-op;Production 开启前必须先验证 SMTP/Redi
---
## Staging 服务器 Docker 凭证配置
Staging 服务器(116.62.226.203)需要配置 ACR 和 Gitea Registry 凭证,否则 docker pull 和 Watchtower 自动更新会失败。
### 凭证文件位置
- Docker 配置文件:`/root/.docker/config.json`
- 包含两个 registry 的认证信息:
- `xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com`(阿里云 ACR)
- `git.xiaoxiajianji.com`(Gitea 容器镜像仓库)
### 服务器迁移后恢复步骤
```bash
# 1. 登录 ACR
docker login xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com -u <ACR_USERNAME>
# 2. 登录 Gitea Registry
docker login git.xiaoxiajianji.com -u xiaoxia -p <GITEA_REGISTRY_TOKEN>
# 3. 重启 Watchtower(确保挂载最新 config.json)
docker restart watchtower
```
### Watchtower 配置
- 容器名:`watchtower`
- 检查间隔:300 秒(5 分钟)
- 监控容器:`xiaoxia-api-staging`、`xiaoxia-worker-staging`、`xiaoxia-web-staging`
- 必须挂载 `-v /root/.docker/config.json:/config.json` 才能拉取私有镜像
- 必须挂载 `-v /var/run/docker.sock:/var/run/docker.sock` 才能管理容器
- 容器使用 `:dev` 稳定 tag,Watchtower 通过检测 `:dev` tag 的 digest 变化来发现更新
### 镜像 Tag 策略
- CI 每次构建推送三种 tag:`${GITHUB_SHA}`(精确版本)、`${GITHUB_REF_NAME}`(分支名)、`:dev`(滚动 tag,仅 develop 分支)
- Staging 容器统一使用 `:dev` tag 启动,确保 Watchtower 能自动发现新版本
- Migration(alembic)使用 commit SHA tag 执行,不依赖 Watchtower
---
## Gitea Actions 约定
- `develop` 分支触发 staging 部署。
+1 -5
View File
@@ -30,10 +30,6 @@ COPY deploy/configs/douyin_cookies.txt /app/configs/douyin_cookies.txt
# 强制升级 yt-dlp 到最新(抖音反爬经常变更,旧版 cookies 支持失效;#1968/#1963)
RUN pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com --upgrade "yt-dlp>=2026.8.19"
# API 启动入口(幂等迁移 + uvicorn)—— #2129: watchtower 自动部署兜底
COPY infra/docker/entrypoint-api.sh /usr/local/bin/entrypoint-api.sh
RUN chmod +x /usr/local/bin/entrypoint-api.sh
# 设置环境变量
ENV PATH="/opt/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
ENV PYTHONPATH=/app:/app/apps/api
@@ -45,4 +41,4 @@ HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"
# API 入口点
ENTRYPOINT ["/usr/local/bin/entrypoint-api.sh"]
CMD ["uvicorn", "apps.api.main:app", "--host", "0.0.0.0", "--port", "8000"]
-22
View File
@@ -1,22 +0,0 @@
#!/bin/bash
# API 启动入口:先幂等执行数据库迁移,再启动传入的 CMD(默认 uvicorn)
# 解决 watchtower 自动拉取新镜像后容器重启、未跑 alembic upgrade head 导致新列缺失 500 的问题(#2129)
set -e
cd /app
echo "[entrypoint-api] Running alembic upgrade head..."
if alembic upgrade head; then
echo "[entrypoint-api] Migrations ok."
else
echo "[entrypoint-api] WARNING: alembic upgrade failed, continuing (existing columns should be fine)..." >&2
fi
# 若有显式 CMD(CI 部署时 docker compose run --rm api sh -c '...' 传入),直接 exec 它
if [ "$#" -gt 0 ]; then
echo "[entrypoint-api] Exec custom command: $*"
exec "$@"
fi
echo "[entrypoint-api] Starting uvicorn..."
exec uvicorn apps.api.main:app --host 0.0.0.0 --port 8000
-11
View File
@@ -18,17 +18,6 @@
set -e
# #2129: 幂等执行数据库迁移(watchtower 自动部署兜底)
# worker 容器独立启动,不能依赖 API 容器先跑迁移
cd /app
echo "[entrypoint-worker] Running alembic upgrade head..."
if alembic upgrade head; then
echo "[entrypoint-worker] Migrations ok."
else
echo "[entrypoint-worker] WARNING: alembic upgrade failed, continuing to start workers..." >&2
fi
cd - >/dev/null
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
# ── 并发计算:显式 env 优先;否则从 WORKER_CONCURRENCY 按比例推导 ──
-1
View File
@@ -34,7 +34,6 @@ ENV APP_VERSION=$APP_VERSION
# 复制文件(按变化频率从低到高排序,最大化层缓存命中)
COPY alembic.ini /app/alembic.ini
COPY alembic/ /app/alembic/
COPY migrations/ /app/migrations/
COPY packages/ /app/packages/
# PR #1844 起,worker 还需要加载 apps.api.app.tasks.lipsync_tts,
@@ -943,23 +943,10 @@ class ViralVideoJobModel(Base):
style_strength = Column(String(20), nullable=False, default="medium")
style_guide = Column(JSON, nullable=True)
style_template_id = Column(String(36), nullable=False, default="", index=True)
# v1.5 音频/视频参数
voice_id = Column(String(200), nullable=False, default="")
voice_source = Column(String(20), nullable=False, default="")
video_ratio = Column(String(10), nullable=False, default="9:16")
video_model = Column(String(100), nullable=False, default="")
# 结果与状态
status = Column(String(30), nullable=False, default="pending", index=True)
current_stage = Column(String(200), nullable=False, default="") # 细粒度阶段 snake_case
phase_message = Column(String(500), nullable=False, default="") # 阶段中文提示文案
heartbeat_at = Column(DateTime, nullable=True, index=True) # worker 心跳,用于僵尸任务超时回收
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="")
copy_result = Column(
JSON, nullable=True
) # v1.6: 编导脚本结构{overview,scene_and_lighting,shots,hard_constraints,negative_prompts,voiceover_script}
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,41 +81,6 @@ def ensure_database_exists(database_url: str) -> None:
admin_engine.dispose()
_VIRAL_VIDEO_BACKFILL_COLS = [
("storyboard", "JSON"),
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
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。
@@ -135,5 +100,4 @@ 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()
@@ -24,7 +24,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
viral_structure=model.viral_structure or "",
marketing_purpose=model.marketing_purpose or "",
bgm_preference=model.bgm_preference or "",
duration=model.duration or 15,
duration=model.duration or 30,
user_copy_text=model.user_copy_text or "",
fusion_level=model.fusion_level or "ai_polish",
reference_audio_path=model.reference_audio_path or "",
@@ -32,19 +32,9 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
style_strength=getattr(model, "style_strength", "medium") or "medium",
style_guide=dict(model.style_guide) if model.style_guide else None,
style_template_id=getattr(model, "style_template_id", "") or "",
voice_id=getattr(model, "voice_id", "") or "",
voice_source=getattr(model, "voice_source", "") or "",
video_ratio=getattr(model, "video_ratio", "9:16") or "9:16",
video_model=getattr(model, "video_model", "") or "",
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
current_stage=getattr(model, "current_stage", "") or "",
phase_message=getattr(model, "phase_message", "") or "",
heartbeat_at=getattr(model, "heartbeat_at", None),
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 "",
copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None,
result_video_url=model.result_video_url or "",
credits_cost=model.credits_cost or 0,
error_msg=model.error_msg or "",
@@ -81,19 +71,9 @@ class SQLAlchemyViralVideoJobRepository:
style_strength=job.style_strength,
style_guide=job.style_guide,
style_template_id=job.style_template_id,
voice_id=job.voice_id,
voice_source=job.voice_source,
video_ratio=job.video_ratio,
video_model=job.video_model,
status=job.status,
current_stage=job.current_stage or "",
phase_message=job.phase_message or "",
heartbeat_at=job.heartbeat_at,
intent_result=job.intent_result,
image_analysis=job.image_analysis,
storyboard=job.storyboard,
generated_copy_text=job.generated_copy_text,
copy_result=job.copy_result,
result_video_url=job.result_video_url,
credits_cost=job.credits_cost,
error_msg=job.error_msg,
@@ -112,14 +92,8 @@ class SQLAlchemyViralVideoJobRepository:
if model is None:
raise ValueError(f"ViralVideoJob {job.id} not found")
model.status = job.status
model.current_stage = job.current_stage or ""
model.phase_message = job.phase_message or ""
model.heartbeat_at = job.heartbeat_at
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.copy_result = job.copy_result
model.result_video_url = job.result_video_url
model.credits_cost = job.credits_cost
model.error_msg = job.error_msg
@@ -127,24 +101,6 @@ 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.voice_id = job.voice_id or ""
model.voice_source = job.voice_source or ""
model.video_ratio = job.video_ratio or "9:16"
model.video_model = job.video_model or ""
model.updated_at = datetime.now(timezone.utc)
self.session.commit()
@@ -170,9 +126,7 @@ class SQLAlchemyViralVideoJobRepository:
self.session.query(ViralVideoJobModel)
.filter(
ViralVideoJobModel.user_id == user_id,
ViralVideoJobModel.status.in_(
["pending", "running", "wait_user_confirm", "image_analyzed", "copy_generated"]
),
ViralVideoJobModel.status.in_(["pending", "running", "wait_user_confirm"]),
)
.count()
)
+2 -5
View File
@@ -90,14 +90,11 @@ class SharedSettings(BaseSettings):
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
doubao_api_key: str = ""
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
doubao_model: str = "doubao-seed-1-6-250615"
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
doubao_timeout: int = 30
doubao_max_retries: int = 2
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
doubao_video_model: str = "doubao-seedance-2-5-260628"
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
+35 -102
View File
@@ -1,11 +1,10 @@
"""ViralVideoJob 领域模型 — 爆款视频任务.
v1.6 重大简化:Seedance 2.5 单次最长30秒,单次调用直接出片,不再分段/拼接/ffmpeg concat。
状态机(三步分步):
pending -> running -> image_analyzed -> running -> copy_generated -> running -> completed
wait_user_confirm -> running -> completed (旧路径兼容)
任意阶段 fail; 任意非终态 cancel.
failed -> pending (retry 重置后重跑)。
状态机:
pending → running → completed
↘ failed → pending (retry)
↘ cancelled
running 中可暂停:running → wait_user_confirm → running (confirm-intent resume)
"""
from __future__ import annotations
@@ -27,10 +26,10 @@ from uuid import uuid4
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"
@@ -38,32 +37,44 @@ class ViralVideoStatus(StrEnum):
class ViralVideoStage(StrEnum):
"""编排流水线阶段枚举(用于 WS 进度推送)。"""
IMAGE_ANALYSIS = "image_analysis"
VIDEO_ANALYSIS = "video_analysis"
INTENT_PARSING = "intent_parsing"
SCRIPT_GENERATION = "script_generation" # v1.6: 编导分镜脚本(融合原 copy_fusion+storyboard+review)
COPY_FUSION = "copy_fusion"
STORYBOARD = "storyboard"
REVIEW = "review"
TTS = "tts"
RENDERING = "rendering" # v1.6: 单次 Seedance 生成(BGM/音效/画面一次出片)
BGM_SELECT = "bgm_select"
RENDERING = "rendering"
MUSETALK = "musetalk"
UPLOADING = "uploading"
class FusionLevel(StrEnum):
"""文案融合级别。"""
AI_FULL = "ai_full"
AI_POLISH = "ai_polish"
USER_PRIMARY = "user_primary"
class StyleStrength(StrEnum):
"""风格强度。"""
LIGHT = "light"
MEDIUM = "medium"
STRICT = "strict"
class PromptType(StrEnum):
"""Prompt 模板类型(与 #2040 seed 对齐)。"""
IMAGE_ANALYSIS = "image_analysis"
INTENT_PARSING = "intent_parsing"
SCRIPT_GENERATION = "script_generation"
COPY_FUSION = "copy_fusion"
STORYBOARD = "storyboard"
REVIEW = "review"
VIDEO_STYLE_INTEGRATION = "video_style_integration"
STYLE_CONSTRAINT = "style_constraint"
@@ -75,17 +86,20 @@ STAGE_LABELS = {
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
ViralVideoStage.INTENT_PARSING: "意图解析",
ViralVideoStage.SCRIPT_GENERATION: "编导脚本生成",
ViralVideoStage.COPY_FUSION: "文案融合",
ViralVideoStage.STORYBOARD: "分镜脚本",
ViralVideoStage.REVIEW: "合规审核",
ViralVideoStage.TTS: "AI 配音",
ViralVideoStage.RENDERING: "视频生成",
ViralVideoStage.BGM_SELECT: "BGM 选择",
ViralVideoStage.RENDERING: "视频渲染",
ViralVideoStage.MUSETALK: "数字人口型",
ViralVideoStage.UPLOADING: "上传发布",
}
@dataclass
class ViralVideoJob:
"""爆款视频任务领域实体(v1.6 单次 Seedance 出片版)。"""
"""爆款视频任务领域实体。"""
user_id: str
images: list[str] = field(default_factory=list)
@@ -95,31 +109,21 @@ class ViralVideoJob:
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = 15 # v1.6: 默认15秒,上限30秒(Seedance 2.5 单次最大30s)
duration: int = 30
user_copy_text: str = ""
fusion_level: str = FusionLevel.AI_POLISH
reference_audio_path: str = ""
# v1.3
reference_video_url: str = ""
style_strength: str = StyleStrength.MEDIUM
style_guide: dict | None = None
style_template_id: str = ""
# v1.5.1 音频/视频参数
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
# v1.4+ 产物
# v1.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜)
image_analysis: dict | None = None
intent_result: dict | None = None
generated_copy_text: str = "" # v1.6: 存 voiceover_script(纯口播对白),字段名兼容
storyboard: list | None = None # v1.6: 存 copy_result.shots,字段名兼容
copy_result: dict | None = None # v1.6: 完整编导脚本结构
# 状态
id: str = field(default_factory=lambda: uuid4().hex)
status: ViralVideoStatus = ViralVideoStatus.PENDING
current_stage: str = "" # 细粒度阶段(ViralVideoStage.value,snake_case)
phase_message: str = "" # 阶段中文提示文案,前端轮询直接展示
heartbeat_at: datetime | None = None # worker 心跳时间,用于超时僵尸任务检测
intent_result: dict | None = None
result_video_url: str = ""
credits_cost: int = 0
error_msg: str = ""
@@ -129,54 +133,13 @@ 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.IMAGE_ANALYZED,
ViralVideoStatus.COPY_GENERATED,
ViralVideoStatus.WAIT_USER_CONFIRM,
ViralVideoStatus.RUNNING,
):
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
raise ValueError(f"Cannot transition from {self.status} to running")
self.status = ViralVideoStatus.RUNNING
now = datetime.now(timezone.utc)
if self.started_at is None:
self.started_at = now
self.heartbeat_at = now
self.updated_at = now
def touch_heartbeat(self) -> None:
"""更新心跳时间(worker 在长任务中周期性调用,用于超时检测)。"""
now = datetime.now(timezone.utc)
if self.started_at is None:
self.started_at = now
self.heartbeat_at = now
self.updated_at = now
def mark_image_analyzed(self) -> None:
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_result: dict) -> None:
"""v1.6 阶段2完成:编导脚本(含 voiceover_script/shots/硬约束/负面词)已生成。"""
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.copy_result = copy_result or {}
if isinstance(copy_result, dict):
self.generated_copy_text = copy_result.get("voiceover_script", "") or ""
shots = copy_result.get("shots") or []
self.storyboard = list(shots) if isinstance(shots, list) else []
self.started_at = datetime.now(timezone.utc)
self.updated_at = datetime.now(timezone.utc)
def mark_wait_user_confirm(self, intent_result: dict) -> None:
@@ -186,25 +149,6 @@ class ViralVideoJob:
self.intent_result = intent_result
self.updated_at = datetime.now(timezone.utc)
def resume_from_image_analyzed(self, **kwargs) -> None:
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:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
if self.status != ViralVideoStatus.COPY_GENERATED:
raise ValueError(f"Cannot resume from {self.status} to render")
if edited_copy and isinstance(self.copy_result, dict):
self.copy_result = {**self.copy_result, "voiceover_script": 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:
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
raise ValueError(f"Cannot resume from {self.status}")
@@ -237,14 +181,3 @@ class ViralVideoJob:
ViralVideoStatus.FAILED,
ViralVideoStatus.CANCELLED,
)
@property
def effective_copy_text(self) -> str:
"""TTS 用的最终口播文案:优先 copy_result.voiceover_script,兼容老字段。"""
if isinstance(self.copy_result, dict) and self.copy_result.get("voiceover_script"):
return self.copy_result["voiceover_script"]
return self.generated_copy_text or self.user_copy_text or "你好,给大家推荐一款好物"
@property
def voiceover_script(self) -> str:
return self.effective_copy_text
+84 -167
View File
@@ -40,8 +40,6 @@ class DoubaoClient:
self.timeout: int = settings.doubao_timeout
self.max_retries: int = settings.doubao_max_retries
self.vision_model: str = settings.doubao_vision_model
self.vision_lite_model: str = settings.doubao_vision_lite_model
self.fast_model: str = settings.doubao_fast_model
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
@@ -94,7 +92,6 @@ class DoubaoClient:
messages: list[dict[str, str]],
temperature: float = 0.7,
max_tokens: int = 1024,
model: str | None = None,
) -> Optional[str]:
"""调用 Chat Completion 接口.
@@ -115,7 +112,7 @@ class DoubaoClient:
"Content-Type": "application/json",
}
payload: dict[str, Any] = {
"model": model or self.model,
"model": self.model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
@@ -157,12 +154,11 @@ class DoubaoClient:
max_tokens: int = 2048,
temperature: float = 0.3,
timeout: int | None = None,
model: str | None = None,
) -> Optional[str]:
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
将 images 附加到最后一条 user message 的 content 中,
使用 vision_model(默认 doubao-1-5-vision-pro-250328)。
使用 vision_model(默认 doubao-1-5-vision-pro-250915)。
Args:
messages: 对话消息列表。最后一条 user message 会被注入图片内容。
@@ -208,7 +204,7 @@ class DoubaoClient:
"Content-Type": "application/json",
}
payload: dict[str, Any] = {
"model": model or self.vision_model,
"model": self.vision_model,
"messages": vision_messages,
"temperature": temperature,
"max_tokens": max_tokens,
@@ -254,21 +250,24 @@ class DoubaoClient:
duration: int = 5,
ratio: str | None = "9:16",
resolution: str = "720p",
generate_audio: bool = True,
generate_audio: bool = False,
watermark: bool = False,
output_dir: str | None = None,
model: str | None = None,
reference_images: list[str] | None = None,
reference_audios: list[str] | None = None,
reference_videos: list[str] | None = None,
) -> str | None:
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
"""调用 Seedance 2.5 文生/图生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
- 首帧图(first_frame 模式)Seedance 2.5 强制 ratio=adaptive;走 omni_reference(参考生视频)模式时才能指定 9:16/1:1 等具体比例。
判定:传了参考音频/视频或 ≥1 张多参考图时,走 omni_reference(首张图 role=reference_image);纯首帧无参考时走 first_frame(ratio 强制 adaptive)。
- 创建任务若因 ratio 报错(HTTP 400),自动回退到 ratio=adaptive 重试一次。
Args:
prompt: 文本提示词
image_url: 首帧参考图 URL(可选,提供则走图生视频)
duration: 视频时长 2~30 秒,默认 5
ratio: 宽高比 16:9/9:16/1:1/4:3/3:4/21:9/adaptive
resolution: 480p/720p/1080p
generate_audio: 是否生成模型自带音效(默认 False,我们自己混 TTS)
watermark: 是否加水印
output_dir: 下载目录,默认 /tmp
Returns:
本地 MP4 文件路径,失败返回 None。
"""
if not self.is_available:
return None
@@ -277,58 +276,24 @@ class DoubaoClient:
settings = get_shared_settings()
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
# 收紧总超时:轮询 8min + 下载 2min = 最长 ~10min,防止出现 20min 卡死
total_timeout = getattr(settings, "doubao_video_timeout", 480) or 480
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
video_model = model or default_video_model
total_timeout = getattr(settings, "doubao_video_timeout", 600) or 600
video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
ref_imgs = [u for u in (reference_images or [])[:9] if u and isinstance(u, str)]
# 判断任务模式:有参考音/视/多图 → omni_reference(支持指定 ratio);纯首帧 → first_frame(ratio=adaptive)
has_extra_refs = bool(ref_audios or ref_videos or ref_imgs)
is_first_frame_mode = bool(image_url) and not has_extra_refs
# 最终 ratio:first_frame 模式强制 adaptive,否则按用户传值(默认 9:16)
final_ratio = "adaptive" if is_first_frame_mode else (ratio or "9:16")
# 构造 content 数组:text + 图 + 音 + 视
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
if image_url:
if has_extra_refs:
# omni_reference:首张图作为 reference_image,允许指定 ratio
content.append(
{
"type": "image_url",
"image_url": {"url": image_url},
"role": "reference_image",
}
)
else:
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
content.append(
{
"type": "image_url",
"image_url": {"url": image_url},
"role": "first_frame",
}
)
for u in ref_imgs:
content.append({"type": "image_url", "image_url": {"url": u}, "role": "reference_image"})
for u in ref_audios:
content.append({"type": "audio_url", "audio_url": {"url": u}, "role": "reference_audio"})
for u in ref_videos:
content.append({"type": "video_url", "video_url": {"url": u}, "role": "reference_video"})
content.append({"type": "image_url", "image_url": {"url": image_url}})
create_payload: dict[str, Any] = {
"model": video_model,
"content": content,
"generate_audio": bool(generate_audio),
"generate_audio": generate_audio,
"duration": int(duration),
"resolution": resolution,
"watermark": bool(watermark),
"ratio": final_ratio,
"watermark": watermark,
}
# Bug #2110: ratio=None 时不传(首帧图生视频跟随原图比例,传 ratio 会 400 InvalidParameter)
if ratio:
create_payload["ratio"] = ratio
headers = {
"Authorization": f"Bearer {self.api_key}",
@@ -336,94 +301,69 @@ class DoubaoClient:
}
create_url = f"{self.base_url}/contents/generations/tasks"
logger.info(
"Seedance 创建任务: model=%s dur=%ds ratio=%s mode=%s gen_audio=%s img=%d aud=%d vid=%d",
"Seedance 创建任务请求: url=%s model=%s duration=%ds ratio=%s gen_audio=%s image_url=%s",
create_url,
video_model,
duration,
final_ratio,
"first_frame" if is_first_frame_mode else "omni_ref",
ratio or "(follow-image)",
generate_audio,
(1 if image_url else 0) + len(ref_imgs),
len(ref_audios),
len(ref_videos),
bool(image_url),
)
def _do_create(payload: dict) -> tuple[str | None, Exception | None, int, str]:
"""返回 (task_id, last_err, status_code, body_text)。"""
last_err: Exception | None = None
for attempt in range(self.max_retries + 1):
# 1) 创建任务(带重试)
task_id: str | None = None
last_error: Exception | None = None
for attempt in range(self.max_retries + 1):
try:
resp = httpx.post(create_url, headers=headers, json=create_payload, timeout=self.timeout)
# 测试环境下 MagicMock().status_code 是 MagicMock,与 int 比较会抛 TypeError;
# 用显式 int() 转换+类型判断,避免误判。
try:
resp = httpx.post(create_url, headers=headers, json=payload, timeout=self.timeout)
sc = int(getattr(resp, "status_code", 0) or 0)
body = (getattr(resp, "text", "") or "")[:1500]
if sc >= 400:
logger.error("Seedance 创建任务 HTTP %d: body=%s", sc, body)
try:
resp.raise_for_status()
except Exception as ee:
last_err = ee
if attempt < self.max_retries:
time.sleep(0.5 * (2**attempt))
continue
return None, last_err, sc, body
data = resp.json()
tid = data.get("id")
if tid:
return tid, None, sc, body
last_err = RuntimeError(f"create ok but no id: {str(data)[:300]}")
except Exception as e:
last_err = e
if attempt < self.max_retries:
wait = 0.5 * (2**attempt)
logger.warning(
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
wait,
attempt + 1,
self.max_retries + 1,
e,
)
time.sleep(wait)
return None, last_err, 0, ""
# 第一次尝试
task_id, last_err, sc, body = _do_create(create_payload)
# ratio 兜底:HTTP 400 且 body 提到 ratio / adaptive → 回退 adaptive 再试一次
if (
not task_id
and sc == 400
and final_ratio != "adaptive"
and (
"ratio" in (body or "").lower()
or "aspect" in (body or "").lower()
or "adaptive" in (body or "").lower()
)
):
logger.warning("Seedance 创建因 ratio 失败,回退 ratio=adaptive 重试")
create_payload["ratio"] = "adaptive"
task_id, last_err, sc2, body2 = _do_create(create_payload)
_status = int(resp.status_code)
except (TypeError, ValueError):
_status = 200
if _status >= 400:
# 把响应体完整打出来(通常含 error.code/message,能直接定位:模型未开通/Key 无权限/模型 ID 错误)
logger.error(
"Seedance 创建任务 HTTP %d: body=%s",
resp.status_code,
(resp.text or "")[:1000],
)
resp.raise_for_status()
data = resp.json()
task_id = data.get("id")
if task_id:
break
last_error = RuntimeError(f"create task returned no id: {str(data)[:200]}")
except Exception as e:
last_error = e
if attempt < self.max_retries:
wait = 0.5 * (2**attempt)
logger.warning(
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s", wait, attempt + 1, self.max_retries + 1, e
)
time.sleep(wait)
if not task_id:
logger.error(
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s body=%s 【排查】"
"1) 方舟控制台已开通 doubao-seedance-2-5-260628;2) API Key 有该模型权限;"
"3) DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3;4) 参考素材 URL 公网可访问。",
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s 【排查建议】"
"1) 确认方舟控制台已开通 Doubao-Seedance-2.5 模型;"
"2) DOUBAO_API_KEY 对应的账号有该模型调用权限;"
"3) DOUBAO_BASE_URL 必须为 https://ark.cn-beijing.volces.com/api/v3;"
"4) 若控制台用「推理接入点」(endpoint),请把 DOUBAO_VIDEO_MODEL 改为 ep-xxx 接入点 ID。",
video_model,
self.base_url,
last_err,
(body or "")[:500],
last_error,
)
return None
logger.info("Seedance 任务已创建: task_id=%s ratio=%s", task_id, create_payload["ratio"])
logger.info("Seedance 任务已创建: task_id=%s model=%s duration=%ds", task_id, video_model, duration)
# 2) 轮询状态
poll_url = f"{create_url}/{task_id}"
deadline = time.time() + total_timeout
video_url: str | None = None
last_status: str = "queued"
poll_count = 0
while time.time() < deadline:
poll_count += 1
try:
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
try:
@@ -438,71 +378,48 @@ class DoubaoClient:
content_obj = data.get("content") or {}
video_url = content_obj.get("video_url")
if video_url:
logger.info("Seedance 任务成功: task_id=%s polls=%d", task_id, poll_count)
break
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
last_error = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
break
if status == "failed":
err = data.get("error") or {}
last_err = RuntimeError(f"task failed: code={err.get('code','')} msg={err.get('message','')}")
logger.error("Seedance 任务失败 task_id=%s: %s", task_id, last_err)
last_error = RuntimeError(f"task failed: {err.get('code','')} {err.get('message','')}")
break
if status in ("expired", "cancelled"):
last_err = RuntimeError(f"task {status}")
logger.error("Seedance 任务 %s: task_id=%s", status, task_id)
last_error = RuntimeError(f"task {status}")
break
# 每 5 次轮询打一次 info 日志,便于观察进度
if poll_count % 5 == 0:
logger.info("Seedance 轮询中: task_id=%s status=%s polls=%d", task_id, status, poll_count)
# queued / running: 继续轮询
except httpx.HTTPStatusError as e:
last_err = e
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[:300])
last_error = e
logger.warning(
"Seedance 轮询 HTTP %d: body=%s",
e.response.status_code,
(e.response.text or "")[:500],
)
except Exception as e:
last_err = e
last_error = e
logger.debug("Seedance 轮询异常: %s", e)
time.sleep(poll_interval)
if not video_url:
logger.error(
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
task_id,
last_status,
poll_count,
last_err,
total_timeout,
)
logger.error("Seedance 任务未成功: task_id=%s status=%s err=%s", task_id, last_status, last_error)
return None
# 3) 下载到本地(下载超时收紧到 120s)
# 3) 下载到本地
try:
out_dir = output_dir or "/tmp"
os.makedirs(out_dir, exist_ok=True)
local_path = f"{out_dir}/seedance_{task_id}_{uuid.uuid4().hex[:8]}.mp4"
download_timeout = 120.0
logger.info(
"Seedance 开始下载: task_id=%s url=%s timeout=%.0fs", task_id, video_url[:120], download_timeout
)
with httpx.stream("GET", video_url, timeout=download_timeout) as r:
with httpx.stream("GET", video_url, timeout=300) as r:
r.raise_for_status()
downloaded = 0
with open(local_path, "wb") as f:
for chunk in r.iter_bytes(chunk_size=1024 * 256):
if chunk:
f.write(chunk)
downloaded += len(chunk)
size = os.path.getsize(local_path)
logger.info("Seedance 视频下载完成: %s size=%d bytes", local_path, size)
if size == 0:
logger.error("Seedance 下载文件大小为 0")
try:
os.remove(local_path)
except Exception:
pass
return None
logger.info("Seedance 视频下载完成: %s (%d bytes)", local_path, os.path.getsize(local_path))
return local_path
except Exception as e:
logger.error("Seedance 视频下载失败: %s", e, exc_info=True)
logger.error("Seedance 视频下载失败: %s", e)
return None
+29 -73
View File
@@ -496,32 +496,16 @@ def run_generate_cover(
# ── 通用 LLM / Vision 调用(#2039 ViralVideoOrchestrator 使用,复用现有豆包客户端)──
def call_llm(
prompt: str,
temperature: float = 0.7,
max_tokens: int = 2048,
model: str | None = None,
system_prompt: str | None = None,
) -> object:
"""调用豆包大模型(文本对话),返回解析后的 JSON(dict/list)或原文字符串;失败返回 None。
Args:
prompt: 用户侧提示。
temperature: 采样温度。
max_tokens: 输出上限(结构化任务默认 2048,长文案可按需加大)。
model: 覆盖默认模型(如 fast_model 提速用),None 走配置默认推理模型。
system_prompt: 覆盖默认 system prompt。
"""
def call_llm(prompt: str, temperature: float = 0.7) -> object:
"""调用豆包大模型(文本对话),返回解析后的 JSON(dict/list)或原文字符串;失败返回 None。"""
client = get_doubao_client()
if not client.is_available:
return None
if system_prompt is None:
system_prompt = "你是专业的短视频内容策划助手。需要结构化输出时请严格使用 JSON。"
messages = [
{"role": "system", "content": system_prompt},
{"role": "system", "content": "你是专业的短视频内容策划助手。需要结构化输出时请严格使用 JSON。"},
{"role": "user", "content": prompt},
]
raw = client.chat_completion(messages, temperature=temperature, max_tokens=max_tokens, model=model)
raw = client.chat_completion(messages, temperature=temperature, max_tokens=4096)
if raw is None:
return None
try:
@@ -530,26 +514,14 @@ def call_llm(
return raw
def call_vision(
image_url: str,
prompt: str,
*,
model: str | None = None,
max_tokens: int = 1024,
temperature: float = 0.2,
timeout: int = 45,
system_prompt: str | None = None,
) -> object:
def call_vision(image_url: str, prompt: str) -> object:
"""调用豆包视觉大模型分析图片,返回解析后的 JSON 或原文字符串;失败返回 None。
Args:
image_url: 可公网访问的图片 URL(直接传给豆包视觉模型,无需本地下载)。
prompt: 用户侧文本提示。
model: 覆盖默认视觉模型(如 vision_lite_model 提速用),None 走配置默认。
max_tokens: 输出上限,商品识别用 800~1200 足够,避免长输出拖慢首 token。
temperature: 温度。
timeout: 单次请求超时(秒)。
system_prompt: 覆盖默认 system prompt(viral-video 商品分析会传专门的详细 prompt)。
Bug #2114 (VLM 牛头不对马嘴根因修复):
之前误走 client.chat_completion(用文本模型 doubao-seed-1.6),多模态 content list 被当成
纯文本发给文本模型 → 模型要么看不到图、要么抛 400,静默被 except 吞掉 → 返回 None →
_step_image_analysis fallback 到 {"name":"未识别"} → 后续文案/分镜完全没图的信息。
现改走 vision_completion,走视觉模型 doubao-1-5-vision-pro-250915。
"""
client = get_doubao_client()
if not client.is_available:
@@ -559,33 +531,28 @@ def call_vision(
logger.warning("[call_vision] 空 image_url,跳过视觉分析")
return None
if system_prompt is None:
system_prompt = (
"你是资深电商视觉分析师。请严格基于用户提供的图片观察回答,"
"图片里没有的信息不要凭空想象或编造;看不清或无法判断时明确说"
"「无法判断」,不要猜测。输出必须是严格 JSON,不要附加 Markdown 或解释文字。"
)
system_prompt = (
"你是资深电商视觉分析师。请严格基于用户提供的图片观察回答,"
"图片里没有的信息不要凭空想象或编造;看不清或无法判断时明确说"
"「图片中无法判断」,不要猜测。输出必须是严格 JSON,不要附加 Markdown 或解释文字。"
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
]
used_model = model or getattr(client, "vision_model", "?")
logger.info(
"[call_vision] 调用豆包视觉模型 model=%s image_url=%s prompt_len=%d max_tokens=%d timeout=%d",
used_model,
"[call_vision] 调用豆包视觉模型 vision_model=%s image_url=%s prompt_len=%d",
getattr(client, "vision_model", "?"),
image_url[:120],
len(prompt),
max_tokens,
timeout,
)
raw = client.vision_completion(
messages=messages,
images=[image_url],
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
model=model,
temperature=0.2,
max_tokens=2048,
timeout=60,
)
if raw is None:
logger.warning("[call_vision] 视觉模型返回 None (image_url=%s)", image_url[:80])
@@ -608,43 +575,32 @@ def call_video_generation(
prompt: str,
*,
image_url: str | None = None,
duration: int = 15,
duration: int = 5,
ratio: str | None = "9:16",
resolution: str = "720p",
output_dir: str | None = None,
model: str | None = None,
generate_audio: bool = True,
reference_images: list[str] | None = None,
reference_audios: list[str] | None = None,
reference_videos: list[str] | None = None,
) -> str | None:
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版),返回本地 MP4 路径;失败返回 None。
"""调用 Seedance 2.5 生成视频段,返回本地 MP4 路径;失败返回 None。
v1.6.1 关键约束(避免 20min 卡死):
- 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image;
- 纯首帧无参考时(first_frame 模式),Seedance 2.5 强制 ratio=adaptive;
传了参考音/视/多图时走 omni_reference 模式,ratio 可指定为 9:16(客户端内部自动判断)。
- ratio 默认 9:16(竖屏),客户端会根据是否有参考自动在 first_frame/adaptive 与 omni/9:16 间切换;
若创建任务因 ratio 报错(HTTP 400),客户端会自动回退到 adaptive 再试一次。
封装 ai_client.video_generation:提交异步任务→轮询→下载到本地。
Bug #2110: 首帧参考图模式下不传 ratio(API 要求跟随首帧图比例,传 ratio=9:16
会返回 400 InvalidParameter)。
"""
client = get_doubao_client()
if not client.is_available:
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
return None
effective_ratio = ratio or "9:16"
# 首帧模式:不强制 ratio,让模型跟随首帧图比例
effective_ratio = None if image_url else ratio
try:
kwargs: dict = dict(
prompt=prompt,
image_url=image_url,
duration=int(duration),
duration=duration,
resolution=resolution,
generate_audio=bool(generate_audio),
generate_audio=False, # 我们自己混 TTS
watermark=False,
output_dir=output_dir,
model=model,
reference_images=reference_images,
reference_audios=reference_audios,
reference_videos=reference_videos,
)
if effective_ratio:
kwargs["ratio"] = effective_ratio
+4 -19
View File
@@ -368,21 +368,6 @@ fi
echo "All images pulled."
# ====== 打稳定 tag(:dev),供 Watchtower 监控 ======
# Watchtower 只能检测同一个 tag 的 digest 变化。
# commit SHA tag 每次构建都不同,Watchtower 无法感知更新。
# 因此每次部署都将最新镜像 tag 为 :dev,容器统一使用 :dev 启动。
DEV_API="${REGISTRY}/xiaoxia-saas-api:dev"
DEV_WORKER="${REGISTRY}/xiaoxia-saas-worker:dev"
DEV_WEB="${REGISTRY}/xiaoxia-saas-web:dev"
docker tag "$REGISTRY_API" "$DEV_API"
docker tag "$REGISTRY_WORKER" "$DEV_WORKER"
docker tag "$REGISTRY_WEB" "$DEV_WEB"
echo "✅ Tagged images as :dev for Watchtower monitoring"
echo " API: $DEV_API"
echo " Worker: $DEV_WORKER"
echo " Web: $DEV_WEB"
# ====== 镜像内容校验 ======
echo ""
echo "=========================================="
@@ -548,7 +533,7 @@ docker run -d \
--health-retries 3 \
--health-start-period 40s \
$LOG_OPTS \
"$DEV_API" &
"$REGISTRY_API" &
PID_API_START=$!
# ── Worker: 通过 compose 启动(单一事实来源)──
@@ -556,7 +541,7 @@ PID_API_START=$!
# healthcheck 匹配 'celery.*worker'(不把 beat 算活)、资源限制 4C/8G。
# WORKER_IMAGE 通过环境变量覆盖镜像 tag(compose.yml 默认 :dev)。
echo "Starting worker via docker compose (from $INFRA_DOCKER_DIR)..."
WORKER_IMAGE="$DEV_WORKER" APP_VERSION="$IMAGE_TAG" compose up -d --no-deps worker &
WORKER_IMAGE="$REGISTRY_WORKER" APP_VERSION="$IMAGE_TAG" compose up -d --no-deps worker &
PID_WORKER_START=$!
# ── Web: 暂保留 docker run(TODO: 后续收敛到 compose)──
@@ -572,7 +557,7 @@ docker run -d \
--health-timeout 5s \
--health-retries 3 \
$LOG_OPTS \
"$DEV_WEB" &
"$REGISTRY_WEB" &
PID_WEB_START=$!
wait $PID_API_START $PID_WORKER_START $PID_WEB_START
@@ -703,5 +688,5 @@ echo "=== Staging deployment complete ==="
echo "API: http://127.0.0.1:8000"
echo "Web: http://127.0.0.1:3001"
echo "Worker: managed by docker compose (project=$COMPOSE_PROJECT)"
echo "Version: $IMAGE_TAG (running as :dev for Watchtower)"
echo "Version: $IMAGE_TAG"
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging
+1 -1
View File
@@ -57,7 +57,7 @@ if [ "$TARGET_ENV" = "staging" ]; then
fi
# 共用 secrets 直接导出(如果存在)
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_FAST_MODEL DOUBAO_BASE_URL DOUBAO_VISION_MODEL DOUBAO_VISION_LITE_MODEL DOUBAO_VISION_USE_LITE WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY GPU_WORKER_TOKEN"
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_BASE_URL DOUBAO_VISION_MODEL WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY GPU_WORKER_TOKEN"
for var in $SHARED_SECRETS; do
value="${!var:-}"
# 已经在环境中了,无需额外操作
-3
View File
@@ -14,13 +14,10 @@ from packages.shared.ai_client import DoubaoClient
class _FakeSettings:
doubao_api_key = "test-key"
doubao_model = "test-model"
doubao_fast_model = "test-fast-model"
doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
doubao_timeout = 10
doubao_max_retries = 0
doubao_vision_model = "test-vision"
doubao_vision_lite_model = "test-vision-lite"
doubao_vision_use_lite = False
doubao_embedding_model = "test-embedding"
+2 -2
View File
@@ -44,7 +44,7 @@ class TestCheckDatabase:
assert result["type"] == "postgresql"
assert result["message"] == "Database connection successful"
mock_psycopg.connect.assert_called_once_with(
"postgresql://test:test@localhost/test", connect_timeout=3
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_cur.execute.assert_called_once_with("SELECT 1")
mock_conn.close.assert_called_once()
@@ -96,7 +96,7 @@ class TestCheckMigrations:
assert result["status"] == "healthy"
assert result["message"] == "Database migrations applied"
mock_psycopg.connect.assert_called_once_with(
"postgresql://test:test@localhost/test", connect_timeout=3
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_conn.close.assert_called_once()
+51 -104
View File
@@ -120,7 +120,7 @@ class TestViralVideoJobDefaults:
job = ViralVideoJob(user_id="u1")
assert job.images == []
assert job.industry == ""
assert job.duration == 15
assert job.duration == 30
assert job.fusion_level == FusionLevel.AI_POLISH
assert job.style_strength == StyleStrength.MEDIUM
assert job.status == ViralVideoStatus.PENDING
@@ -145,10 +145,13 @@ class TestViralVideoStage:
"image_analysis",
"video_analysis",
"intent_parsing",
"script_generation",
"copy_fusion",
"storyboard",
"review",
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
]
actual_order = [s.value for s in ViralVideoStage]
@@ -168,7 +171,7 @@ class TestViralVideoSchemas:
assert req.images == ["https://example.com/img.jpg"]
assert req.fusion_level == "ai_polish"
assert req.style_strength == "medium"
assert req.duration == 15
assert req.duration == 30
def test_create_request_empty_images_raises(self):
from app.schemas.viral_video import CreateViralVideoRequest
@@ -362,10 +365,9 @@ class TestViralVideoPipeline:
industry="美妆",
target_customer="年轻女性",
marketing_purpose="品牌推广",
duration=15,
duration=30,
user_copy_text="这款产品超好用",
fusion_level="ai_polish",
video_ratio="9:16",
)
@patch("packages.shared.ai_service.call_vision")
@@ -403,110 +405,63 @@ class TestViralVideoPipeline:
assert "intent" in result
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_returns_copy_result(self, mock_llm, mock_job):
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
def test_copy_fusion_ai_polish(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_copy_fusion
mock_llm.return_value = {
"overview": {"theme": "口红推荐", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台,柔和自然光",
"shots": [
{
"time_range": "0-5秒",
"shot_type_angle_movement": "近景平视,缓慢推镜",
"scene_and_dialogue": "女主微笑展示口红:大家好,今天分享一款口红",
"action_details": "手持口红特写",
"audio_bgm": "轻快流行BGM",
"transition": "硬切",
"reference_image_index": 0,
},
{
"time_range": "5-15秒",
"shot_type_angle_movement": "特写,固定镜头",
"scene_and_dialogue": "涂抹口红:颜色特别好看很显白",
"action_details": "嘴唇涂抹特写",
"audio_bgm": "轻快BGM继续",
"transition": "结束",
"reference_image_index": 1,
},
],
"hard_constraints": ["无字幕无水印"],
"negative_prompts": ["字幕", "水印"],
"voiceover_script": "大家好,今天分享一款口红,颜色特别好看很显白。",
}
result = _step_script_generation(
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
)
assert isinstance(result, dict)
assert "voiceover_script" in result
assert "shots" in result
assert isinstance(result["shots"], list)
assert len(result["shots"]) == 2
assert result["overview"]["total_duration"] == 15
# final_copy 必须 = voiceover_script(向后兼容)
assert result.get("final_copy") == result["voiceover_script"]
mock_llm.return_value = "融合后的文案内容"
result = _step_copy_fusion(mock_job, {"intent": "推广"}, {"products": []})
assert isinstance(result, str)
assert len(result) > 0
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_fallback(self, mock_llm, mock_job):
"""LLM 返回异常时使用兜底脚本(不会抛错)。"""
from apps.worker.worker_app.tasks.viral_video import _fallback_script
def test_storyboard_generation(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
result = _fallback_script(mock_job)
assert isinstance(result, dict)
assert result["voiceover_script"]
assert len(result["shots"]) >= 1
mock_llm.return_value = [
{"order": 0, "type": "product_shot", "duration": 10},
{"order": 1, "type": "closing", "duration": 5},
]
result = _step_storyboard(mock_job, "测试文案", {})
assert isinstance(result, list)
assert len(result) == 2
@patch("packages.shared.ai_service.call_llm")
def test_review_pass_v16(self, mock_llm, mock_job):
"""v1.6 _step_review 接收 copy_result dict。"""
def test_review_pass(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_review
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
cr = {"voiceover_script": "大家好", "shots": []}
result = _step_review(mock_job, cr)
result = _step_review(mock_job, "测试文案", [])
assert result["passed"] is True
def test_assemble_seedance_prompt(self, mock_job):
"""编导脚本必须能拼出完整的 Seedance prompt,含总览/场景/逐镜头/约束。"""
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
def test_bgm_select(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
cr = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [
{
"time_range": "0-15秒",
"shot_type_angle_movement": "中景平视",
"scene_and_dialogue": "你好分享",
"action_details": "展示",
"audio_bgm": "BGM",
"transition": "结束",
"reference_image_index": 0,
}
],
"hard_constraints": ["无字幕"],
"negative_prompts": ["水印"],
}
prompt = _assemble_seedance_prompt(cr, mock_job)
assert "【视频总览】" in prompt
assert "【逐镜头时间轴】" in prompt
assert "【硬性约束】" in prompt
assert "【负面提示词】" in prompt
assert "0-15秒" in prompt
# P1: BGM 素材未就绪前 _step_bgm_select 统一返回 None(跳过 BGM 混音)
mock_job.bgm_preference = "upbeat"
bgm = _step_bgm_select(mock_job)
assert bgm is None
def test_bgm_select_default(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
mock_job.bgm_preference = ""
bgm = _step_bgm_select(mock_job)
assert bgm is None
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
class TestPipelineIntegration:
"""v1.6 流水线端到端集成测试(mock 外部依赖):TTS+单次 Seedance+上传。"""
"""流水线端到端集成测试(mock 外部依赖)。"""
@patch("apps.worker.worker_app.tasks.viral_video._step_upload")
@patch("apps.worker.worker_app.tasks.viral_video._step_render")
@patch("apps.worker.worker_app.tasks.viral_video._upload_tts_to_oss")
@patch("apps.worker.worker_app.tasks.viral_video._step_bgm_select")
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
@patch("apps.worker.worker_app.tasks.viral_video._step_review")
@patch("apps.worker.worker_app.tasks.viral_video._step_script_generation")
@patch("apps.worker.worker_app.tasks.viral_video._step_storyboard")
@patch("apps.worker.worker_app.tasks.viral_video._step_copy_fusion")
@patch("apps.worker.worker_app.tasks.viral_video._step_intent_parsing")
@patch("apps.worker.worker_app.tasks.viral_video._step_video_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
@@ -519,46 +474,38 @@ class TestPipelineIntegration:
mock_img_analysis,
mock_video_analysis,
mock_intent,
mock_script,
mock_copy_fusion,
mock_storyboard,
mock_review,
mock_tts,
mock_tts_upload,
mock_bgm,
mock_render,
mock_upload,
):
"""v1.6: TTS整段合成 → 上传TTS到OSS → 单次 Seedance → 上传成片。"""
"""测试 resume 流水线能从确认状态走到完成。"""
from apps.worker.worker_app.tasks.viral_video import (
resume_viral_video_pipeline,
)
# 构造 mock job
job = ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
status=ViralVideoStatus.RUNNING,
intent_result={"intent": "推广"},
duration=15,
video_ratio="9:16",
)
mock_repo = MagicMock()
mock_session = MagicMock()
mock_get_repo.return_value = (mock_session, mock_repo, job)
# v1.6: 如果没有 copy_result 会现场补生成
mock_intent.return_value = {"intent": "推广", "key_messages": [], "tone": "亲切"}
mock_script.return_value = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [],
"hard_constraints": [],
"negative_prompts": [],
"voiceover_script": "大家好,分享一款口红。",
"final_copy": "大家好,分享一款口红。",
}
# 设置各步骤返回值
mock_copy_fusion.return_value = "融合文案"
mock_storyboard.return_value = [{"order": 0, "duration": 10}]
mock_review.return_value = {"passed": True, "score": 90}
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
mock_tts_upload.return_value = None
mock_tts.return_value = None # P1: TTS 返回 Path|None,mock 用 None 跳过混音
mock_bgm.return_value = None # P1: BGM 未就绪前返回 None
mock_render.return_value = "/tmp/video.mp4"
mock_upload.return_value = "https://oss.example.com/final.mp4"
+52 -128
View File
@@ -67,70 +67,49 @@ class TestImageAnalysisField:
# ── P0-1: storyboard 规范化 ────────────────────────────────────────
class TestScriptGenerationV16:
"""v1.6 编导分镜脚本生成相关纯函数测试。"""
class TestStoryboardNormalize:
def test_normalize_fills_defaults(self):
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
def test_fallback_script_has_required_fields(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _fallback_script
raw = [{"order": 0, "description": "镜头一"}]
out = _normalize_storyboard(raw, total_duration=10, n_segments=1, copy_text="文案")
assert len(out) == 1
assert out[0]["duration"] >= 3
assert out[0]["ken_burns"] in {"zoom_in", "zoom_out", "pan_left", "pan_right", "static"}
assert out[0]["type"] == "product_shot"
assert out[0]["text"] == ""
out = _fallback_script(mock_job)
assert isinstance(out, dict)
assert "overview" in out
assert "shots" in out
assert "voiceover_script" in out
assert "hard_constraints" in out
assert "negative_prompts" in out
assert out["overview"]["total_duration"] == mock_job.duration
assert out["final_copy"] == out["voiceover_script"]
assert len(out["shots"]) >= 1
def test_normalize_scales_to_total_duration(self):
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
def test_safe_json_loads_parses_fenced_code(self):
from apps.worker.worker_app.tasks.viral_video import _safe_json_loads
raw = [
{"order": 0, "duration": 10, "description": "a"},
{"order": 1, "duration": 10, "description": "b"},
]
out = _normalize_storyboard(raw, total_duration=10, n_segments=2, copy_text="x")
total = sum(s["duration"] for s in out)
assert total == 10
fenced = '```json\n{"voiceover_script": "你好", "shots": []}\n```'
out = _safe_json_loads(fenced)
assert out is not None
assert out["voiceover_script"] == "你好"
def test_fallback_storyboard(self):
from apps.worker.worker_app.tasks.viral_video import _fallback_storyboard
def test_safe_json_loads_handles_none(self):
from apps.worker.worker_app.tasks.viral_video import _safe_json_loads
out = _fallback_storyboard("文案", total_duration=15, n_segments=3)
assert len(out) == 3
assert sum(s["duration"] for s in out) == 15
assert all(s["duration"] >= 3 for s in out)
assert _safe_json_loads(None) is None
assert _safe_json_loads("not json") is None
def test_storyboard_llm_list(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
def test_validate_normalize_fills_defaults(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _validate_and_normalize_script
raw = {"voiceover_script": "你好", "shots": [{"scene_and_dialogue": "测试"}]}
out = _validate_and_normalize_script(raw, mock_job)
assert out["voiceover_script"] == "你好"
assert len(out["shots"]) == 1
assert out["shots"][0]["shot_type_angle_movement"]
assert out["overview"]["total_duration"] == mock_job.duration
def test_assemble_seedance_prompt_contains_sections(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
cr = {
"overview": {"theme": "测试", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮",
"shots": [
{
"time_range": "0-15秒",
"shot_type_angle_movement": "中景",
"scene_and_dialogue": "你好",
"action_details": "展示",
"audio_bgm": "BGM",
"transition": "结束",
"reference_image_index": 0,
}
],
"hard_constraints": ["无字幕"],
"negative_prompts": ["水印"],
}
p = _assemble_seedance_prompt(cr, mock_job)
for key in ("【视频总览】", "【场景与光线】", "【逐镜头时间轴】", "【硬性约束】", "【负面提示词】"):
assert key in p
with patch("packages.shared.ai_service.call_llm") as mock_llm:
mock_llm.return_value = [
{"order": 0, "description": "产品特写", "duration": 5, "text": "t1"},
{"order": 1, "description": "使用场景", "duration": 5, "text": "t2"},
{"order": 2, "description": "CTA", "duration": 5, "text": "t3"},
]
result = _step_storyboard(mock_job, "文案", {"products": []})
assert len(result) == 3
assert all("description" in s for s in result)
# ── P1: TTS 返回 Path|None ────────────────────────────────────────
@@ -171,24 +150,12 @@ class TestTTSPath:
# ── P1: BGM 跳过 / MuseTalk 无 persona 跳过 ───────────────────────
class TestDurationClamp:
"""v1.6 mark_copy_generated 派生字段 + duration clamp。"""
class TestBGMSkip:
def test_bgm_returns_none(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
def test_mark_copy_generated_derives_fields(self):
job = ViralVideoJob(user_id="u1", duration=15)
cr = {
"overview": {"theme": "x", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "亮",
"shots": [{"time_range": "0-15秒", "scene_and_dialogue": "对白"}],
"voiceover_script": "你好",
"hard_constraints": [],
"negative_prompts": [],
}
job.mark_copy_generated(cr)
assert job.copy_result is cr
assert job.generated_copy_text == "你好"
assert job.storyboard == cr["shots"]
assert job.effective_copy_text == "你好"
mock_job.bgm_preference = "upbeat"
assert _step_bgm_select(mock_job) is None
# ── P0-1: call_video_generation 参数构造 ──────────────────────────
@@ -221,61 +188,23 @@ class TestCallVideoGeneration:
assert kwargs["prompt"] == "测试"
assert kwargs["image_url"] == "https://img/x.jpg"
assert kwargs["duration"] == 5
assert kwargs["generate_audio"] is True
# ── P0-1: _step_render 占位片段生成 ──────────────────────────────
class TestCallVideoGenerationV16:
"""v1.6 call_video_generation 透传 reference_audios/reference_images 等参数到 client。"""
class TestPlaceholderClip:
def test_make_placeholder_clip(self, tmp_path):
import shutil
def test_passes_reference_params_to_client(self, tmp_path):
from packages.shared.ai_service import call_video_generation
from apps.worker.worker_app.tasks.viral_video import _make_placeholder_clip, _probe_ok
out = tmp_path / "v.mp4"
out.write_bytes(b"fake")
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
mock_client = MagicMock()
mock_client.is_available = True
mock_client.video_generation.return_value = str(out)
mock_get.return_value = mock_client
result = call_video_generation(
prompt="测试",
image_url="https://img/x.jpg",
duration=15,
ratio="9:16",
reference_images=["https://img/r1.jpg"],
reference_audios=["https://oss/tts.mp3"],
reference_videos=["https://oss/ref.mp4"],
generate_audio=True,
model="doubao-seedance-2-5-260628",
)
assert result == str(out)
kwargs = mock_client.video_generation.call_args.kwargs
# 图生视频也必须传 ratio(避免首帧方图导致默认输出 1:1)
assert kwargs.get("ratio") == "9:16", f"ratio 应透传,got {kwargs.get('ratio')!r}"
assert kwargs["image_url"] == "https://img/x.jpg"
assert kwargs["reference_audios"] == ["https://oss/tts.mp3"]
assert kwargs["reference_images"] == ["https://img/r1.jpg"]
assert kwargs["reference_videos"] == ["https://oss/ref.mp4"]
assert kwargs["generate_audio"] is True
assert kwargs["model"] == "doubao-seedance-2-5-260628"
if not shutil.which("ffmpeg"):
pytest.skip("ffmpeg not available")
def test_ratio_passed_when_no_image(self, tmp_path):
from packages.shared.ai_service import call_video_generation
out = tmp_path / "v.mp4"
out.write_bytes(b"fake")
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
mock_client = MagicMock()
mock_client.is_available = True
mock_client.video_generation.return_value = str(out)
mock_get.return_value = mock_client
call_video_generation(prompt="测试", duration=10, ratio="16:9")
kwargs = mock_client.video_generation.call_args.kwargs
assert kwargs["ratio"] == "16:9"
assert kwargs["image_url"] is None
out = _make_placeholder_clip(tmp_path, 0, 3)
assert out.exists()
assert _probe_ok(str(out))
# ── P0-1: DoubaoClient.video_generation 在不可用时返回 None ───────
@@ -295,15 +224,10 @@ class TestDoubaoClientVideoGen:
class TestResumeReadsImageAnalysis:
def test_resume_uses_persisted_image_analysis(self):
"""resume/render pipeline 应从 job.image_analysis 读(v1.5 _run_render_pipeline 共享渲染逻辑)。"""
"""resume_pipeline 应从 job.image_analysis 读(P0-3 持久化)。"""
import inspect
from apps.worker.worker_app.tasks import viral_video as vv
# v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis
src = inspect.getsource(vv._run_render_pipeline)
src = inspect.getsource(vv.resume_viral_video_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
+1 -226
View File
@@ -34,7 +34,7 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
"viral_structure": "",
"marketing_purpose": "",
"bgm_preference": "",
"duration": 15,
"duration": 30,
"user_copy_text": "",
"fusion_level": "ai_polish",
"reference_audio_path": "",
@@ -51,22 +51,7 @@ 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,
"copy_result": None,
"generated_copy_text": "",
"voice_id": "",
"voice_source": "",
"voice_mode": "global",
"video_ratio": "9:16",
"video_model": "",
"credits_cost": 0,
"current_stage": "",
"phase_message": "",
"updated_at": None,
"is_terminal": False,
"effective_copy_text": "",
"voiceover_script": "",
}.items():
setattr(job, k, kwargs.pop(k, v))
return job
@@ -189,213 +174,3 @@ 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
def test_generate_copy_persists_voice_and_ratio(self):
"""generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。"""
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-gc3", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED)
repo = MagicMock()
repo.get.return_value = job
req = GenerateCopyRequest(
voice_id="cosy_voice_001",
voice_source="library",
video_ratio="16:9",
)
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task"),
):
vv_mod.generate_copy("job-gc3", req, authenticated_user=user, session=session)
assert job.voice_id == "cosy_voice_001"
assert job.voice_source == "library"
assert job.video_ratio == "16:9"
class TestAnalyzeImagesPersist:
def test_analyze_images_persists_voice_and_ratio(self):
"""analyze-images 创建任务时应带上 voice/video_ratio 字段。"""
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()
saved = {}
class FakeRepo:
def save(self, job):
saved["job"] = job
def get(self, jid):
return None
req = AnalyzeImagesRequest(
images=["img-1"],
voice_id="preset_v1",
voice_source="preset",
video_ratio="1:1",
)
with (
patch.object(vv_mod, "_get_job_repo", return_value=FakeRepo()),
patch.object(vv_mod.celery_app, "send_task"),
):
resp = vv_mod.analyze_images(req, authenticated_user=user, session=session)
job = saved["job"]
assert job.voice_id == "preset_v1"
assert job.voice_source == "preset"
assert job.video_ratio == "1:1"
assert resp.images == ["img-1"]
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