Files
xiaoxia-saas/packages/adapters/sqlalchemy_impl/viral_video_repository.py
T
xiaoxia d08835ec9f 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
2026-09-30 19:06:33 +08:00

203 lines
7.4 KiB
Python
Executable File

"""爆款视频任务 SQLAlchemy 仓储实现。"""
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import (
ViralVideoJobModel,
ViralVideoPromptTemplateModel,
ViralVideoStyleTemplateModel,
)
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
"""ORM → 领域实体。"""
return ViralVideoJob(
id=model.id,
user_id=model.user_id,
images=list(model.images or []),
industry=model.industry or "",
target_customer=model.target_customer or "",
persona_id=model.persona_id or "",
viral_structure=model.viral_structure or "",
marketing_purpose=model.marketing_purpose or "",
bgm_preference=model.bgm_preference or "",
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 "",
reference_video_url=getattr(model, "reference_video_url", "") or "",
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 "",
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 "",
retry_count=model.retry_count or 0,
started_at=model.started_at,
completed_at=model.completed_at,
created_at=model.created_at,
updated_at=model.updated_at,
)
class SQLAlchemyViralVideoJobRepository:
"""爆款视频任务仓储。"""
def __init__(self, session: Session):
self.session = session
def save(self, job: ViralVideoJob) -> ViralVideoJob:
model = ViralVideoJobModel(
id=job.id,
user_id=job.user_id,
images=job.images,
industry=job.industry,
target_customer=job.target_customer,
persona_id=job.persona_id,
viral_structure=job.viral_structure,
marketing_purpose=job.marketing_purpose,
bgm_preference=job.bgm_preference,
duration=job.duration,
user_copy_text=job.user_copy_text,
fusion_level=job.fusion_level,
reference_audio_path=job.reference_audio_path,
reference_video_url=job.reference_video_url,
style_strength=job.style_strength,
style_guide=job.style_guide,
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,
retry_count=job.retry_count,
started_at=job.started_at,
completed_at=job.completed_at,
created_at=job.created_at,
updated_at=job.updated_at,
)
self.session.add(model)
self.session.commit()
return job
def update(self, job: ViralVideoJob) -> None:
model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job.id).first()
if model is None:
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
model.retry_count = job.retry_count
model.started_at = job.started_at
model.completed_at = job.completed_at
model.style_guide = job.style_guide
model.updated_at = datetime.now(timezone.utc)
self.session.commit()
def get(self, job_id: str) -> ViralVideoJob | None:
model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job_id).first()
if model is None:
return None
return _to_domain(model)
def list_by_user(self, user_id: str, limit: int = 50, offset: int = 0) -> list[ViralVideoJob]:
models = (
self.session.query(ViralVideoJobModel)
.filter(ViralVideoJobModel.user_id == user_id)
.order_by(ViralVideoJobModel.created_at.desc())
.offset(offset)
.limit(limit)
.all()
)
return [_to_domain(m) for m in models]
def count_pending_by_user(self, user_id: str) -> int:
return (
self.session.query(ViralVideoJobModel)
.filter(
ViralVideoJobModel.user_id == user_id,
ViralVideoJobModel.status.in_(["pending", "running", "wait_user_confirm"]),
)
.count()
)
class SQLAlchemyViralVideoStyleTemplateRepository:
"""风格模板仓储。"""
def __init__(self, session: Session):
self.session = session
def list_all(self) -> list[dict]:
models = (
self.session.query(ViralVideoStyleTemplateModel)
.order_by(ViralVideoStyleTemplateModel.sort_order.asc())
.all()
)
return [
{
"id": m.id,
"name": m.name,
"description": m.description or "",
"thumbnail_url": m.thumbnail_url or "",
"style_config": dict(m.style_config) if m.style_config else {},
"is_system": m.is_system,
}
for m in models
]
def get(self, template_id: str) -> dict | None:
model = (
self.session.query(ViralVideoStyleTemplateModel)
.filter(ViralVideoStyleTemplateModel.id == template_id)
.first()
)
if model is None:
return None
return {
"id": model.id,
"name": model.name,
"description": model.description or "",
"thumbnail_url": model.thumbnail_url or "",
"style_config": dict(model.style_config) if model.style_config else {},
"is_system": model.is_system,
}
class SQLAlchemyViralVideoPromptTemplateRepository:
"""Prompt 模板仓储(由 #2040 seed,这里只读取)。"""
def __init__(self, session: Session):
self.session = session
def get_active_by_type(self, prompt_type: str) -> dict | None:
model = (
self.session.query(ViralVideoPromptTemplateModel)
.filter(
ViralVideoPromptTemplateModel.prompt_type == prompt_type,
ViralVideoPromptTemplateModel.is_active.is_(True),
)
.order_by(ViralVideoPromptTemplateModel.version.desc())
.first()
)
if model is None:
return None
return {
"id": model.id,
"prompt_type": model.prompt_type,
"name": model.name,
"content": model.content,
"variables": list(model.variables or []),
"version": model.version,
}