fix(viral-video): #2106 P0 渲染阻塞修复 + Seedance 2.5 对接 + image_analysis 持久化
P0-1: Seedance 2.5 视频生成对接 - packages/shared/ai_client.py: DoubaoClient 新增 video_generation(prompt, image_url, duration, ratio, ...) 方法:走方舟 /contents/generations/tasks 异步任务(submit→poll→download),返回本地 MP4 路径 - packages/shared/ai_service.py: 新增 call_video_generation() 高层封装 - packages/config/base.py: 新增 doubao_video_model/doubao_video_timeout/doubao_video_poll_interval 配置 - apps/worker/.../viral_video.py _step_render 重写:按 storyboard 分镜逐段调 Seedance 生成短视频 → ffmpeg concat 拼接 → 混入 TTS 音频(-map 0:v/1:a -shortest -c:v copy) - 单分镜失败自动用 ffmpeg color 源占位片段兜底,保证 concat 不中断 - 新增 _normalize_storyboard/_fallback_storyboard/_build_segment_prompt/_probe_ok/_make_placeholder_clip 等辅助函数 P0-2: _step_video_analysis import 路径修复 - from worker_app.tasks.viral_video_analyzer → from viral_video.video_analyzer import analyze_video_style (文件在 apps/worker/viral_video/video_analyzer.py,worker PYTHONPATH 包含 apps/worker) - ImportError 仍兜底返回占位 style_guide,不阻塞流水线 P0-3: image_analysis 持久化 - packages/domain/viral_video.py: ViralVideoJob 新增 image_analysis: dict|None 字段 - packages/adapters/sqlalchemy_impl/models.py: viral_video_jobs 加 image_analysis JSON 列 - packages/adapters/sqlalchemy_impl/viral_video_repository.py: _to_domain/save/update 同步该字段 - alembic/versions/087_viral_video_image_analysis.py: migration 087 - run_viral_video_pipeline: step1 后立即 job.image_analysis = image_analysis 并 _save_job - resume_viral_video_pipeline: 从 job.image_analysis 读取,不再硬编码空 dict P1 顺手修复: - _step_tts: 返回值统一为 Path|None,Path 不存在/ImportError/synthesize 失败均返回 None - _step_bgm_select: BGM 素材未就绪前统一返回 None,渲染时跳过 BGM 混音 - _step_musetalk: GPU 端点未就绪前(即使有 persona_id)也直接跳过,不发 HTTP 请求 - TTS 工厂 import 路径修正: worker_app.services.tts_service_factory → services.tts_service_factory - credits_cost 赋值保留 TODO(等 credits.deduct() 总开关) - tests/unit/test_viral_video.py: BGM/TTS mock 对齐新返回语义 - tests/unit/test_viral_video_p0.py: 新增 15 个单测覆盖 video_analysis/storyboard 规范化/ TTS Path 处理/BGM None/MuseTalk 跳过/call_video_generation 委托/placeholder clip/resume 读 job
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
"""viral video add image_analysis column
|
||||
|
||||
Revision ID: 087_viral_video_image_analysis
|
||||
Revises: 086_add_viral_video_tables
|
||||
Create Date: 2026-09-30
|
||||
|
||||
#2106 爆款视频 P0:持久化图片分析结果(image_analysis JSON),供 resume 阶段使用。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "087_viral_video_image_analysis"
|
||||
down_revision = "086_add_viral_video_tables"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("viral_video_jobs", sa.Column("image_analysis", sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("viral_video_jobs", "image_analysis")
|
||||
@@ -97,6 +97,9 @@ def _save_job(repo, job, session):
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _step_image_analysis(job: ViralVideoJob) -> dict:
|
||||
"""步骤 1: 图片 VLM 分析 — 识别产品特征、场景、卖点。"""
|
||||
try:
|
||||
@@ -126,13 +129,13 @@ def _step_video_analysis(job: ViralVideoJob) -> dict | None:
|
||||
return None
|
||||
|
||||
try:
|
||||
# 尝试导入 video_analyzer(由 #2051 提供)
|
||||
from worker_app.tasks.viral_video_analyzer import analyze_video_style
|
||||
# P0-2: 修正 import 路径(video_analyzer.py 在 apps/worker/viral_video/ 下,worker PYTHONPATH 含 apps/worker)
|
||||
from viral_video.video_analyzer import analyze_video_style
|
||||
|
||||
style_guide = analyze_video_style(job.reference_video_url)
|
||||
return style_guide
|
||||
except ImportError:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪,使用占位风格分析")
|
||||
except ImportError as e:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪(%s),使用占位风格分析", e)
|
||||
return {
|
||||
"cut_speed": "medium",
|
||||
"transition": "cross_dissolve",
|
||||
@@ -229,42 +232,120 @@ def _step_copy_fusion(job: ViralVideoJob, intent: dict, image_analysis: dict) ->
|
||||
|
||||
|
||||
def _step_storyboard(job: ViralVideoJob, copy_text: str, image_analysis: dict) -> list[dict]:
|
||||
"""步骤 4: 分镜脚本生成。"""
|
||||
"""步骤 4: 分镜脚本生成。每个分镜独立一段视频,段内时长建议 3~6 秒。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return [{"order": 0, "type": "product_shot", "text": copy_text[:50], "duration": job.duration}]
|
||||
return [
|
||||
{
|
||||
"order": 0,
|
||||
"type": "product_shot",
|
||||
"text": copy_text[:50],
|
||||
"duration": min(5, job.duration),
|
||||
"description": "产品展示",
|
||||
"ken_burns": "zoom_in",
|
||||
"transition": "cut",
|
||||
}
|
||||
]
|
||||
|
||||
prompt = f"""请根据以下文案生成短视频分镜脚本:
|
||||
products_hint = ""
|
||||
products = image_analysis.get("products", []) if image_analysis else []
|
||||
if products:
|
||||
p0 = products[0] if isinstance(products[0], dict) else {}
|
||||
feats = p0.get("features", []) if isinstance(p0, dict) else []
|
||||
products_hint = f"\n首帧参考产品特征:{p0.get('name','')} - {', '.join(feats[:3])}"
|
||||
|
||||
seg_seconds = 5
|
||||
n_segments = max(2, min(6, max(1, job.duration // seg_seconds)))
|
||||
ratio = "9:16"
|
||||
|
||||
prompt = f"""请根据以下文案生成爆款短视频分镜脚本,共 {n_segments} 个分镜:
|
||||
|
||||
文案内容:{copy_text}
|
||||
视频时长:{job.duration}秒
|
||||
视频总时长:{job.duration}秒(每个分镜 3~6 秒,总和约等于总时长)
|
||||
风格强度:{job.style_strength}
|
||||
输出宽高比:{ratio}{products_hint}
|
||||
|
||||
请以JSON数组格式返回分镜列表,每个分镜包含:
|
||||
- order: 序号
|
||||
- type: 镜头类型(product_shot/text_card/scene_transition/closing)
|
||||
- description: 画面描述
|
||||
- text: 配音/字幕文本
|
||||
- duration: 时长(秒)
|
||||
- ken_burns: 运镜方式(zoom_in/zoom_out/pan_left/pan_right/none)
|
||||
- transition: 转场方式(cut/dissolve/wipe/fade)"""
|
||||
请以 JSON 数组格式返回分镜列表,每个分镜包含:
|
||||
- order: 序号(从0开始)
|
||||
- type: 镜头类型(product_shot/close_up/scene/action/text_card/closing)
|
||||
- description: 画面详细描述(中文,含主体、动作、场景、运镜、光影,用于AI视频生成prompt)
|
||||
- text: 该分镜配音/字幕文本
|
||||
- duration: 时长(秒,3~6秒的整数)
|
||||
- ken_burns: 运镜方式(zoom_in/zoom_out/pan_left/pan_right/static)
|
||||
- transition: 与下一分镜的转场(cut/dissolve/fade)"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
if isinstance(result, list):
|
||||
return result
|
||||
# 尝试从字符串中解析 JSON
|
||||
import json
|
||||
return _normalize_storyboard(result, job.duration, n_segments, copy_text)
|
||||
import json as _json
|
||||
|
||||
return (
|
||||
json.loads(result)
|
||||
if isinstance(result, str)
|
||||
else [{"order": 0, "text": copy_text, "duration": job.duration}]
|
||||
)
|
||||
parsed = _json.loads(result) if isinstance(result, str) else result
|
||||
if isinstance(parsed, list):
|
||||
return _normalize_storyboard(parsed, job.duration, n_segments, copy_text)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 分镜生成失败: %s", e)
|
||||
return [{"order": 0, "type": "product_shot", "text": copy_text[:100], "duration": job.duration}]
|
||||
|
||||
return _fallback_storyboard(copy_text, job.duration, n_segments)
|
||||
|
||||
|
||||
def _normalize_storyboard(raw: list, total_duration: int, n_segments: int, copy_text: str) -> list[dict]:
|
||||
"""规范化 LLM 输出的分镜:填充缺省字段、保证总时长合理。"""
|
||||
out: list[dict] = []
|
||||
for i, item in enumerate(raw):
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
dur = int(item.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(3, min(8, dur))
|
||||
out.append(
|
||||
{
|
||||
"order": int(item.get("order", i)),
|
||||
"type": str(item.get("type", "product_shot")),
|
||||
"description": str(item.get("description", copy_text[:80])),
|
||||
"text": str(item.get("text", "")),
|
||||
"duration": dur,
|
||||
"ken_burns": str(item.get("ken_burns", "zoom_in")),
|
||||
"transition": str(item.get("transition", "cut")),
|
||||
}
|
||||
)
|
||||
if not out:
|
||||
return _fallback_storyboard(copy_text, total_duration, n_segments)
|
||||
out = out[:n_segments]
|
||||
total = sum(s["duration"] for s in out)
|
||||
if total > 0 and total != total_duration:
|
||||
scale = total_duration / total
|
||||
acc = 0
|
||||
for s in out[:-1]:
|
||||
s["duration"] = max(3, min(8, round(s["duration"] * scale)))
|
||||
acc += s["duration"]
|
||||
out[-1]["duration"] = max(3, total_duration - acc)
|
||||
return out
|
||||
|
||||
|
||||
def _fallback_storyboard(copy_text: str, total_duration: int, n_segments: int) -> list[dict]:
|
||||
if n_segments <= 0:
|
||||
n_segments = 1
|
||||
dur = total_duration // n_segments
|
||||
remainder = total_duration - dur * n_segments
|
||||
out = []
|
||||
for i in range(n_segments):
|
||||
d = dur + (remainder if i == n_segments - 1 else 0)
|
||||
out.append(
|
||||
{
|
||||
"order": i,
|
||||
"type": "product_shot",
|
||||
"description": f"产品展示镜头 {i + 1}:{copy_text[:40]}",
|
||||
"text": copy_text,
|
||||
"duration": max(3, d),
|
||||
"ken_burns": "zoom_in" if i % 2 == 0 else "pan_left",
|
||||
"transition": "cut",
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _step_review(job: ViralVideoJob, copy_text: str, storyboard: list[dict]) -> dict:
|
||||
@@ -296,92 +377,230 @@ def _step_review(job: ViralVideoJob, copy_text: str, storyboard: list[dict]) ->
|
||||
return {"passed": True, "score": 75, "details": {d: "默认通过" for d in dimensions}}
|
||||
|
||||
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str) -> str:
|
||||
"""步骤 6: CosyVoice 配音。"""
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str):
|
||||
"""步骤 6: CosyVoice 配音。P1:返回 Path;失败返回 None。"""
|
||||
try:
|
||||
from worker_app.services.tts_service_factory import get_tts_service
|
||||
from pathlib import Path as _Path
|
||||
from services.tts_service_factory import get_tts_service
|
||||
|
||||
tts_service = get_tts_service()
|
||||
# 简化调用,实际需要更详细的参数
|
||||
audio_url = tts_service.synthesize(text=copy_text, voice_id=job.persona_id or "default")
|
||||
return audio_url
|
||||
result = tts_service.synthesize(text=copy_text, voice_id=job.persona_id or "default")
|
||||
if result is None:
|
||||
return None
|
||||
p = _Path(result) if not isinstance(result, _Path) else result
|
||||
if p.exists():
|
||||
return p
|
||||
logger.warning("[爆款视频] TTS 返回路径不存在: %s", p)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 配音失败: %s", e)
|
||||
return ""
|
||||
return None
|
||||
|
||||
|
||||
def _step_bgm_select(job: ViralVideoJob) -> str:
|
||||
"""步骤 7: BGM 选择。"""
|
||||
# 基于 bgm_preference 和 marketing_purpose 匹配预设 BGM
|
||||
bgm_map = {
|
||||
"upbeat": "bgm_upbeat_01.mp3",
|
||||
"calm": "bgm_calm_01.mp3",
|
||||
"energetic": "bgm_energetic_01.mp3",
|
||||
"emotional": "bgm_emotional_01.mp3",
|
||||
def _step_bgm_select(job: ViralVideoJob):
|
||||
"""步骤 7: BGM 选择。P1:素材未就绪前返回 None,跳过 BGM 混音。"""
|
||||
return None
|
||||
|
||||
|
||||
def _build_segment_prompt(seg: dict, job: ViralVideoJob, style_hint: str) -> str:
|
||||
desc = seg.get("description") or seg.get("text") or "产品展示"
|
||||
ken_burns = seg.get("ken_burns", "zoom_in")
|
||||
cam_map = {
|
||||
"zoom_in": "缓慢推镜放大",
|
||||
"zoom_out": "缓慢拉镜缩小",
|
||||
"pan_left": "镜头向左平移",
|
||||
"pan_right": "镜头向右平移",
|
||||
"static": "固定镜头",
|
||||
}
|
||||
preference = job.bgm_preference.lower()
|
||||
for key, bgm in bgm_map.items():
|
||||
if key in preference:
|
||||
return bgm
|
||||
return "bgm_default.mp3"
|
||||
camera = cam_map.get(ken_burns, "缓慢运镜")
|
||||
parts = [
|
||||
f"{desc}。",
|
||||
f"运镜:{camera}。",
|
||||
"画面流畅、电影感光影、高清细节,9:16竖屏,适合短视频。",
|
||||
]
|
||||
if style_hint:
|
||||
parts.append(f"参考风格:{style_hint}")
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _step_render(job: ViralVideoJob, storyboard: list[dict], audio_url: str, bgm: str) -> str:
|
||||
"""步骤 8: UnifiedRenderService 渲染。"""
|
||||
try:
|
||||
from video_processing.render_adapter import build_render_plan
|
||||
from video_processing.unified_render_service import UnifiedRenderService
|
||||
def _step_render(job, storyboard, tts_path, bgm):
|
||||
"""步骤 8: 渲染(P0-1 核心重写)。
|
||||
|
||||
render_plan = build_render_plan(
|
||||
images=job.images,
|
||||
storyboard=storyboard,
|
||||
audio_url=audio_url,
|
||||
bgm=bgm,
|
||||
duration=job.duration,
|
||||
style_guide=job.style_guide,
|
||||
每个 storyboard 分镜 → Seedance 2.5 生成短视频段(无声)→ 下载 → ffmpeg concat → 混入 TTS。
|
||||
返回最终视频本地路径字符串。
|
||||
"""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
from video_processing.concat_engine import concat_video_files
|
||||
from packages.shared.ffmpeg_utils import run_ffmpeg
|
||||
|
||||
if not storyboard:
|
||||
raise ValueError("storyboard is empty")
|
||||
|
||||
style_hint = ""
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、色调{job.style_guide.get('color_grade','')}"
|
||||
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
|
||||
logger.info("[爆款视频] 开始渲染,分镜数=%d, tmpdir=%s", len(storyboard), tmpdir)
|
||||
|
||||
seg_paths: list[str] = []
|
||||
first_image = job.images[0] if job.images else None
|
||||
n_total = len(storyboard)
|
||||
for i, seg in enumerate(storyboard):
|
||||
try:
|
||||
dur = int(seg.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(2, min(12, dur))
|
||||
prompt = _build_segment_prompt(seg, job, style_hint)
|
||||
_emit_progress(
|
||||
job.id,
|
||||
ViralVideoStage.RENDERING,
|
||||
80.0 + (i + 1) / max(n_total, 1) * 5.0,
|
||||
f"正在生成分镜 {i + 1}/{n_total} ({dur}s)...",
|
||||
)
|
||||
logger.info("[爆款视频] 分镜 %d/%d dur=%ds prompt=%s", i + 1, n_total, dur, prompt[:80])
|
||||
seg_path = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image if i == 0 else None,
|
||||
duration=dur,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmpdir),
|
||||
)
|
||||
if not seg_path or not Path(seg_path).exists():
|
||||
logger.warning("[爆款视频] 分镜 %d 生成失败,使用占位片段", i + 1)
|
||||
seg_path = str(_make_placeholder_clip(tmpdir, i, dur))
|
||||
seg_paths.append(seg_path)
|
||||
|
||||
render_svc = UnifiedRenderService()
|
||||
output_path = render_svc.render(render_plan)
|
||||
return output_path
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 86.0, "正在拼接分镜...")
|
||||
concat_out = tmpdir / "concat_raw.mp4"
|
||||
try:
|
||||
concat_video_files(seg_paths, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 渲染失败: %s", e, exc_info=True)
|
||||
raise
|
||||
logger.error("[爆款视频] concat 失败: %s,降级过滤无效片段", e, exc_info=True)
|
||||
valid = [p for p in seg_paths if _probe_ok(p)]
|
||||
if not valid:
|
||||
raise RuntimeError(f"所有分镜片段均无效: {e}") from e
|
||||
concat_video_files(valid, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
|
||||
final_path = concat_out
|
||||
|
||||
if tts_path is not None:
|
||||
tts_p = Path(tts_path) if not isinstance(tts_path, Path) else tts_path
|
||||
if tts_p.exists():
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 87.5, "正在合成配音...")
|
||||
mixed_out = tmpdir / "final_with_audio.mp4"
|
||||
try:
|
||||
run_ffmpeg(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(concat_out),
|
||||
"-i",
|
||||
str(tts_p),
|
||||
"-c:v",
|
||||
"copy",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"192k",
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-shortest",
|
||||
str(mixed_out),
|
||||
]
|
||||
)
|
||||
if mixed_out.exists() and mixed_out.stat().st_size > 0:
|
||||
final_path = mixed_out
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 混音失败,使用无声视频: %s", e)
|
||||
|
||||
logger.info("[爆款视频] 渲染完成: %s size=%d", final_path, final_path.stat().st_size if final_path.exists() else 0)
|
||||
return str(final_path)
|
||||
|
||||
|
||||
def _probe_ok(video_path: str) -> bool:
|
||||
from pathlib import Path as _Path
|
||||
import subprocess
|
||||
|
||||
try:
|
||||
if not _Path(video_path).exists():
|
||||
return False
|
||||
r = subprocess.run(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=codec_type",
|
||||
"-of",
|
||||
"csv=p=0",
|
||||
video_path,
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=10,
|
||||
)
|
||||
return r.returncode == 0 and b"video" in r.stdout
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _make_placeholder_clip(tmpdir, idx: int, duration: int):
|
||||
import subprocess
|
||||
|
||||
out = tmpdir / f"placeholder_{idx}.mp4"
|
||||
try:
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"color=c=0x202030:s=720x1280:d={max(duration,2)}:r=24",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"anullsrc=r=44100:cl=stereo:d={max(duration,2)}",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-preset",
|
||||
"ultrafast",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(out),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=60,
|
||||
check=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 占位片段生成失败: %s", e)
|
||||
return out
|
||||
|
||||
|
||||
def _step_musetalk(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 9: 数字人口型(MuseTalk)。"""
|
||||
# MuseTalk 集成由现有 GPU worker 处理
|
||||
# 这里调用现有接口
|
||||
try:
|
||||
# 如果不需要数字人,直接跳过
|
||||
if not job.persona_id:
|
||||
return video_path
|
||||
|
||||
# 调用 GPU worker 的 MuseTalk 接口
|
||||
import requests
|
||||
|
||||
gpu_worker_url = os.environ.get("GPU_WORKER_URL", "http://localhost:8900")
|
||||
resp = requests.post(
|
||||
f"{gpu_worker_url}/api/v1/gpu/lipsync",
|
||||
json={
|
||||
"video_path": video_path,
|
||||
"audio_path": job.reference_audio_path,
|
||||
"persona_id": job.persona_id,
|
||||
},
|
||||
timeout=300,
|
||||
)
|
||||
if resp.ok:
|
||||
result = resp.json()
|
||||
return result.get("output_path", video_path)
|
||||
return video_path
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] MuseTalk 处理失败,使用原始视频: %s", e)
|
||||
"""步骤 9: 数字人口型。P1:MuseTalk GPU 端点未就绪前统一跳过。"""
|
||||
if not job.persona_id:
|
||||
return video_path
|
||||
logger.info("[爆款视频] persona_id=%s 已设置,但 MuseTalk GPU 端点未就绪,跳过口型同步", job.persona_id)
|
||||
return video_path
|
||||
|
||||
|
||||
def _step_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 10: OSS 上传 + 扣点。"""
|
||||
"""步骤 10: OSS 上传。"""
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
@@ -397,7 +616,7 @@ def _step_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
|
||||
@shared_task(bind=True, max_retries=2, name="worker.run_viral_video_pipeline")
|
||||
def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""爆款视频 10 步流水线编排器。"""
|
||||
"""爆款视频 10 步流水线编排器(前半段:图片分析→风格分析→意图解析,然后 WAIT_USER_CONFIRM)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
@@ -405,7 +624,6 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
logger.error("[爆款视频] 任务不存在: %s", job_id)
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
# 标记运行中
|
||||
job.mark_running()
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 5.0, "开始图片分析")
|
||||
@@ -413,9 +631,13 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
# ── Step 1: 图片 VLM 分析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 10.0, "正在分析产品图片...")
|
||||
image_analysis = _step_image_analysis(job)
|
||||
# P0-3: 持久化 image_analysis 到 job,供 resume 阶段使用
|
||||
job.image_analysis = image_analysis
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 15.0, "图片分析完成", {"result": image_analysis})
|
||||
|
||||
# ── Step 1.5: 视频风格分析(v1.3) ──
|
||||
style_guide = None
|
||||
if job.reference_video_url or job.style_template_id:
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 20.0, "正在分析参考视频风格...")
|
||||
style_guide = _step_video_analysis(job)
|
||||
@@ -428,14 +650,11 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"风格分析完成",
|
||||
{"style_analyzed": True, "style_guide": style_guide},
|
||||
)
|
||||
else:
|
||||
style_guide = None
|
||||
|
||||
# ── Step 2: 意图解析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 30.0, "正在解析文案意图...")
|
||||
intent_result = _step_intent_parsing(job, image_analysis)
|
||||
|
||||
# 进入等待用户确认状态
|
||||
job.mark_wait_user_confirm(intent_result)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(
|
||||
@@ -454,8 +673,6 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
|
||||
# 这里流水线暂停,等待 confirm-intent API 调用 resume
|
||||
# resume 后由 resume_viral_video_pipeline 继续
|
||||
return {"ok": True, "job_id": job_id, "status": "wait_user_confirm", "intent_result": intent_result}
|
||||
|
||||
except Retry:
|
||||
@@ -504,43 +721,44 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
if job.status != ViralVideoStatus.RUNNING:
|
||||
return {"ok": False, "error": f"unexpected status: {job.status}"}
|
||||
|
||||
# P0-3: 从 job 读取 image_analysis(run_pipeline 阶段已持久化)
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...")
|
||||
|
||||
# ── Step 3: 文案融合 ──
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, {"products": []})
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成")
|
||||
|
||||
# ── Step 4: 分镜脚本 ──
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...")
|
||||
storyboard = _step_storyboard(job, copy_text, {})
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成")
|
||||
storyboard = _step_storyboard(job, copy_text, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)})
|
||||
|
||||
# ── Step 5: 合规审核 ──
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...")
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
if not review_result.get("passed", True):
|
||||
# 自动重写 1 次
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, {"products": []})
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
|
||||
|
||||
# ── Step 6: CosyVoice 配音 ──
|
||||
# ── Step 6: CosyVoice 配音(返回 Path | None) ──
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...")
|
||||
audio_url = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成")
|
||||
tts_path = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None})
|
||||
|
||||
# ── Step 7: BGM 选择 ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "正在选择BGM...")
|
||||
# ── Step 7: BGM 选择(P1:暂返回 None,跳过) ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)")
|
||||
bgm = _step_bgm_select(job)
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 78.0, "BGM选择完成")
|
||||
|
||||
# ── Step 8: 渲染 ──
|
||||
# ── Step 8: 渲染(逐分镜 Seedance → concat → 混 TTS) ──
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...")
|
||||
video_path = _step_render(job, storyboard, audio_url, bgm)
|
||||
video_path = _step_render(job, storyboard, tts_path, bgm)
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成")
|
||||
|
||||
# ── Step 9: MuseTalk 数字人口型 ──
|
||||
# ── Step 9: MuseTalk(P1:端点未就绪前跳过) ──
|
||||
_emit_progress(job_id, ViralVideoStage.MUSETALK, 90.0, "正在处理数字人口型...")
|
||||
final_video_path = _step_musetalk(job, video_path)
|
||||
_emit_progress(job_id, ViralVideoStage.MUSETALK, 93.0, "数字人处理完成")
|
||||
@@ -549,11 +767,9 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...")
|
||||
video_url = _step_upload(job, final_video_path)
|
||||
|
||||
# 扣点
|
||||
job.credits_cost = CREDITS_VIRAL_VIDEO_COST
|
||||
# TODO: 调用 credits.deduct() 实际扣点
|
||||
# TODO: 调用 credits.deduct() 实际扣点(#1895 总开关为 false 时不扣,保留 TODO)
|
||||
|
||||
# 完成
|
||||
job.mark_completed(video_url)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url})
|
||||
|
||||
@@ -946,6 +946,7 @@ class ViralVideoJobModel(Base):
|
||||
# 结果与状态
|
||||
status = Column(String(30), nullable=False, default="pending", index=True)
|
||||
intent_result = Column(JSON, nullable=True)
|
||||
image_analysis = Column(JSON, nullable=True)
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
|
||||
@@ -34,6 +34,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
style_template_id=getattr(model, "style_template_id", "") or "",
|
||||
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
|
||||
intent_result=dict(model.intent_result) if model.intent_result else None,
|
||||
image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
credits_cost=model.credits_cost or 0,
|
||||
error_msg=model.error_msg or "",
|
||||
@@ -72,6 +73,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
style_template_id=job.style_template_id,
|
||||
status=job.status,
|
||||
intent_result=job.intent_result,
|
||||
image_analysis=job.image_analysis,
|
||||
result_video_url=job.result_video_url,
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
@@ -91,6 +93,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
raise ValueError(f"ViralVideoJob {job.id} not found")
|
||||
model.status = job.status
|
||||
model.intent_result = job.intent_result
|
||||
model.image_analysis = job.image_analysis
|
||||
model.result_video_url = job.result_video_url
|
||||
model.credits_cost = job.credits_cost
|
||||
model.error_msg = job.error_msg
|
||||
|
||||
@@ -96,6 +96,9 @@ class SharedSettings(BaseSettings):
|
||||
doubao_max_retries: int = 2
|
||||
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 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
|
||||
@@ -118,6 +118,8 @@ class ViralVideoJob:
|
||||
style_strength: str = StyleStrength.MEDIUM
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
# v1.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜)
|
||||
image_analysis: dict | None = None
|
||||
# 状态
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
status: ViralVideoStatus = ViralVideoStatus.PENDING
|
||||
|
||||
@@ -238,6 +238,147 @@ class DoubaoClient:
|
||||
logger.error("豆包视觉API调用最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
# ── 视频生成(Seedance 2.5,异步任务)────────────────────────────
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str = "9:16",
|
||||
resolution: str = "720p",
|
||||
generate_audio: bool = False,
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 文生/图生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
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
|
||||
if not prompt or not prompt.strip():
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||||
total_timeout = getattr(settings, "doubao_video_timeout", 600) or 600
|
||||
video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
|
||||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
|
||||
if image_url:
|
||||
content.append({"type": "image_url", "image_url": {"url": image_url}})
|
||||
|
||||
create_payload: dict[str, Any] = {
|
||||
"model": video_model,
|
||||
"content": content,
|
||||
"generate_audio": generate_audio,
|
||||
"ratio": ratio,
|
||||
"duration": int(duration),
|
||||
"resolution": resolution,
|
||||
"watermark": watermark,
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
create_url = f"{self.base_url}/contents/generations/tasks"
|
||||
|
||||
# 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)
|
||||
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 创建任务最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
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"
|
||||
while time.time() < deadline:
|
||||
try:
|
||||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
status = data.get("status", "")
|
||||
last_status = status
|
||||
if status == "succeeded":
|
||||
content_obj = data.get("content") or {}
|
||||
video_url = content_obj.get("video_url")
|
||||
if video_url:
|
||||
break
|
||||
last_error = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||||
break
|
||||
if status == "failed":
|
||||
err = data.get("error") or {}
|
||||
last_error = RuntimeError(f"task failed: {err.get('code','')} {err.get('message','')}")
|
||||
break
|
||||
if status in ("expired", "cancelled"):
|
||||
last_error = RuntimeError(f"task {status}")
|
||||
break
|
||||
# queued / running: 继续轮询
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
logger.debug("Seedance 轮询异常: %s", e)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if not video_url:
|
||||
logger.error("Seedance 任务未成功: task_id=%s status=%s err=%s", task_id, last_status, last_error)
|
||||
return None
|
||||
|
||||
# 3) 下载到本地
|
||||
try:
|
||||
import os as _os
|
||||
import uuid as _uuid
|
||||
|
||||
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"
|
||||
with httpx.stream("GET", video_url, timeout=300) as r:
|
||||
r.raise_for_status()
|
||||
with open(local_path, "wb") as f:
|
||||
for chunk in r.iter_bytes(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
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)
|
||||
return None
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -536,3 +536,36 @@ def call_vision(image_url: str, prompt: str) -> object:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return raw
|
||||
|
||||
|
||||
def call_video_generation(
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str = "9:16",
|
||||
resolution: str = "720p",
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生成视频段,返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
封装 ai_client.video_generation:提交异步任务→轮询→下载到本地。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
try:
|
||||
return client.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
generate_audio=False, # 我们自己混 TTS
|
||||
watermark=False,
|
||||
output_dir=output_dir,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
return None
|
||||
|
||||
@@ -436,16 +436,17 @@ class TestViralVideoPipeline:
|
||||
def test_bgm_select(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
# P1: BGM 素材未就绪前 _step_bgm_select 统一返回 None(跳过 BGM 混音)
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
bgm = _step_bgm_select(mock_job)
|
||||
assert "upbeat" in bgm
|
||||
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 == "bgm_default.mp3"
|
||||
assert bgm is None
|
||||
|
||||
|
||||
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
|
||||
@@ -505,8 +506,8 @@ class TestPipelineIntegration:
|
||||
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 = "https://audio.mp3"
|
||||
mock_bgm.return_value = "bgm_default.mp3"
|
||||
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_musetalk.return_value = "/tmp/video_final.mp4"
|
||||
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
||||
|
||||
@@ -0,0 +1,244 @@
|
||||
"""#2106 P0 修复单测:Seedance 对接、image_analysis 持久化、TTS Path 统一、BGM/MuseTalk 跳过。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# worker 容器 PYTHONPATH 包含 apps/worker(worker 侧代码使用顶层包名 services/、viral_video/)
|
||||
_WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
|
||||
if str(_WORKER_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_WORKER_ROOT))
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_job():
|
||||
return ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=15,
|
||||
user_copy_text="测试文案",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
|
||||
|
||||
# ── P0-2: _step_video_analysis import 路径 ──────────────────────────
|
||||
|
||||
|
||||
class TestVideoAnalysisImport:
|
||||
def test_no_reference_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = ""
|
||||
assert _step_video_analysis(mock_job) is None
|
||||
|
||||
def test_with_reference_returns_dict_or_none(self, mock_job):
|
||||
"""有参考视频 URL 时,不管分析成功/失败/占位,返回 dict(不抛异常)。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = "https://example.com/ref.mp4"
|
||||
result = _step_video_analysis(mock_job)
|
||||
# 允许占位/失败/真实返回,但绝不能抛异常
|
||||
assert result is None or isinstance(result, dict)
|
||||
|
||||
|
||||
# ── P0-3: image_analysis 字段 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestImageAnalysisField:
|
||||
def test_default_none(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
assert job.image_analysis is None
|
||||
|
||||
def test_persist_and_read(self, mock_job):
|
||||
mock_job.image_analysis = {"products": [{"name": "口红"}]}
|
||||
assert mock_job.image_analysis["products"][0]["name"] == "口红"
|
||||
|
||||
|
||||
# ── P0-1: storyboard 规范化 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStoryboardNormalize:
|
||||
def test_normalize_fills_defaults(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
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"] == ""
|
||||
|
||||
def test_normalize_scales_to_total_duration(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
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
|
||||
|
||||
def test_fallback_storyboard(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _fallback_storyboard
|
||||
|
||||
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)
|
||||
|
||||
def test_storyboard_llm_list(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
|
||||
|
||||
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 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTTSPath:
|
||||
def test_tts_returns_none_on_import_error(self, mock_job):
|
||||
"""get_tts_service 抛 ImportError 时 _step_tts 返回 None。"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
with patch("services.tts_service_factory.get_tts_service", side_effect=ImportError("no tts")):
|
||||
assert vv._step_tts(mock_job, "文案") is None
|
||||
|
||||
def test_tts_returns_none_when_path_not_exists(self, mock_job, tmp_path):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
fake_service = MagicMock()
|
||||
fake_service.synthesize.return_value = str(tmp_path / "not_exist.mp3")
|
||||
with patch("services.tts_service_factory.get_tts_service", return_value=fake_service):
|
||||
assert vv._step_tts(mock_job, "文案") is None
|
||||
|
||||
def test_tts_returns_path_when_exists(self, mock_job, tmp_path):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
audio = tmp_path / "voice.mp3"
|
||||
audio.write_bytes(b"ID3fake")
|
||||
fake_service = MagicMock()
|
||||
fake_service.synthesize.return_value = audio
|
||||
with patch("services.tts_service_factory.get_tts_service", return_value=fake_service):
|
||||
result = vv._step_tts(mock_job, "文案")
|
||||
assert isinstance(result, Path)
|
||||
assert result.exists()
|
||||
|
||||
|
||||
# ── P1: BGM 跳过 / MuseTalk 无 persona 跳过 ───────────────────────
|
||||
|
||||
|
||||
class TestBGMSkip:
|
||||
def test_bgm_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
assert _step_bgm_select(mock_job) is None
|
||||
|
||||
|
||||
class TestMuseTalkSkip:
|
||||
def test_musetalk_no_persona_passes_through(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_musetalk
|
||||
|
||||
mock_job.persona_id = ""
|
||||
assert _step_musetalk(mock_job, "/tmp/video.mp4") == "/tmp/video.mp4"
|
||||
|
||||
def test_musetalk_with_persona_still_skips(self, mock_job):
|
||||
"""GPU 端点未就绪前,即使有 persona 也直接返回原路径。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_musetalk
|
||||
|
||||
mock_job.persona_id = "persona-1"
|
||||
assert _step_musetalk(mock_job, "/tmp/video.mp4") == "/tmp/video.mp4"
|
||||
|
||||
|
||||
# ── P0-1: call_video_generation 参数构造 ──────────────────────────
|
||||
|
||||
|
||||
class TestCallVideoGeneration:
|
||||
def test_returns_none_when_client_unavailable(self):
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get.return_value = mock_client
|
||||
assert call_video_generation("prompt") is None
|
||||
|
||||
def test_delegates_to_client(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
|
||||
result = call_video_generation(prompt="测试", image_url="https://img/x.jpg", duration=5, ratio="9:16")
|
||||
assert result == str(out)
|
||||
mock_client.video_generation.assert_called_once()
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
assert kwargs["prompt"] == "测试"
|
||||
assert kwargs["image_url"] == "https://img/x.jpg"
|
||||
assert kwargs["duration"] == 5
|
||||
|
||||
|
||||
# ── P0-1: _step_render 占位片段生成 ──────────────────────────────
|
||||
|
||||
|
||||
class TestPlaceholderClip:
|
||||
def test_make_placeholder_clip(self, tmp_path):
|
||||
import shutil
|
||||
|
||||
from apps.worker.worker_app.tasks.viral_video import _make_placeholder_clip, _probe_ok
|
||||
|
||||
if not shutil.which("ffmpeg"):
|
||||
pytest.skip("ffmpeg not available")
|
||||
|
||||
out = _make_placeholder_clip(tmp_path, 0, 3)
|
||||
assert out.exists()
|
||||
assert _probe_ok(str(out))
|
||||
|
||||
|
||||
# ── P0-1: DoubaoClient.video_generation 在不可用时返回 None ───────
|
||||
|
||||
|
||||
class TestDoubaoClientVideoGen:
|
||||
def test_unavailable_returns_none(self):
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = "" # is_available -> False
|
||||
assert client.video_generation("prompt") is None
|
||||
|
||||
|
||||
# ── P0-3: resume 从 job 读 image_analysis ────────────────────────
|
||||
|
||||
|
||||
class TestResumeReadsImageAnalysis:
|
||||
def test_resume_uses_persisted_image_analysis(self):
|
||||
"""resume_pipeline 应从 job.image_analysis 读(P0-3 持久化)。"""
|
||||
import inspect
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "job.image_analysis" in src
|
||||
Reference in New Issue
Block a user