Compare commits
31 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5e61dbe4f9 | |||
| 22e04d65a7 | |||
| 6ff57b2feb | |||
| 2981d20d5b | |||
| 6cddd72910 | |||
| 6d5c44d6be | |||
| 665a3063b6 | |||
| 24724dca9f | |||
| d08835ec9f | |||
| 77ce4a1a0d | |||
| 7ad722e6c6 | |||
| b54dda6526 | |||
| 69da326ed6 | |||
| f7f600d091 | |||
| bf9249da19 | |||
| ca834b23cb | |||
| 37f7aa3329 | |||
| 794f5f374b | |||
| 34305974ad | |||
| e83a7cad2e | |||
| c45a2ce9b1 | |||
| 9814fcdc22 | |||
| 6636dc45f7 | |||
| e11e4f0e99 | |||
| a7d6ba473b | |||
| 966da04c9c | |||
| fdeb792bab | |||
| ff1d878c62 | |||
| f19be5fd09 | |||
| eeb8a05b69 | |||
| 0a004db1bd |
+100
@@ -0,0 +1,100 @@
|
||||
"""add viral video tables
|
||||
|
||||
Revision ID: 086_add_viral_video_tables
|
||||
Revises: 085_atom_clip_caption_embedding
|
||||
Create Date: 2026-09-28
|
||||
|
||||
新增爆款视频相关表:
|
||||
- viral_video_jobs: 爆款视频任务
|
||||
- viral_video_style_templates: 风格模板配置
|
||||
- viral_video_prompt_templates: Prompt 模板(由 #2040 seed)
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "086_add_viral_video_tables"
|
||||
down_revision = "085_atom_clip_caption_embedding"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# viral_video_jobs
|
||||
op.create_table(
|
||||
"viral_video_jobs",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, index=True),
|
||||
sa.Column("images", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("industry", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("target_customer", sa.String(500), nullable=False, server_default=""),
|
||||
sa.Column("persona_id", sa.String(36), nullable=False, server_default=""),
|
||||
sa.Column("viral_structure", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("marketing_purpose", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("bgm_preference", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("duration", sa.Integer(), nullable=False, server_default="30"),
|
||||
sa.Column("user_copy_text", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("fusion_level", sa.String(20), nullable=False, server_default="ai_polish"),
|
||||
sa.Column("reference_audio_path", sa.String(1000), nullable=False, server_default=""),
|
||||
# v1.3 新增
|
||||
sa.Column("reference_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_strength", sa.String(20), nullable=False, server_default="medium"),
|
||||
sa.Column("style_guide", sa.JSON(), nullable=True),
|
||||
sa.Column("style_template_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
# 状态与结果
|
||||
sa.Column("status", sa.String(30), nullable=False, server_default="pending", index=True),
|
||||
sa.Column("intent_result", sa.JSON(), nullable=True),
|
||||
sa.Column("result_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("credits_cost", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error_msg", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("retry_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_style_templates
|
||||
op.create_table(
|
||||
"viral_video_style_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("thumbnail_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_config", sa.JSON(), nullable=False, server_default="{}"),
|
||||
sa.Column("is_system", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("sort_order", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_prompt_templates
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("prompt_type", sa.String(50), nullable=False, index=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("variables", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("version", sa.Integer(), nullable=False, server_default="1"),
|
||||
sa.Column("is_active", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# Seed 默认风格模板
|
||||
op.execute("""
|
||||
INSERT INTO viral_video_style_templates (id, name, description, style_config, is_system, sort_order)
|
||||
VALUES
|
||||
('style-tpl-001', '快节奏冲击', '高频切镜+动感BGM,适合食品饮料等快消品', '{"cut_speed": "fast", "transition": "jump_cut", "energy": "high"}', true, 1),
|
||||
('style-tpl-002', '质感慢镜', '慢节奏+电影感调色,适合美妆护肤珠宝', '{"cut_speed": "slow", "transition": "dissolve", "energy": "low", "color_grade": "cinematic"}', true, 2),
|
||||
('style-tpl-003', '口播种草', '数字人口播+产品特写穿插', '{"cut_speed": "medium", "transition": "cross_dissolve", "has_talking_head": true}', true, 3),
|
||||
('style-tpl-004', '场景叙事', '多场景切换+故事线叙述', '{"cut_speed": "medium", "transition": "wipe", "narrative": true}', true, 4)
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
op.drop_table("viral_video_style_templates")
|
||||
op.drop_table("viral_video_jobs")
|
||||
@@ -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")
|
||||
@@ -36,6 +36,7 @@ from app.api.routes.titles import router as titles_router
|
||||
from app.api.routes.tts import router as tts_router
|
||||
from app.api.routes.upload import router as upload_router
|
||||
from app.api.routes.videos import router as videos_router
|
||||
from app.api.routes.viral_video import router as viral_video_router
|
||||
from app.api.routes.voice_clones import router as voice_clones_router
|
||||
from app.api.routes.voices import router as voices_router
|
||||
from fastapi import APIRouter
|
||||
@@ -240,3 +241,4 @@ api_router.include_router(
|
||||
prefix="/gpu",
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
api_router.include_router(viral_video_router, prefix="/viral-video", tags=["爆款视频"])
|
||||
|
||||
@@ -0,0 +1,549 @@
|
||||
"""爆款视频 API 路由。
|
||||
|
||||
端点:
|
||||
POST /api/v1/viral-video/generate 创建爆款视频任务
|
||||
GET /api/v1/viral-video/{job_id} 查询任务状态
|
||||
GET /api/v1/viral-video/history 历史记录
|
||||
POST /api/v1/viral-video/{job_id}/retry 重试失败任务
|
||||
POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案
|
||||
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
|
||||
GET /api/v1/viral-video/style-templates 获取风格模板列表
|
||||
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送(订阅 Redis pub/sub)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
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 (
|
||||
AnalyzeStyleRequest,
|
||||
AnalyzeStyleResponse,
|
||||
ConfirmIntentRequest,
|
||||
CreateViralVideoRequest,
|
||||
StyleTemplateListResponse,
|
||||
StyleTemplateResponse,
|
||||
ViralVideoHistoryResponse,
|
||||
ViralVideoJobResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _to_response(job) -> ViralVideoJobResponse:
|
||||
return ViralVideoJobResponse(
|
||||
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,
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _get_job_repo(session: Session) -> SQLAlchemyViralVideoJobRepository:
|
||||
return SQLAlchemyViralVideoJobRepository(session)
|
||||
|
||||
|
||||
def _get_style_repo(session: Session) -> SQLAlchemyViralVideoStyleTemplateRepository:
|
||||
return SQLAlchemyViralVideoStyleTemplateRepository(session)
|
||||
|
||||
|
||||
# ── Endpoints ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/generate", response_model=ViralVideoJobResponse)
|
||||
def create_viral_video(
|
||||
request: CreateViralVideoRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""创建爆款视频任务,入队 Celery 编排器。"""
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
repo = _get_job_repo(session)
|
||||
|
||||
# 创建领域实体
|
||||
job = ViralVideoJob(
|
||||
user_id=authenticated_user.user.id,
|
||||
images=list(request.images),
|
||||
industry=request.industry,
|
||||
target_customer=request.target_customer,
|
||||
persona_id=request.persona_id,
|
||||
viral_structure=request.viral_structure,
|
||||
marketing_purpose=request.marketing_purpose,
|
||||
bgm_preference=request.bgm_preference,
|
||||
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,
|
||||
)
|
||||
|
||||
# 持久化
|
||||
repo.save(job)
|
||||
|
||||
# 入队 Celery 任务
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 任务已入队: job_id=%s user_id=%s", job.id, job.user_id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 入队失败: %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,
|
||||
offset: int = 0,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoHistoryResponse:
|
||||
"""获取用户的爆款视频历史列表。"""
|
||||
repo = _get_job_repo(session)
|
||||
jobs = repo.list_by_user(authenticated_user.user.id, limit=limit, offset=offset)
|
||||
items = [_to_response(j) for j in jobs]
|
||||
return ViralVideoHistoryResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.get("/style-templates", response_model=StyleTemplateListResponse)
|
||||
def list_style_templates(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> StyleTemplateListResponse:
|
||||
"""获取风格模板列表。"""
|
||||
repo = _get_style_repo(session)
|
||||
templates = repo.list_all()
|
||||
items = [
|
||||
StyleTemplateResponse(
|
||||
id=t["id"],
|
||||
name=t["name"],
|
||||
description=t["description"],
|
||||
thumbnail_url=t["thumbnail_url"],
|
||||
style_config=t["style_config"],
|
||||
)
|
||||
for t in templates
|
||||
]
|
||||
return StyleTemplateListResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
|
||||
def get_viral_video_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""查询爆款视频任务状态。"""
|
||||
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="无权查看此任务")
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/retry", response_model=ViralVideoJobResponse)
|
||||
def retry_viral_video_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""重试失败的爆款视频任务。"""
|
||||
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.FAILED:
|
||||
raise HTTPException(status_code=409, detail="只有失败的任务可以重试")
|
||||
|
||||
# 重置状态
|
||||
job.retry_count += 1
|
||||
job.status = ViralVideoStatus.PENDING
|
||||
job.error_msg = ""
|
||||
job.started_at = None
|
||||
job.completed_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", job.id, job.retry_count)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"重试入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/confirm-intent", response_model=ViralVideoJobResponse)
|
||||
def confirm_intent(
|
||||
job_id: str,
|
||||
request: ConfirmIntentRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""用户确认/修改 AI 生成的意图文案,恢复流水线。"""
|
||||
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.WAIT_USER_CONFIRM:
|
||||
raise HTTPException(status_code=409, detail="任务当前不在等待确认状态")
|
||||
|
||||
# 更新文案
|
||||
if request.confirmed_copy:
|
||||
job.user_copy_text = request.confirmed_copy
|
||||
|
||||
# 恢复流水线
|
||||
job.resume_from_confirm()
|
||||
repo.update(job)
|
||||
|
||||
# 从断点恢复 Celery 任务
|
||||
try:
|
||||
celery_app.send_task("worker.resume_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 意图确认,恢复流水线: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"恢复流水线失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/analyze-style", response_model=AnalyzeStyleResponse)
|
||||
def analyze_style(
|
||||
job_id: str,
|
||||
request: AnalyzeStyleRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> AnalyzeStyleResponse:
|
||||
"""触发参考视频风格分析(独立步骤,可在生成前单独调用)。"""
|
||||
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="无权操作此任务")
|
||||
|
||||
# 更新参考视频 URL
|
||||
job.reference_video_url = request.reference_video_url
|
||||
if request.style_template_id:
|
||||
job.style_template_id = request.style_template_id
|
||||
repo.update(job)
|
||||
|
||||
# 入队风格分析任务
|
||||
try:
|
||||
celery_app.send_task("worker.run_video_style_analysis", args=[job.id])
|
||||
logger.info("[爆款视频] 风格分析入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 风格分析入队失败: %s", e, exc_info=True)
|
||||
|
||||
return AnalyzeStyleResponse(
|
||||
job_id=job.id,
|
||||
status="analyzing",
|
||||
style_guide=None,
|
||||
)
|
||||
|
||||
|
||||
# ── WebSocket 进度推送 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ws_authenticate_user(token: str):
|
||||
"""从 token 字符串解析用户(复用 HTTP Bearer 的解码 + 黑名单逻辑)。
|
||||
|
||||
WebSocket 握手阶段不能发自定义 Authorization header,
|
||||
因此统一通过 query 参数 ``?token=...`` 传 JWT。
|
||||
"""
|
||||
from app.auth import _decode_user_token
|
||||
from app.dependencies import get_user_repository
|
||||
|
||||
if not token:
|
||||
return None
|
||||
try:
|
||||
payload = _decode_user_token(token)
|
||||
except Exception:
|
||||
return None
|
||||
user_id = payload.get("sub")
|
||||
if not isinstance(user_id, str) or not user_id:
|
||||
return None
|
||||
# 同步场景下手动拉 repository 实例
|
||||
from app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
user_repo = get_user_repository(session)
|
||||
user = user_repo.find_by_id(user_id)
|
||||
return user
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
async def _run_pubsub_forwarder(
|
||||
websocket, redis_lib, settings, job_id: str
|
||||
) -> None: # pragma: no cover - integration tested (real Redis + thread)
|
||||
"""订阅 Redis 频道并把消息桥接到 WebSocket,终态消息后自动关闭。
|
||||
|
||||
该函数封装了线程 + asyncio.Queue 桥接逻辑,在单测中可被整体替换为桩,
|
||||
避免引入真实 Redis 与线程调度的不确定性。
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
|
||||
r = redis_lib.from_url(settings.REDIS_URL, decode_responses=True)
|
||||
pubsub = r.pubsub(ignore_subscribe_messages=True)
|
||||
channel = f"viral_video:{job_id}"
|
||||
pubsub.subscribe(channel)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
queue: asyncio.Queue = asyncio.Queue(maxsize=64)
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
def _reader() -> None:
|
||||
try:
|
||||
while not stop_event.is_set():
|
||||
msg = pubsub.get_message(timeout=0.5)
|
||||
if msg is None or msg.get("type") != "message":
|
||||
continue
|
||||
raw = msg.get("data")
|
||||
if not isinstance(raw, str):
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(raw)
|
||||
except Exception:
|
||||
payload = {"type": "viral_video:progress", "data": {"raw": raw}}
|
||||
loop.call_soon_threadsafe(queue.put_nowait, payload)
|
||||
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
|
||||
loop.call_soon_threadsafe(stop_event.set)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频WS] pubsub reader 异常退出: %s", e)
|
||||
loop.call_soon_threadsafe(stop_event.set)
|
||||
|
||||
try:
|
||||
reader_thread = threading.Thread(target=_reader, name=f"viral-video-ws-{job_id}", daemon=True)
|
||||
reader_thread.start()
|
||||
|
||||
while not stop_event.is_set():
|
||||
try:
|
||||
payload = await asyncio.wait_for(queue.get(), timeout=1.0)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
try:
|
||||
await websocket.send_json(payload)
|
||||
except Exception:
|
||||
break
|
||||
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
|
||||
break
|
||||
except WebSocketDisconnect:
|
||||
logger.info("[爆款视频WS] 客户端断开: job_id=%s", job_id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频WS] 转发异常: %s", e, exc_info=True)
|
||||
try:
|
||||
await websocket.send_json({"type": "viral_video:error", "message": f"服务异常: {e}"})
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
stop_event.set()
|
||||
try:
|
||||
pubsub.unsubscribe(channel)
|
||||
pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
r.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@router.websocket("/ws/{job_id}")
|
||||
async def viral_video_websocket(websocket: WebSocket, job_id: str) -> None:
|
||||
"""WebSocket 桥接:订阅 Redis `viral_video:{job_id}` 频道并转发给前端。
|
||||
|
||||
认证:通过 ``?token=<jwt>`` query 参数传 JWT(浏览器 WS 握手不支持自定义 header)。
|
||||
事件类型:
|
||||
- viral_video:progress 中间进度(progress: 0-100)
|
||||
- viral_video:wait_user 等待用户确认意图文案
|
||||
- viral_video:completed 任务完成(data.video_url)
|
||||
- viral_video:failed 任务失败(data.error)
|
||||
- viral_video:error 服务端错误(如鉴权失败 / job 不存在 / 无权限)
|
||||
"""
|
||||
|
||||
import redis as redis_lib
|
||||
from app.config import settings
|
||||
|
||||
# ── 1. 鉴权 ──────────────────────────────────────────────────────
|
||||
token = websocket.query_params.get("token", "")
|
||||
user = _ws_authenticate_user(token)
|
||||
if user is None:
|
||||
await websocket.close(code=4401, reason="Unauthorized")
|
||||
return
|
||||
|
||||
# ── 2. 校验 job 归属 ─────────────────────────────────────────────
|
||||
from app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
job_repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
await websocket.close(code=4404, reason="Job not found")
|
||||
return
|
||||
if job.user_id != user.id:
|
||||
await websocket.close(code=4403, reason="Forbidden")
|
||||
return
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
await websocket.accept()
|
||||
|
||||
# ── 3. 发送一条初始状态(前端连接后立即拿到当前进度) ────────────
|
||||
try:
|
||||
session = SessionLocal()
|
||||
job_repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is not None:
|
||||
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
initial = {
|
||||
"type": "viral_video:progress",
|
||||
"job_id": job_id,
|
||||
"stage": _stage_from_status(job),
|
||||
"progress": _estimate_progress(job),
|
||||
"message": _initial_message(job),
|
||||
"data": {"status": status_val},
|
||||
}
|
||||
await websocket.send_json(initial)
|
||||
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
|
||||
if job.is_terminal:
|
||||
is_completed = status_val == "completed"
|
||||
terminal_type = "viral_video:completed" if is_completed else "viral_video:failed"
|
||||
terminal_data = (
|
||||
{"video_url": job.result_video_url or ""} if is_completed else {"error": job.error_msg or ""}
|
||||
)
|
||||
await websocket.send_json(
|
||||
{
|
||||
"type": terminal_type,
|
||||
"job_id": job_id,
|
||||
"stage": "",
|
||||
"progress": 100 if is_completed else 0,
|
||||
"message": "视频生成完成" if is_completed else "任务失败",
|
||||
"data": terminal_data,
|
||||
}
|
||||
)
|
||||
await websocket.close()
|
||||
return
|
||||
session.close()
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频WS] 发送初始状态失败: %s", e)
|
||||
try:
|
||||
session.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ── 4. 订阅 Redis 频道并转发 ─────────────────────────────────────
|
||||
# redis-py 的 pubsub 是同步阻塞的,放到线程里跑,通过 asyncio.Queue 桥接到 event loop。
|
||||
# 该段依赖真实 Redis + 线程调度,属于集成测试范围,单测通过桩替换。
|
||||
await _run_pubsub_forwarder(websocket, redis_lib, settings, job_id)
|
||||
|
||||
|
||||
def _job_status(job) -> str:
|
||||
return job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
|
||||
|
||||
# 初始快照的 stage 推断:领域对象不持久化 stage,
|
||||
# 只能根据 status 给一个占位,后续 worker 推送的真实进度事件会覆盖。
|
||||
_STATUS_STAGE = {
|
||||
"pending": "",
|
||||
"running": "",
|
||||
"wait_user_confirm": "intent_parsing",
|
||||
"completed": "uploading",
|
||||
"failed": "",
|
||||
"cancelled": "",
|
||||
}
|
||||
|
||||
_STATUS_PROGRESS = {
|
||||
"pending": 0.0,
|
||||
"running": 5.0,
|
||||
"wait_user_confirm": 35.0,
|
||||
"completed": 100.0,
|
||||
"failed": 0.0,
|
||||
"cancelled": 0.0,
|
||||
}
|
||||
|
||||
_STATUS_MESSAGE = {
|
||||
"pending": "任务已创建,等待执行",
|
||||
"running": "任务执行中",
|
||||
"wait_user_confirm": "等待用户确认意图文案",
|
||||
"completed": "视频生成完成",
|
||||
"failed": "任务失败",
|
||||
"cancelled": "任务已取消",
|
||||
}
|
||||
|
||||
|
||||
def _stage_from_status(job) -> str:
|
||||
return _STATUS_STAGE.get(_job_status(job), "")
|
||||
|
||||
|
||||
def _estimate_progress(job) -> float:
|
||||
"""根据 status 粗略估算百分比(0-100),用于连接初始快照;
|
||||
连接建立后由 Redis 推送的真实事件持续更新。
|
||||
"""
|
||||
return _STATUS_PROGRESS.get(_job_status(job), 5.0)
|
||||
|
||||
|
||||
def _initial_message(job) -> str:
|
||||
"""给新连接的前端一个可读的初始状态文案。"""
|
||||
status_val = _job_status(job)
|
||||
if status_val == "failed" and job.error_msg:
|
||||
return f"任务失败: {job.error_msg}"
|
||||
return _STATUS_MESSAGE.get(status_val, "任务准备中")
|
||||
+2
-2
@@ -7,9 +7,9 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.schema_guard import assert_auto_create_schema_allowed
|
||||
|
||||
ensure_database_exists(settings.DATABASE_URL)
|
||||
ensure_database_exists(settings.effective_database_url)
|
||||
engine, SessionLocal = build_session_factory(
|
||||
settings.DATABASE_URL,
|
||||
settings.effective_database_url,
|
||||
pool_size=settings.DATABASE_POOL_SIZE,
|
||||
max_overflow=settings.DATABASE_MAX_OVERFLOW,
|
||||
pool_timeout=settings.DATABASE_POOL_TIMEOUT,
|
||||
|
||||
@@ -56,7 +56,7 @@ from packages.adapters.sqlalchemy_impl.voice_library_repository import (
|
||||
from packages.ports.tag_repository import TagRepository
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
|
||||
_engine, _SessionLocal = build_session_factory(settings.effective_database_url)
|
||||
|
||||
|
||||
def get_db_session() -> Generator[Session, None, None]:
|
||||
|
||||
Executable
+156
@@ -0,0 +1,156 @@
|
||||
"""爆款视频 API schemas。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
# ── 枚举常量 ─────────────────────────────────────────────────────────────
|
||||
|
||||
VALID_FUSION_LEVELS = ("ai_full", "ai_polish", "user_primary")
|
||||
VALID_STYLE_STRENGTHS = ("light", "medium", "strict")
|
||||
VALID_STAGES = (
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"review",
|
||||
"tts",
|
||||
"bgm_select",
|
||||
"rendering",
|
||||
"musetalk",
|
||||
"uploading",
|
||||
)
|
||||
|
||||
|
||||
# ── Request Schemas ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class CreateViralVideoRequest(BaseModel):
|
||||
"""创建爆款视频任务请求。"""
|
||||
|
||||
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 _validate_fusion_level(cls, v: str) -> str:
|
||||
if v not in VALID_FUSION_LEVELS:
|
||||
raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一")
|
||||
return v
|
||||
|
||||
@field_validator("style_strength")
|
||||
@classmethod
|
||||
def _validate_style_strength(cls, v: str) -> str:
|
||||
if v not in VALID_STYLE_STRENGTHS:
|
||||
raise ValueError(f"style_strength 必须是 {VALID_STYLE_STRENGTHS} 之一")
|
||||
return v
|
||||
|
||||
|
||||
class ConfirmIntentRequest(BaseModel):
|
||||
"""确认意图请求(confirm-intent)。"""
|
||||
|
||||
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 = Field(default="", description="风格模板 ID(可选覆盖)")
|
||||
|
||||
|
||||
# ── Response Schemas ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ViralVideoJobResponse(BaseModel):
|
||||
"""爆款视频任务响应。"""
|
||||
|
||||
id: str
|
||||
user_id: str
|
||||
images: list[str] = Field(default_factory=list)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
duration: int = 30
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = "ai_polish"
|
||||
reference_audio_path: str = ""
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = "medium"
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
status: str
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class ViralVideoHistoryResponse(BaseModel):
|
||||
"""历史记录列表响应。"""
|
||||
|
||||
items: list[ViralVideoJobResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class StyleTemplateResponse(BaseModel):
|
||||
"""风格模板响应。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
thumbnail_url: str = ""
|
||||
style_config: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class StyleTemplateListResponse(BaseModel):
|
||||
"""风格模板列表响应。"""
|
||||
|
||||
items: list[StyleTemplateResponse]
|
||||
|
||||
|
||||
class AnalyzeStyleResponse(BaseModel):
|
||||
"""风格分析结果响应。"""
|
||||
|
||||
job_id: str
|
||||
status: str
|
||||
style_guide: dict | None = None
|
||||
|
||||
|
||||
# ── WebSocket 事件 Schema ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class WSProgressEvent(BaseModel):
|
||||
"""WebSocket 进度推送事件。"""
|
||||
|
||||
type: str = "viral_video:progress"
|
||||
job_id: str
|
||||
stage: str
|
||||
progress: float = Field(ge=0.0, le=100.0)
|
||||
message: str = ""
|
||||
data: dict = Field(default_factory=dict)
|
||||
@@ -0,0 +1,47 @@
|
||||
import apiClient from "@/api/client"
|
||||
import type {
|
||||
GenerateViralVideoRequest,
|
||||
HistoryResponse,
|
||||
StyleTemplate,
|
||||
ViralVideoJob,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
export function generateViralVideo(payload: GenerateViralVideoRequest) {
|
||||
return apiClient.post<ViralVideoJob>("/viral-video/generate", payload).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 查询单个任务 */
|
||||
export function getViralVideoJob(id: string) {
|
||||
return apiClient.get<ViralVideoJob>(`/viral-video/${id}`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 用户确认/修改 AI 理解的意图后继续 */
|
||||
export function confirmViralVideoIntent(
|
||||
id: string,
|
||||
payload: { confirmed_copy?: string; edits?: Record<string, unknown> },
|
||||
) {
|
||||
return apiClient
|
||||
.post<ViralVideoJob>(`/viral-video/${id}/confirm-intent`, payload)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 重试失败任务 */
|
||||
export function retryViralVideo(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/retry`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 历史记录(分页) */
|
||||
export function getViralVideoHistory(params?: { page?: number; page_size?: number }) {
|
||||
return apiClient.get<HistoryResponse>("/viral-video/history", { params }).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 预设风格模板 */
|
||||
export function getViralStyleTemplates() {
|
||||
return apiClient.get<StyleTemplate[]>("/viral-video/style-templates").then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 上传参考视频后触发风格分析(返回带 style_guide 的任务详情) */
|
||||
export function analyzeViralStyle(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
export type FusionLevel = "full_ai" | "polish" | "as_is"
|
||||
export const FUSION_LEVELS: { value: FusionLevel; label: string; desc: string }[] = [
|
||||
{ value: "full_ai", label: "AI 全写", desc: "给我方向,全由AI创作" },
|
||||
{ value: "polish", label: "AI润色", desc: "我写草稿,AI帮我润色" },
|
||||
{ value: "as_is", label: "按我写的来", desc: "几乎不改我的文案" },
|
||||
]
|
||||
|
||||
export type StyleStrength = "light" | "medium" | "strict"
|
||||
export const STYLE_STRENGTHS: { value: StyleStrength; label: string }[] = [
|
||||
{ value: "light", label: "轻度借鉴" },
|
||||
{ value: "medium", label: "中度参考" },
|
||||
{ value: "strict", label: "深度模仿" },
|
||||
]
|
||||
|
||||
export type ViralVideoStatus =
|
||||
"pending" | "running" | "wait_user_confirm" | "completed" | "failed" | "cancelled"
|
||||
|
||||
export type ViralVideoStage =
|
||||
| "image_analysis"
|
||||
| "video_analysis"
|
||||
| "intent_parsing"
|
||||
| "copy_fusion"
|
||||
| "storyboard"
|
||||
| "review"
|
||||
| "tts"
|
||||
| "bgm_select"
|
||||
| "rendering"
|
||||
| "musetalk"
|
||||
| "uploading"
|
||||
|
||||
export interface StageDisplay {
|
||||
label: string
|
||||
/** 插值到的总体进度百分比 */
|
||||
pct: number
|
||||
}
|
||||
|
||||
export const STAGE_DISPLAYS: Record<ViralVideoStage, StageDisplay> = {
|
||||
image_analysis: { label: "图片分析", pct: 15 },
|
||||
video_analysis: { label: "参考视频风格分析", pct: 25 },
|
||||
intent_parsing: { label: "意图理解", pct: 35 },
|
||||
copy_fusion: { label: "文案融合创作", pct: 50 },
|
||||
storyboard: { label: "分镜生成", pct: 60 },
|
||||
review: { label: "AI审核", pct: 70 },
|
||||
tts: { label: "配音生成", pct: 78 },
|
||||
bgm_select: { label: "BGM匹配", pct: 85 },
|
||||
rendering: { label: "视频渲染", pct: 92 },
|
||||
musetalk: { label: "口型同步", pct: 97 },
|
||||
uploading: { label: "上传成片", pct: 100 },
|
||||
}
|
||||
|
||||
export interface StyleTemplate {
|
||||
id: string
|
||||
name: string
|
||||
description?: string
|
||||
preview_url?: string
|
||||
tags?: string[]
|
||||
}
|
||||
|
||||
export interface IntentResult {
|
||||
product: string
|
||||
selling_points: string[]
|
||||
target_audience: string
|
||||
tone: string
|
||||
structure: string
|
||||
duration: number
|
||||
suggested_title?: string
|
||||
suggested_copy?: string
|
||||
}
|
||||
|
||||
export interface ViralVideoJob {
|
||||
id: string
|
||||
status: ViralVideoStatus
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
style_guide?: string
|
||||
user_copy_text?: string
|
||||
final_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_mode?: "global" | "per_video"
|
||||
bgm_preference?: string
|
||||
intent_result?: IntentResult
|
||||
intent_text?: string
|
||||
progress_stage?: ViralVideoStage
|
||||
progress_percent?: number
|
||||
progress_message?: string
|
||||
output_url?: string
|
||||
error_message?: string
|
||||
credits_cost?: number
|
||||
created_at?: string
|
||||
updated_at?: string
|
||||
}
|
||||
|
||||
export interface GenerateViralVideoRequest {
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
bgm_preference?: string
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
language?: string
|
||||
persona_id?: string
|
||||
viral_structure?: string
|
||||
marketing_purpose?: string
|
||||
duration?: number
|
||||
video_model?: string
|
||||
video_ratio?: string
|
||||
}
|
||||
|
||||
export interface HistoryResponse {
|
||||
items: ViralVideoJob[]
|
||||
total: number
|
||||
page: number
|
||||
page_size: number
|
||||
}
|
||||
@@ -2,6 +2,7 @@
|
||||
export const ROUTE_TITLE_MAP: Record<string, string> = {
|
||||
"/app/dashboard": "首页",
|
||||
"/app/generate": "智能剪辑",
|
||||
"/app/viral-video": "爆款视频",
|
||||
"/app/assets": "视频库",
|
||||
"/app/voices": "配音库",
|
||||
"/app/products": "成片库",
|
||||
|
||||
@@ -18,6 +18,7 @@ import {
|
||||
ThunderboltOutlined,
|
||||
UnorderedListOutlined,
|
||||
UserOutlined,
|
||||
FireOutlined,
|
||||
} from "@ant-design/icons"
|
||||
|
||||
/** 导航项类型 */
|
||||
@@ -76,6 +77,12 @@ export const NAV_ITEMS: NavItem[] = [
|
||||
path: "/app/ai-avatar",
|
||||
icon: React.createElement(UserOutlined),
|
||||
},
|
||||
{
|
||||
key: "viral-video",
|
||||
label: "爆款视频",
|
||||
path: "/app/viral-video",
|
||||
icon: React.createElement(FireOutlined),
|
||||
},
|
||||
{
|
||||
key: "history",
|
||||
label: "任务历史",
|
||||
@@ -142,6 +149,12 @@ export const NAV_GROUPS: NavGroup[] = [
|
||||
path: "/app/ai-avatar",
|
||||
icon: React.createElement(UserOutlined),
|
||||
},
|
||||
{
|
||||
key: "viral-video",
|
||||
label: "爆款视频",
|
||||
path: "/app/viral-video",
|
||||
icon: React.createElement(FireOutlined),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
||||
@@ -11,7 +11,12 @@
|
||||
* 防止长标题在窄列里溢出导致与相邻卡片进度条视觉重叠。
|
||||
*/
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import {
|
||||
LoadingOutlined,
|
||||
CheckCircleFilled,
|
||||
CloseCircleOutlined,
|
||||
ClockCircleOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
@@ -61,6 +66,11 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#ef4444" }}
|
||||
/>
|
||||
) : task.status === "queued" ? (
|
||||
<ClockCircleOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#faad14" }}
|
||||
/>
|
||||
) : (
|
||||
<LoadingOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
@@ -85,6 +95,21 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
<div className="xx-batch-gen-card-pct">{Math.round(task.progress)}%</div>
|
||||
</>
|
||||
)}
|
||||
{task.status === "queued" && (
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 8,
|
||||
color: "var(--text-secondary, #faad14)",
|
||||
fontSize: 13,
|
||||
padding: "8px 0",
|
||||
}}
|
||||
>
|
||||
<ClockCircleOutlined />
|
||||
<span>排队等待中,前面任务完成后自动开始渲染</span>
|
||||
</div>
|
||||
)}
|
||||
{(task.status === "completed" || task.status === "awaiting_cover") && video && (
|
||||
// 竖屏自适应容器(#1750):成片固定 1080×1920(9:16),
|
||||
// 视频按真实宽高比 contain 显示,黑底居中,杜绝横屏播放器左右大黑边
|
||||
|
||||
@@ -12,6 +12,7 @@ import Step2MaterialSelect from "../components/Step2MaterialSelect"
|
||||
import Step4TitleSettings from "../components/Step4TitleSettings"
|
||||
import Step6CoverSettings from "../components/Step6CoverSettings"
|
||||
import BatchGenerationGrid from "./BatchGenerationGrid"
|
||||
import Step3VoiceWithMode from "./Step3VoiceWithMode"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
import type { TitleTemplate } from "@/components/title/template-types"
|
||||
@@ -151,6 +152,12 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedCoverTemplate,
|
||||
onSelectedCoverTemplateChange,
|
||||
onConfirmGenerate,
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
voiceModePerVideo,
|
||||
onVoiceModePerVideoChange,
|
||||
voiceLibraryIds,
|
||||
onVoiceLibraryIdsChange,
|
||||
} = props
|
||||
|
||||
switch (currentStep) {
|
||||
@@ -184,32 +191,46 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
)
|
||||
case 3:
|
||||
return (
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1 ? props.selectedVariantIds?.length || 1 : 1
|
||||
}
|
||||
/>
|
||||
<>
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1
|
||||
? props.selectedVariantIds?.length || 1
|
||||
: 1
|
||||
}
|
||||
/>
|
||||
{/* 批量配音选择:共用/独立切换(#2096) */}
|
||||
<Step3VoiceWithMode
|
||||
previewCount={previewCount}
|
||||
selectedVoice={selectedVoice}
|
||||
onSelectedVoiceChange={onSelectedVoiceChange}
|
||||
voiceModePerVideo={voiceModePerVideo}
|
||||
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
|
||||
voiceLibraryIds={voiceLibraryIds}
|
||||
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
|
||||
/>
|
||||
</>
|
||||
)
|
||||
case 4:
|
||||
return (
|
||||
|
||||
@@ -157,7 +157,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
/>
|
||||
</div>
|
||||
) : (
|
||||
/* ── 批量:N 个独立标题输入框 ── */
|
||||
/* ── 批量:N 个独立标题输入框(两列布局 #2096) ── */
|
||||
<div className="xx-batch-titles">
|
||||
<div
|
||||
style={{
|
||||
@@ -170,17 +170,25 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
为每个视频输入独立标题。标题样式(字体/颜色/位置)全局统一。
|
||||
</div>
|
||||
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: 640 }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(2, minmax(0, 1fr))",
|
||||
gap: 16,
|
||||
}}
|
||||
>
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: "100%" }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ export interface BatchTaskState {
|
||||
taskId: string
|
||||
/** 变体序号(0-based,与标题/封面数组对齐) */
|
||||
variantIndex: number
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed"
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed" | "queued"
|
||||
progress: number
|
||||
error: string | null
|
||||
/** 完成后的成片视频 */
|
||||
@@ -374,5 +374,42 @@ export function useGenerationPolling({
|
||||
}
|
||||
}, [])
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
/**
|
||||
* 批量队列模式:逐任务追加到轮询队列(支持串行提交、429 排队重试场景)。
|
||||
* 与 startPollingBatch 不同的是:
|
||||
* - 不会 reset batchContextRef;多次调用会累积
|
||||
* - 不触发整体 onComplete / onFailed(完成判定交给外层 useEffect 按状态聚合)
|
||||
* - 仍通过 onBatchTaskUpdate 回传单任务状态
|
||||
*/
|
||||
const pollBatchTaskQueued = useCallback(
|
||||
(taskId: string, variantIndex: number) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.set(taskId, variantIndex)
|
||||
onBatchTaskUpdate?.(taskId, {
|
||||
taskId,
|
||||
variantIndex,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
})
|
||||
pollSingleTask(taskId, Date.now(), {
|
||||
onTaskProgress: (pct) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: pct })
|
||||
},
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
const finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
onBatchTaskUpdate?.(taskId, { status: finalStatus, progress: 100, videos })
|
||||
},
|
||||
onTaskFailed: (msg) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "failed", error: msg })
|
||||
},
|
||||
}).catch(() => {
|
||||
/* onTaskFailed 已处理 */
|
||||
})
|
||||
},
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, pollBatchTaskQueued, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
* 视频生成 Hook
|
||||
* 封装视频生成的核心逻辑、状态管理、轮询等
|
||||
*/
|
||||
import { useState, useCallback, useEffect } from "react"
|
||||
import { useState, useCallback, useEffect, useRef } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor"
|
||||
import { createGenerationTask } from "@/api/tasks/tasks"
|
||||
import type { CreateGenerationTaskRequest } from "@/api/tasks/types"
|
||||
import type { UseGenerateVideoProps } from "./generate-video/types"
|
||||
import { getGenerationPhase } from "./generate-video/phase"
|
||||
import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling"
|
||||
@@ -15,6 +17,26 @@ import { extractBackendError, translateError } from "./generate-video/errorUtils
|
||||
|
||||
export type GenerationCompleteStatus = "completed" | "awaiting_cover" | null
|
||||
|
||||
/** 判断是否是用户队列已满 429(需要排队重试而非直接报错) */
|
||||
function isUserQueueFullError(err: unknown): { waitMs: number } | null {
|
||||
if (!axios.isAxiosError(err)) return null
|
||||
if (err.response?.status !== 429 && err.response?.status !== 503) return null
|
||||
const detail = (err.response?.data as { detail?: unknown })?.detail
|
||||
const code =
|
||||
typeof detail === "object" && detail !== null ? (detail as { code?: string }).code : undefined
|
||||
if (code === "USER_QUEUE_FULL" || code === "SYSTEM_QUEUE_FULL") {
|
||||
const waitSec =
|
||||
typeof detail === "object" && detail !== null
|
||||
? Number((detail as { estimated_wait_seconds?: number }).estimated_wait_seconds) || 0
|
||||
: 0
|
||||
return { waitMs: Math.max(15_000, waitSec * 1000 || 30_000) }
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/** sleep */
|
||||
const sleep = (ms: number) => new Promise<void>((r) => setTimeout(r, ms))
|
||||
|
||||
export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { selectedTemplate, onGenerationSuccess } = props
|
||||
|
||||
@@ -31,6 +53,22 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
|
||||
|
||||
/** 排队中重试的定时器,unmount / 新提交时清理 */
|
||||
const queueTimersRef = useRef<number[]>([])
|
||||
const cancelledRef = useRef(false)
|
||||
|
||||
const clearQueueTimers = useCallback(() => {
|
||||
queueTimersRef.current.forEach((id) => clearTimeout(id))
|
||||
queueTimersRef.current = []
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
cancelledRef.current = true
|
||||
clearQueueTimers()
|
||||
}
|
||||
}, [clearQueueTimers])
|
||||
|
||||
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
|
||||
setBatchTasks((prev) => {
|
||||
const list = prev || []
|
||||
@@ -58,7 +96,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
(videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => {
|
||||
setGenerating(false)
|
||||
setGenerated(true)
|
||||
const finalStatus: GenerationCompleteStatus = taskStatus ?? "completed"
|
||||
setCompletionStatus(finalStatus)
|
||||
@@ -81,21 +118,30 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
[onGenerationSuccess],
|
||||
)
|
||||
const handleFailed = useCallback((errorMsg: string) => {
|
||||
setGenerating(false)
|
||||
setGenerateError(errorMsg)
|
||||
}, [])
|
||||
|
||||
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频。
|
||||
当全部任务都已结束(completed/awaiting_cover/failed)且无排队/渲染中任务时,关闭 generating。 */
|
||||
useEffect(() => {
|
||||
if (batchTasks.length === 0) return
|
||||
const byVariant = new Map<number, GeneratedVideo>()
|
||||
let hasQueued = false
|
||||
let hasRunning = false
|
||||
let hasSuccess = false
|
||||
let allDone = true
|
||||
batchTasks.forEach((t) => {
|
||||
if (
|
||||
t.status === "completed" ||
|
||||
(t.status === "awaiting_cover" && t.videos && t.videos.length > 0)
|
||||
) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
if (t.status === "queued") hasQueued = true
|
||||
else if (t.status === "running") hasRunning = true
|
||||
if (t.status === "completed" || t.status === "awaiting_cover") {
|
||||
hasSuccess = true
|
||||
if (t.videos && t.videos.length > 0) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
}
|
||||
}
|
||||
if (t.status !== "completed" && t.status !== "awaiting_cover" && t.status !== "failed") {
|
||||
allDone = false
|
||||
}
|
||||
})
|
||||
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
|
||||
@@ -105,17 +151,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
return ordered
|
||||
})
|
||||
if (allDone && !hasQueued && !hasRunning) {
|
||||
setGenerating(false)
|
||||
if (hasSuccess) {
|
||||
setGenerated(true)
|
||||
setCompletionStatus("awaiting_cover")
|
||||
}
|
||||
}
|
||||
}, [batchTasks])
|
||||
|
||||
const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({
|
||||
const { startPolling, pollBatchTaskQueued, retryTask, clearTimer } = useGenerationPolling({
|
||||
onProgress: handleProgress,
|
||||
onComplete: handleComplete,
|
||||
onFailed: handleFailed,
|
||||
onBatchTaskUpdate: handleBatchTaskUpdate,
|
||||
})
|
||||
|
||||
/** 根据 props 构造基础 payload(批量/单任务共用的字段) */
|
||||
const buildBasePayload = useCallback((): Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> => {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
const bgmConfig = {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
}
|
||||
|
||||
const titleConfig = props.titleSettings?.title
|
||||
? {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color: props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
}
|
||||
: undefined
|
||||
|
||||
const payload: Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> = {
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: bgmConfig as CreateGenerationTaskRequest["bgm_config"],
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(titleConfig ? ({ title_config: titleConfig } as Record<string, unknown>) : {}),
|
||||
}
|
||||
|
||||
return payload
|
||||
}, [props, selectedTemplate])
|
||||
|
||||
/* ── 生成视频 ──
|
||||
返回 true 表示任务创建成功并已开始轮询;false 表示校验未通过或创建失败 */
|
||||
返回 true 表示任务创建成功并已开始轮询(含排队中);false 表示校验未通过或创建失败 */
|
||||
const generate = useCallback(async (): Promise<boolean> => {
|
||||
const errorMsg = validateGenerateInputs(props)
|
||||
if (errorMsg) {
|
||||
@@ -123,6 +337,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
return false
|
||||
}
|
||||
|
||||
cancelledRef.current = false
|
||||
clearQueueTimers()
|
||||
setGenerating(true)
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
@@ -133,25 +349,16 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setCurrentTaskId("")
|
||||
clearTimer()
|
||||
|
||||
const basePayload = buildBasePayload()
|
||||
const assetIds = basePayload.asset_ids
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
|
||||
try {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
|
||||
// from-assets 已由 useStep2Materials 在用户选素材时(debounce 800ms)调用,
|
||||
// 后端已改为异步秒级返回,这里做一次轻量兜底:
|
||||
// 单次查 clips,已有则直接放行;没有则再调一次 from-assets。
|
||||
// from-assets 兜底:片段不存在则补一次
|
||||
if (assetIds.length > 0 && selectedTemplate) {
|
||||
try {
|
||||
const clipList = await getEditPlanClips(selectedTemplate, { limit: 500 })
|
||||
if (clipList.items.length === 0) {
|
||||
// 片段不存在(极端情况:useStep2Materials 的 debounce 还没触发)
|
||||
// 手动补一次 from-assets(后端秒级返回)
|
||||
await createClipsFromAssets(selectedTemplate, assetIds, "main")
|
||||
}
|
||||
} catch {
|
||||
@@ -159,221 +366,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
}
|
||||
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
const hide = message.loading(
|
||||
isBatch ? `正在生成 ${props.previewCount} 个视频...` : "正在生成预览视频...",
|
||||
0,
|
||||
)
|
||||
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
|
||||
const indexes =
|
||||
isBatch && props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = isBatch ? indexes.length : 1
|
||||
|
||||
// 标题文字数组:批量时按勾选顺序
|
||||
const titlesArr =
|
||||
isBatch && (props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: []
|
||||
// 配音数组:独立配音模式按勾选顺序;否则不传(回退共用 voice_library_id)
|
||||
const voiceArr =
|
||||
isBatch && props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map((i) => props.variantVoiceLibraryIds![i] || voiceLibraryId)
|
||||
: []
|
||||
// 封面数组:批量时按勾选顺序(未设置封面的变体传空串,后端回退智能封面)
|
||||
const coversArr =
|
||||
isBatch && props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: []
|
||||
// #1744 变体 plan 数组:预览阶段后端独立选片产出的 plan id,按勾选顺序回传,
|
||||
// 后端直接关联这些 plan 渲染(不再重新选片)→ 预览所见即成片。
|
||||
// 全部为空(降级本地模拟/后端端点未上线)时不传,后端走自身独立选片。
|
||||
const variantPlansArr =
|
||||
isBatch && props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: []
|
||||
const hasVariantPlans = variantPlansArr.some((id) => !!id)
|
||||
|
||||
try {
|
||||
const taskResp = await createGenerationTask({
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
},
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(isBatch ? { count: batchCount } : {}),
|
||||
...(titlesArr.length ? { titles: titlesArr } : {}),
|
||||
...(voiceArr.length ? { voice_library_ids: voiceArr } : {}),
|
||||
...(coversArr.length ? { cover_urls: coversArr } : {}),
|
||||
...(hasVariantPlans ? { variant_plan_ids: variantPlansArr } : {}),
|
||||
...(props.titleSettings?.title
|
||||
? {
|
||||
title_config: {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color:
|
||||
props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
})
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
if (taskIds.length > 1) {
|
||||
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
|
||||
setCurrentTaskId("")
|
||||
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
|
||||
} else {
|
||||
if (!isBatch) {
|
||||
/* ── 单视频:原逻辑(一次提交 count=1) ── */
|
||||
const hide = message.loading("正在生成预览视频...", 0)
|
||||
try {
|
||||
const taskResp = await createGenerationTask({ ...basePayload, count: 1 })
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
setCurrentTaskId(taskIds[0])
|
||||
startPolling(taskIds[0])
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
}
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
return true
|
||||
}
|
||||
|
||||
/* ── 批量:支持任意数量视频,按队列容量串行提交,429 自动排队重试 ── */
|
||||
const indexes = props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = indexes.length
|
||||
|
||||
const titlesAll =
|
||||
(props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: indexes.map(() => props.titleSettings?.title || "")
|
||||
const voiceArrAll =
|
||||
props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map(
|
||||
(i) => props.variantVoiceLibraryIds![i] || basePayload.voice_library_id || "",
|
||||
)
|
||||
: []
|
||||
const coversAll = props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: indexes.map(() => "")
|
||||
const plansAll = props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: indexes.map(() => "")
|
||||
|
||||
const hasAnyVoice = voiceArrAll.some((v) => !!v)
|
||||
const hasAnyCover = coversAll.some((u) => !!u)
|
||||
const hasAnyPlan = plansAll.some((id) => !!id)
|
||||
|
||||
// 先用占位 ID 把所有变体卡片置为 queued,UI 可见
|
||||
const placeholderIds = indexes.map((_, i) => `__queued_${Date.now()}_${i}`)
|
||||
const initialTasks: BatchTaskState[] = indexes.map((variantIndex, i) => ({
|
||||
taskId: placeholderIds[i],
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
}))
|
||||
setBatchTasks(initialTasks)
|
||||
|
||||
message.loading({
|
||||
content: `已提交 ${batchCount} 个视频任务,系统按队列容量依次渲染…`,
|
||||
key: "batch-gen",
|
||||
duration: 3,
|
||||
})
|
||||
|
||||
/** 将占位 taskId 更新为真实 taskId(卡片引用同一对象) */
|
||||
const replacePlaceholder = (placeholderId: string, realTaskId: string) => {
|
||||
setBatchTasks((prev) => {
|
||||
const idx = prev.findIndex((t) => t.taskId === placeholderId)
|
||||
if (idx === -1) return prev
|
||||
const next = [...prev]
|
||||
next[idx] = { ...next[idx], taskId: realTaskId }
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
/** 提交某一索引的单任务(count=1),成功后返回真实 taskId;429/503 则返回 waitMs */
|
||||
const submitOne = async (
|
||||
i: number,
|
||||
): Promise<{ queued: true; waitMs: number } | { queued: false; taskId: string }> => {
|
||||
const body: CreateGenerationTaskRequest = {
|
||||
...basePayload,
|
||||
count: 1,
|
||||
titles: [titlesAll[i] || ""],
|
||||
...(hasAnyVoice
|
||||
? { voice_library_ids: [voiceArrAll[i] || basePayload.voice_library_id || ""] }
|
||||
: {}),
|
||||
...(hasAnyCover ? { cover_urls: [coversAll[i] || ""] } : {}),
|
||||
...(hasAnyPlan && plansAll[i] ? { variant_plan_ids: [plansAll[i]] } : {}),
|
||||
}
|
||||
try {
|
||||
const resp = await createGenerationTask(body)
|
||||
const item = resp.items?.[0]
|
||||
const tid = item?.id
|
||||
if (!tid) throw new Error("创建任务成功但未返回任务 ID")
|
||||
return { queued: false, taskId: tid }
|
||||
} catch (err) {
|
||||
const q = isUserQueueFullError(err)
|
||||
if (q) return { queued: true, waitMs: q.waitMs }
|
||||
throw err
|
||||
}
|
||||
}
|
||||
|
||||
// 串行提交:每次提交一个;429/503 则等待后重试;其它错误立即标记该任务失败
|
||||
let fatalErr: unknown = null
|
||||
for (let i = 0; i < batchCount; i++) {
|
||||
if (cancelledRef.current) return false
|
||||
const variantIndex = indexes[i]
|
||||
const placeholderId = placeholderIds[i]
|
||||
let attempt = 0
|
||||
let submitted = false
|
||||
while (!submitted) {
|
||||
if (cancelledRef.current) return false
|
||||
attempt++
|
||||
try {
|
||||
const result = await submitOne(i)
|
||||
if (!result.queued) {
|
||||
replacePlaceholder(placeholderId, result.taskId)
|
||||
// 先更新到 running,再启动单任务增量轮询(不触发整体 onComplete)
|
||||
pollBatchTaskQueued(result.taskId, variantIndex)
|
||||
submitted = true
|
||||
} else {
|
||||
// 排队:保持 queued 状态,等待后重试
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
})
|
||||
if (attempt === 1) {
|
||||
message.info({
|
||||
content: `队列繁忙,${Math.round(result.waitMs / 1000)} 秒后自动继续提交后续视频…`,
|
||||
key: "batch-gen",
|
||||
duration: 4,
|
||||
})
|
||||
}
|
||||
await sleep(Math.min(result.waitMs, 60_000))
|
||||
}
|
||||
} catch (err) {
|
||||
// 非限流错误:该任务标记失败,继续后续任务(不阻断整个批量)
|
||||
console.error("[batch generate] 任务提交失败:", err)
|
||||
const msg = translateError(extractBackendError(err))
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "failed",
|
||||
error: msg,
|
||||
progress: 0,
|
||||
})
|
||||
submitted = true
|
||||
if (!fatalErr) fatalErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (fatalErr) {
|
||||
// 有任务失败但其余已成功,整体不 throw;由 UI 展示单个失败卡片
|
||||
}
|
||||
return true
|
||||
} catch (err: unknown) {
|
||||
console.error("[handleGenerate] 生成失败:", err)
|
||||
setGenerating(false)
|
||||
const backendMsg = extractBackendError(err)
|
||||
console.error("[handleGenerate] 错误信息:", backendMsg, "完整错误:", err)
|
||||
const finalMsg = translateError(backendMsg)
|
||||
setGenerateError(finalMsg)
|
||||
setGenerating(false)
|
||||
message.error(finalMsg)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
|
||||
}, [
|
||||
props,
|
||||
clearTimer,
|
||||
startPolling,
|
||||
selectedTemplate,
|
||||
buildBasePayload,
|
||||
handleBatchTaskUpdate,
|
||||
clearQueueTimers,
|
||||
pollBatchTaskQueued,
|
||||
])
|
||||
|
||||
const retry = useCallback(() => {
|
||||
setGenerateError(null)
|
||||
@@ -383,9 +554,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 第5步:单独重试某个失败任务 */
|
||||
const retryBatchTask = useCallback(
|
||||
(taskId: string) => {
|
||||
handleBatchTaskUpdate(taskId, { status: "running", progress: 0, error: null, videos: [] })
|
||||
retryTask(taskId)
|
||||
},
|
||||
[retryTask],
|
||||
[retryTask, handleBatchTaskUpdate],
|
||||
)
|
||||
|
||||
const dismissError = useCallback(() => {
|
||||
|
||||
@@ -0,0 +1,855 @@
|
||||
/* ============================================================
|
||||
爆款视频创作页 - 浅色紫调(对齐 AI 数字人页视觉规范)
|
||||
布局(参考 ui-ref-step-layout.png 三列等宽 STEP 向导):
|
||||
.vv-tabs 顶栏多任务 Tab(生成1 × / + 新建)
|
||||
.vv-grid 三列等宽 grid(1fr 1fr 1fr,gap 16)
|
||||
├── .vv-col 左:STEP 1 上传素材(图片+参考视频+配音)
|
||||
├── .vv-col 中:STEP 2 生成视频文案(融合Tab+参数+文案I/O+AI摘要)
|
||||
└── .vv-col 右:STEP 3 生成视频(预览+参数+进度+扣点+按钮)
|
||||
可折叠模块:.vv-section > .vv-section-head[aria-expanded] + .vv-section-body
|
||||
============================================================ */
|
||||
|
||||
.vv-page {
|
||||
padding: 16px;
|
||||
background: #f5f6fa;
|
||||
min-height: calc(100vh - 56px);
|
||||
}
|
||||
|
||||
/* ── 顶部任务 Tab 栏 ─────────────────────────────────── */
|
||||
.vv-tabs {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
margin-bottom: 14px;
|
||||
padding: 6px 8px;
|
||||
background: #fff;
|
||||
border-radius: 10px;
|
||||
border: 1px solid #e5e7eb;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
.vv-tab {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
padding: 6px 12px;
|
||||
border-radius: 6px;
|
||||
font-size: 13px;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
border: 1px solid transparent;
|
||||
background: transparent;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-tab:hover {
|
||||
background: #f3f4f6;
|
||||
color: #374151;
|
||||
}
|
||||
.vv-tab.active {
|
||||
background: #f3f0ff;
|
||||
color: #7c3aed;
|
||||
border-color: #d8cafc;
|
||||
font-weight: 500;
|
||||
}
|
||||
.vv-tab .vv-tab-close {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-tab .vv-tab-close:hover {
|
||||
background: rgba(0, 0, 0, 0.08);
|
||||
color: #374151;
|
||||
}
|
||||
.vv-tab-new {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
padding: 6px 10px;
|
||||
border-radius: 6px;
|
||||
font-size: 13px;
|
||||
color: #7c3aed;
|
||||
cursor: pointer;
|
||||
border: 1px dashed #d8cafc;
|
||||
background: transparent;
|
||||
}
|
||||
.vv-tab-new:hover {
|
||||
background: #f3f0ff;
|
||||
}
|
||||
.vv-tabs-right {
|
||||
margin-left: auto;
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
/* ── 三列等宽网格 ───────────────────────────────────── */
|
||||
.vv-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(0, 1fr));
|
||||
gap: 16px;
|
||||
align-items: start;
|
||||
}
|
||||
@media (max-width: 1280px) {
|
||||
.vv-grid {
|
||||
grid-template-columns: repeat(2, minmax(0, 1fr));
|
||||
}
|
||||
}
|
||||
@media (max-width: 900px) {
|
||||
.vv-grid {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
|
||||
.vv-col {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 14px;
|
||||
}
|
||||
|
||||
/* ── 可折叠模块(对齐 AI 数字人卡块) ────────────────── */
|
||||
.vv-section {
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 12px;
|
||||
overflow: hidden;
|
||||
}
|
||||
.vv-section-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 16px;
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
border-bottom: 1px solid #f3f4f6;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-head {
|
||||
border-bottom: none;
|
||||
}
|
||||
.vv-section-title {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
color: #111827;
|
||||
}
|
||||
.vv-section-title .vv-step-badge {
|
||||
width: 22px;
|
||||
height: 22px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
color: #fff;
|
||||
font-size: 12px;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-section-arrow {
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-arrow {
|
||||
transform: rotate(-90deg);
|
||||
}
|
||||
.vv-section-body {
|
||||
padding: 14px 16px 16px;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-body {
|
||||
display: none;
|
||||
}
|
||||
|
||||
/* ── 通用表单元素(对齐 AI 数字人样式) ─────────────── */
|
||||
.vv-label {
|
||||
display: block;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.vv-input,
|
||||
.vv-select,
|
||||
.vv-textarea {
|
||||
width: 100%;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 8px;
|
||||
padding: 9px 12px;
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
background: #fff;
|
||||
outline: none;
|
||||
transition:
|
||||
border-color 0.15s,
|
||||
box-shadow 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-textarea {
|
||||
line-height: 1.6;
|
||||
resize: vertical;
|
||||
min-height: 90px;
|
||||
}
|
||||
.vv-input:focus,
|
||||
.vv-select:focus,
|
||||
.vv-textarea:focus {
|
||||
border-color: #7c3aed;
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.1);
|
||||
}
|
||||
.vv-input::placeholder,
|
||||
.vv-textarea::placeholder {
|
||||
color: #d1d5db;
|
||||
}
|
||||
.vv-form-row {
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.vv-form-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, minmax(0, 1fr));
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
/* ── Tab 分段(对齐"系统预设/我的音色"样式) ────────── */
|
||||
.vv-seg-tabs {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
margin-bottom: 12px;
|
||||
border-radius: 8px;
|
||||
padding: 3px;
|
||||
background: #f5f6fa;
|
||||
}
|
||||
.vv-seg-tab {
|
||||
flex: 1;
|
||||
padding: 7px 10px;
|
||||
font-size: 13px;
|
||||
text-align: center;
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
color: #6b7280;
|
||||
background: transparent;
|
||||
border: 1px solid transparent;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-seg-tab:hover {
|
||||
color: #374151;
|
||||
}
|
||||
.vv-seg-tab.active {
|
||||
background: #fff;
|
||||
color: #7c3aed;
|
||||
border-color: #7c3aed;
|
||||
font-weight: 500;
|
||||
box-shadow: 0 1px 2px rgba(124, 58, 237, 0.06);
|
||||
}
|
||||
|
||||
/* 融合强度大按钮(选中紫色描边+浅紫底) */
|
||||
.vv-fusion-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(0, 1fr));
|
||||
gap: 8px;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.vv-fusion-btn {
|
||||
padding: 10px 8px;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
border-radius: 8px;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
color: #6b7280;
|
||||
text-align: center;
|
||||
line-height: 1.4;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-fusion-btn strong {
|
||||
display: block;
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
margin-bottom: 2px;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-fusion-btn:hover {
|
||||
border-color: #d8cafc;
|
||||
}
|
||||
.vv-fusion-btn.active {
|
||||
background: #f3f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-fusion-btn.active strong {
|
||||
color: #7c3aed;
|
||||
}
|
||||
|
||||
/* 风格强度小分段按钮(三档,参考配音风格按钮) */
|
||||
.vv-pill-row {
|
||||
display: flex;
|
||||
gap: 6px;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
.vv-pill {
|
||||
padding: 6px 12px;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
border-radius: 6px;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
font-family: inherit;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-pill:hover {
|
||||
border-color: #d8cafc;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-pill.active {
|
||||
background: #f3f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
/* ── 上传区(浅灰虚线框) ───────────────────────────── */
|
||||
.vv-upload {
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 10px;
|
||||
padding: 20px;
|
||||
text-align: center;
|
||||
cursor: pointer;
|
||||
transition:
|
||||
border-color 0.2s,
|
||||
background 0.2s;
|
||||
background: #fafbfc;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-upload:hover,
|
||||
.vv-upload.dragover {
|
||||
border-color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-upload-icon {
|
||||
font-size: 28px;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.vv-upload small {
|
||||
display: block;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
/* 图片网格 */
|
||||
.vv-img-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(90px, 1fr));
|
||||
gap: 8px;
|
||||
margin-top: 12px;
|
||||
}
|
||||
.vv-img-item {
|
||||
position: relative;
|
||||
aspect-ratio: 1;
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
border: 1px solid #e5e7eb;
|
||||
cursor: grab;
|
||||
background: #f5f6fa;
|
||||
}
|
||||
.vv-img-item.dragging {
|
||||
opacity: 0.4;
|
||||
}
|
||||
.vv-img-item img {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: cover;
|
||||
}
|
||||
.vv-img-badge {
|
||||
position: absolute;
|
||||
top: 4px;
|
||||
left: 4px;
|
||||
background: rgba(124, 58, 237, 0.9);
|
||||
color: #fff;
|
||||
font-size: 11px;
|
||||
font-weight: 600;
|
||||
padding: 1px 6px;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-img-del {
|
||||
position: absolute;
|
||||
top: 4px;
|
||||
right: 4px;
|
||||
width: 20px;
|
||||
height: 20px;
|
||||
border-radius: 50%;
|
||||
background: rgba(239, 68, 68, 0.9);
|
||||
color: #fff;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
font-size: 12px;
|
||||
line-height: 1;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
opacity: 0;
|
||||
transition: opacity 0.15s;
|
||||
}
|
||||
.vv-img-item:hover .vv-img-del {
|
||||
opacity: 1;
|
||||
}
|
||||
.vv-img-add {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 8px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #9ca3af;
|
||||
cursor: pointer;
|
||||
background: #fafbfc;
|
||||
font-size: 11px;
|
||||
gap: 2px;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-img-add:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
}
|
||||
.vv-progress-mini {
|
||||
position: absolute;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.35);
|
||||
color: #fff;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 11px;
|
||||
}
|
||||
|
||||
/* 参考视频预览 */
|
||||
.vv-video-preview {
|
||||
width: 100%;
|
||||
aspect-ratio: 16/9;
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
background: #000;
|
||||
margin-top: 10px;
|
||||
border: 1px solid #e5e7eb;
|
||||
}
|
||||
.vv-video-preview video {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: cover;
|
||||
}
|
||||
.vv-video-ph {
|
||||
width: 100%;
|
||||
aspect-ratio: 16/9;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 8px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #9ca3af;
|
||||
cursor: pointer;
|
||||
background: #fafbfc;
|
||||
font-size: 12px;
|
||||
gap: 4px;
|
||||
margin-top: 10px;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-video-ph:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
}
|
||||
|
||||
/* ── 音色列表(参考 AI 数字人「龙小淳」卡片) ──────── */
|
||||
.vv-voice-tabs {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-voice-list {
|
||||
max-height: 280px;
|
||||
overflow-y: auto;
|
||||
border: 1px solid #f3f4f6;
|
||||
border-radius: 8px;
|
||||
}
|
||||
.vv-voice-list::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
.vv-voice-list::-webkit-scrollbar-thumb {
|
||||
background: #e5e7eb;
|
||||
border-radius: 3px;
|
||||
}
|
||||
.vv-voice-item {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 10px 12px;
|
||||
border-bottom: 1px solid #f3f4f6;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
.vv-voice-item:last-child {
|
||||
border-bottom: none;
|
||||
}
|
||||
.vv-voice-item:hover {
|
||||
background: #f9fafb;
|
||||
}
|
||||
.vv-voice-item.selected {
|
||||
background: #f3f0ff;
|
||||
}
|
||||
.vv-voice-radio {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
border: 2px solid #d1d5db;
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
.vv-voice-item.selected .vv-voice-radio {
|
||||
border-color: #7c3aed;
|
||||
}
|
||||
.vv-voice-item.selected .vv-voice-radio::after {
|
||||
content: "";
|
||||
width: 8px;
|
||||
height: 8px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
}
|
||||
.vv-voice-info {
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
}
|
||||
.vv-voice-name {
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
font-weight: 500;
|
||||
}
|
||||
.vv-voice-desc {
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 2px;
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
}
|
||||
.vv-voice-play {
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border-radius: 50%;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 12px;
|
||||
flex-shrink: 0;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-voice-play:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-voice-play.playing {
|
||||
background: #7c3aed;
|
||||
border-color: #7c3aed;
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
/* ── AI 摘要确认卡(黄色高亮) ─────────────────────── */
|
||||
.vv-intent {
|
||||
background: #fffbeb;
|
||||
border: 1px solid #fcd34d;
|
||||
border-radius: 10px;
|
||||
padding: 14px;
|
||||
margin-top: 12px;
|
||||
}
|
||||
.vv-intent-title {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: #b45309;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-intent-row {
|
||||
margin-bottom: 8px;
|
||||
font-size: 13px;
|
||||
line-height: 1.6;
|
||||
}
|
||||
.vv-intent-row .vv-k {
|
||||
font-size: 12px;
|
||||
color: #92400e;
|
||||
margin-bottom: 3px;
|
||||
}
|
||||
.vv-intent-row .vv-v {
|
||||
color: #111827;
|
||||
}
|
||||
.vv-chip {
|
||||
display: inline-block;
|
||||
padding: 2px 8px;
|
||||
background: #fef3c7;
|
||||
border-radius: 4px;
|
||||
font-size: 11px;
|
||||
color: #92400e;
|
||||
margin-right: 4px;
|
||||
margin-bottom: 3px;
|
||||
}
|
||||
|
||||
/* ── 9:16 预览区 ────────────────────────────────────── */
|
||||
.vv-preview {
|
||||
width: 100%;
|
||||
aspect-ratio: 9/16;
|
||||
max-height: 480px;
|
||||
border-radius: 12px;
|
||||
background: #fff;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
.vv-preview video {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: contain;
|
||||
background: #000;
|
||||
}
|
||||
.vv-preview-placeholder {
|
||||
text-align: center;
|
||||
color: #9ca3af;
|
||||
padding: 20px;
|
||||
}
|
||||
.vv-preview-placeholder .ph-icon {
|
||||
font-size: 40px;
|
||||
opacity: 0.4;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
.vv-preview-placeholder .ph-txt {
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
/* ── 阶段进度列表 ──────────────────────────────────── */
|
||||
.vv-stages {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 4px;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-stage {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
padding: 5px 0;
|
||||
}
|
||||
.vv-stage-dot {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
border: 2px solid #e5e7eb;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 9px;
|
||||
background: #fff;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
.vv-stage.active {
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-stage.active .vv-stage-dot {
|
||||
border-color: #7c3aed;
|
||||
background: #f3f0ff;
|
||||
}
|
||||
.vv-stage.done {
|
||||
color: #10b981;
|
||||
}
|
||||
.vv-stage.done .vv-stage-dot {
|
||||
border-color: #10b981;
|
||||
background: #10b981;
|
||||
color: #fff;
|
||||
}
|
||||
.vv-stage.failed {
|
||||
color: #ef4444;
|
||||
}
|
||||
.vv-stage.failed .vv-stage-dot {
|
||||
border-color: #ef4444;
|
||||
background: #ef4444;
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
.vv-progress-bar {
|
||||
width: 100%;
|
||||
height: 6px;
|
||||
background: #f3f4f6;
|
||||
border-radius: 3px;
|
||||
overflow: hidden;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-progress-fill {
|
||||
height: 100%;
|
||||
background: linear-gradient(90deg, #7c3aed, #a855f7);
|
||||
border-radius: 3px;
|
||||
transition: width 0.4s ease;
|
||||
}
|
||||
.vv-progress-meta {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 6px;
|
||||
}
|
||||
.vv-progress-pct {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
/* ── 扣点 & 按钮 ────────────────────────────────────── */
|
||||
.vv-credits {
|
||||
background: #f9fafb;
|
||||
border: 1px solid #f3f4f6;
|
||||
border-radius: 8px;
|
||||
padding: 10px 12px;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
}
|
||||
.vv-credits strong {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.vv-btn {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 6px;
|
||||
padding: 10px 16px;
|
||||
border-radius: 8px;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
cursor: pointer;
|
||||
border: 1px solid transparent;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-btn-primary {
|
||||
width: 100%;
|
||||
background: linear-gradient(135deg, #7c3aed, #a855f7);
|
||||
color: #fff;
|
||||
padding: 12px;
|
||||
font-size: 14px;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-btn-primary:hover:not(:disabled) {
|
||||
box-shadow: 0 4px 14px rgba(124, 58, 237, 0.3);
|
||||
transform: translateY(-1px);
|
||||
}
|
||||
.vv-btn-primary:disabled {
|
||||
opacity: 0.5;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.vv-btn-ghost {
|
||||
background: #fff;
|
||||
color: #6b7280;
|
||||
border-color: #e5e7eb;
|
||||
}
|
||||
.vv-btn-ghost:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-btn-warn {
|
||||
background: #fef3c7;
|
||||
color: #b45309;
|
||||
border-color: #fcd34d;
|
||||
}
|
||||
.vv-btn-sm {
|
||||
padding: 6px 12px;
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
.vv-error {
|
||||
background: #fef2f2;
|
||||
border: 1px solid #fecaca;
|
||||
border-radius: 8px;
|
||||
padding: 10px 12px;
|
||||
color: #dc2626;
|
||||
font-size: 12px;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-spinner {
|
||||
width: 14px;
|
||||
height: 14px;
|
||||
border: 2px solid rgba(255, 255, 255, 0.3);
|
||||
border-top-color: #fff;
|
||||
border-radius: 50%;
|
||||
animation: vvspin 0.8s linear infinite;
|
||||
display: inline-block;
|
||||
}
|
||||
@keyframes vvspin {
|
||||
to {
|
||||
transform: rotate(360deg);
|
||||
}
|
||||
}
|
||||
|
||||
.vv-muted {
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
}
|
||||
.vv-meta {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 4px;
|
||||
}
|
||||
.vv-file-name {
|
||||
font-size: 12px;
|
||||
color: #374151;
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-link-btn {
|
||||
background: none;
|
||||
border: none;
|
||||
color: #7c3aed;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-link-btn:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
.vv-flex {
|
||||
display: flex;
|
||||
}
|
||||
.vv-between {
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
}
|
||||
.vv-gap-8 {
|
||||
gap: 8px;
|
||||
}
|
||||
.vv-mt-8 {
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-mt-12 {
|
||||
margin-top: 12px;
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,72 @@
|
||||
import { useCallback, useEffect, useRef } from "react"
|
||||
import { getViralVideoJob } from "@/api/viral-video"
|
||||
import type { ViralVideoJob, ViralVideoStatus } from "@/api/viral-video/types"
|
||||
|
||||
const TERMINAL: ViralVideoStatus[] = ["completed", "failed", "cancelled"]
|
||||
|
||||
export interface UseViralVideoPollingOptions {
|
||||
/** 轮询间隔(毫秒),默认 1500 */
|
||||
intervalMs?: number
|
||||
}
|
||||
|
||||
/**
|
||||
* 爆款视频任务 HTTP 轮询 hook(后端暂未暴露 WS 桥,轮询兜底)。
|
||||
* 任务进入终态(completed/failed/cancelled)后自动停止。
|
||||
*/
|
||||
export function useViralVideoPolling(
|
||||
jobId: string | null | undefined,
|
||||
onUpdate: (job: ViralVideoJob) => void,
|
||||
options: UseViralVideoPollingOptions = {},
|
||||
) {
|
||||
const { intervalMs = 1500 } = options
|
||||
const timerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
const stoppedRef = useRef(false)
|
||||
const failCountRef = useRef(0)
|
||||
|
||||
const stop = useCallback(() => {
|
||||
stoppedRef.current = true
|
||||
if (timerRef.current) {
|
||||
clearTimeout(timerRef.current)
|
||||
timerRef.current = null
|
||||
}
|
||||
}, [])
|
||||
|
||||
const pollOnce = useCallback(
|
||||
async (id: string) => {
|
||||
try {
|
||||
const job = await getViralVideoJob(id)
|
||||
failCountRef.current = 0
|
||||
onUpdate(job)
|
||||
if (TERMINAL.includes(job.status)) {
|
||||
stop()
|
||||
return
|
||||
}
|
||||
if (stoppedRef.current) return
|
||||
// 后端在做 GPU 推理/合成阶段拉长间隔
|
||||
const inRender = job.progress_stage === "rendering" || job.progress_stage === "musetalk"
|
||||
const nextDelay = inRender ? 3000 : intervalMs
|
||||
timerRef.current = setTimeout(() => pollOnce(id), nextDelay)
|
||||
} catch (err) {
|
||||
failCountRef.current += 1
|
||||
if (stoppedRef.current) return
|
||||
// 指数退避,最多退到 10s
|
||||
const delay = Math.min(intervalMs * 2 ** Math.min(failCountRef.current, 3), 10000)
|
||||
timerRef.current = setTimeout(() => pollOnce(id), delay)
|
||||
}
|
||||
},
|
||||
[intervalMs, onUpdate, stop],
|
||||
)
|
||||
|
||||
useEffect(() => {
|
||||
stoppedRef.current = false
|
||||
failCountRef.current = 0
|
||||
if (!jobId) {
|
||||
stop()
|
||||
return
|
||||
}
|
||||
pollOnce(jobId)
|
||||
return stop
|
||||
}, [jobId, pollOnce, stop])
|
||||
|
||||
return { stop }
|
||||
}
|
||||
@@ -52,6 +52,10 @@ const appChildren: RouteObject[] = [
|
||||
path: "ai-avatar",
|
||||
lazy: lazyRoute(() => import("@/pages/ai-avatar/AiAvatarPage")),
|
||||
},
|
||||
{
|
||||
path: "viral-video",
|
||||
lazy: lazyRoute(() => import("@/pages/viral-video/ViralVideoPage")),
|
||||
},
|
||||
{
|
||||
path: "voice-clone",
|
||||
lazy: lazyRoute(() => import("@/pages/voice-clone/VoiceClone")),
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""爆款视频 Worker 侧模块(#2039/#2040/#2051)。
|
||||
|
||||
video_analyzer(#2051):参考视频风格分析 6 步管线,输出 style_guide + clips 渲染参数映射。
|
||||
#2040 的 prompt 系统(prompts/prompt_store/llm_runner)由 #2040 分支提供,本文件不依赖它。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import (
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_ANALYSIS_TIMEOUT",
|
||||
"MAX_REFERENCE_DURATION_SEC",
|
||||
"MAX_REFERENCE_SIZE_MB",
|
||||
"STYLE_GUIDE_SCHEMA",
|
||||
"analyze_video_style",
|
||||
"build_render_params_for_clip",
|
||||
"map_bgm_bpm",
|
||||
"map_camera_to_ken_burns",
|
||||
"map_color_to_video_filter",
|
||||
"map_transition_to_xfade",
|
||||
]
|
||||
@@ -0,0 +1,961 @@
|
||||
"""参考爆款视频风格分析模块(#2051,v1.3)。
|
||||
|
||||
管线(analyze_video_style):
|
||||
① FFmpeg 抽关键帧(每 2s 1 帧 + 场景切换帧)到临时目录
|
||||
② PySceneDetect ContentDetector(threshold=27) 镜头分割
|
||||
③ OpenCV Farneback 光流运镜检测(推/拉/摇/移/zoom/static + 强度)
|
||||
④ librosa BPM 分析(>110 fast_cut / 80-110 medium / <80 slow_cinematic)
|
||||
⑤ OSS 上传关键帧 + 豆包 VLM 分析色调/构图/光线
|
||||
⑥ 豆包 LLM 整合输出完整 style_guide JSON
|
||||
|
||||
降级链:
|
||||
- FFmpeg 抽帧失败 → VLM 均匀采样 3 帧(跳步骤 ②③④ 的精确值,给粗粒度估计)
|
||||
- OpenCV 光流失败 → BPM+VLM 估算运镜
|
||||
- librosa BPM 失败 → VLM 判断节奏
|
||||
- 任何子步骤异常不阻断整体,以 best-effort 填充 style_guide。
|
||||
|
||||
资源约束:
|
||||
- 参考视频 ≤60s 且 ≤100MB;分析总超时 ≤60s;临时帧 try/finally 清理。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
import tempfile
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 资源约束 ─────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_REFERENCE_DURATION_SEC = 60
|
||||
MAX_REFERENCE_SIZE_MB = 100
|
||||
DEFAULT_ANALYSIS_TIMEOUT = 60 # 秒
|
||||
KEYFRAME_INTERVAL_SEC = 2
|
||||
SCENEDETECT_THRESHOLD = 27
|
||||
VLM_SAMPLE_FRAMES = 5 # 上传给 VLM 的关键帧上限
|
||||
FARNEBACK_PARAMS = dict(pyr_scale=0.5, levels=3, winsize=15, iterations=3, poly_n=5, poly_sigma=1.2, flags=0)
|
||||
|
||||
# ── style_guide 输出 schema(最小校验参考,不强制 jsonschema 依赖) ───────
|
||||
|
||||
STYLE_GUIDE_SCHEMA: dict[str, Any] = {
|
||||
"style_name": str,
|
||||
"avg_shot_duration": float,
|
||||
"shot_count": int,
|
||||
"pace": str, # fast_cut | medium | slow_cinematic
|
||||
"bpm": int,
|
||||
"camera_movements": list,
|
||||
"transitions": list,
|
||||
"color_palette": list,
|
||||
"color_tone": str, # warm | cool | high_sat | low_sat | vintage | fresh | dramatic | bright
|
||||
"color_filter": str, # none | warm_vintage | cool_fresh | high_contrast | soft_pastel | dramatic_cinematic
|
||||
"composition": dict,
|
||||
"lighting": str,
|
||||
"mood": str,
|
||||
"visual_keywords": list,
|
||||
"ken_burns_params": dict,
|
||||
"transition_map": dict,
|
||||
"video_filter_eq_params": dict,
|
||||
"ken_burns_direction_hint": str,
|
||||
}
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ShotBoundary:
|
||||
"""一段镜头(帧号区间)。"""
|
||||
|
||||
index: int
|
||||
start_sec: float
|
||||
end_sec: float
|
||||
movement: str = (
|
||||
"static" # push_in | pull_out | pan_left | pan_right | tilt_up | tilt_down | static | zoom_in | zoom_out
|
||||
)
|
||||
intensity: str = "low" # low | medium | high
|
||||
transition: str = "hard_cut" # 到下一个镜头的转场
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnalysisArtifacts:
|
||||
"""中间产物(降级路径用)。"""
|
||||
|
||||
frames_dir: Path
|
||||
frame_paths: list[Path] = field(default_factory=list)
|
||||
shots: list[ShotBoundary] = field(default_factory=list)
|
||||
bpm: int = 0
|
||||
vlm_descriptions: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _strip_code_fence(text: str) -> str:
|
||||
"""移除 markdown 代码块围栏,返回纯文本。"""
|
||||
t = text.strip()
|
||||
for fence in ("```json", "```JSON", "```"):
|
||||
if t.startswith(fence):
|
||||
t = t[len(fence) :].lstrip()
|
||||
if t.endswith("```"):
|
||||
t = t[:-3].rstrip()
|
||||
return t
|
||||
|
||||
|
||||
# ── FFmpeg / ffprobe ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ffmpeg_bin() -> str:
|
||||
return shutil.which("ffmpeg") or "ffmpeg"
|
||||
|
||||
|
||||
def _ffprobe_bin() -> str:
|
||||
return shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
def _probe_duration(video_path: str | Path) -> float:
|
||||
"""用 ffprobe 取视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
_ffprobe_bin(),
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(video_path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
text=True,
|
||||
) # nosec B603
|
||||
return float(out.strip() or 0)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("ffprobe 时长探测失败 %s: %s", video_path, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _extract_keyframes(video_path: Path, out_dir: Path, interval: int = KEYFRAME_INTERVAL_SEC) -> list[Path]:
|
||||
"""按固定间隔抽帧;同时检测场景切换帧(select='gt(scene,...)')。"""
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
# 固定间隔
|
||||
fixed_tpl = str(out_dir / "f_%04d.jpg")
|
||||
cmd_fixed = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
f"fps=1/{interval}",
|
||||
"-q:v",
|
||||
"3",
|
||||
fixed_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_fixed, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
# 场景切换帧(独立命名,scene_ 前缀)
|
||||
scene_tpl = str(out_dir / "scene_%04d.jpg")
|
||||
cmd_scene = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
"select='gt(scene,0.35)',showinfo",
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
"3",
|
||||
scene_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_scene, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
frames = sorted(out_dir.glob("f_*.jpg")) + sorted(out_dir.glob("scene_*.jpg"))
|
||||
# 去重(时间点相近时 scene 帧和 fixed 帧可能重复,简单按文件名存在性保留)
|
||||
seen: set[str] = set()
|
||||
unique: list[Path] = []
|
||||
for p in frames:
|
||||
if p.name not in seen:
|
||||
seen.add(p.name)
|
||||
unique.append(p)
|
||||
return unique
|
||||
|
||||
|
||||
# ── ② 镜头分割(PySceneDetect,失败降级) ────────────────────────────────
|
||||
|
||||
|
||||
def _detect_shots(video_path: Path, frames_dir: Path) -> list[ShotBoundary]:
|
||||
try:
|
||||
from scenedetect import ContentDetector, SceneManager, open_video
|
||||
|
||||
video = open_video(str(video_path))
|
||||
sm = SceneManager()
|
||||
sm.add_detector(ContentDetector(threshold=SCENEDETECT_THRESHOLD))
|
||||
sm.detect_scenes(video)
|
||||
scenes = sm.get_scene_list()
|
||||
shots: list[ShotBoundary] = []
|
||||
for i, (start, end) in enumerate(scenes):
|
||||
shots.append(
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=start.get_seconds(),
|
||||
end_sec=end.get_seconds(),
|
||||
)
|
||||
)
|
||||
if shots:
|
||||
return shots
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("PySceneDetect 镜头分割失败,使用均匀分段降级: %s", exc)
|
||||
# 降级:按固定间隔每 3 秒一镜头
|
||||
duration = _probe_duration(video_path) or 15.0
|
||||
dur = max(3.0, min(duration, float(MAX_REFERENCE_DURATION_SEC)))
|
||||
shots = []
|
||||
seg = 3.0
|
||||
i = 0
|
||||
t = 0.0
|
||||
while t < dur - 0.1:
|
||||
shots.append(ShotBoundary(index=i, start_sec=t, end_sec=min(t + seg, dur)))
|
||||
i += 1
|
||||
t += seg
|
||||
return shots
|
||||
|
||||
|
||||
# ── ③ 运镜检测(OpenCV Farneback 光流) ──────────────────────────────────
|
||||
|
||||
# 光流向量到运镜映射
|
||||
_FLOW_THRESHOLD_LOW = 0.3
|
||||
_FLOW_THRESHOLD_HIGH = 1.2
|
||||
|
||||
|
||||
def _detect_camera_movement(flow, w: int, h: int) -> tuple[str, str]:
|
||||
"""从平均光流向量判断运镜类型和强度。"""
|
||||
import numpy as np # noqa: PLC0415 - numpy 已在 requirements 中
|
||||
|
||||
fx = float(np.median(flow[..., 0]))
|
||||
fy = float(np.median(flow[..., 1]))
|
||||
trans_mag = math.hypot(fx, fy)
|
||||
# 发散/收敛判断 zoom:比较边缘流沿径向外指的平均分量(稳健版)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
rx, ry = (xs - cx) / max(cx, 1.0), (ys - cy) / max(cy, 1.0)
|
||||
rmag = np.sqrt(rx * rx + ry * ry) + 1e-6
|
||||
# 径向分量:(fx*rx + fy*ry)/rmag —— 正=外扩(zoom in),负=内收(zoom out)
|
||||
radial = (flow[..., 0] * rx + flow[..., 1] * ry) / rmag
|
||||
# 只看边缘带(|r|>0.5),且减去平移贡献:径向减去平均平移投影
|
||||
edge_mask = (rmag > 0.5).astype(np.float32)
|
||||
if edge_mask.sum() > 10:
|
||||
trans_radial = (fx * rx + fy * ry) / rmag
|
||||
zoom_signal = float(np.mean((radial - trans_radial)[edge_mask > 0]))
|
||||
else:
|
||||
zoom_signal = 0.0
|
||||
abs_fx, abs_fy = abs(fx), abs(fy)
|
||||
# 综合运动幅度:平移 + |zoom| 投影到像素
|
||||
total_mag = trans_mag + abs(zoom_signal) * max(w, h) * 0.3
|
||||
if total_mag < _FLOW_THRESHOLD_LOW:
|
||||
return "static", "low"
|
||||
intensity = "high" if total_mag > _FLOW_THRESHOLD_HIGH else "medium"
|
||||
# zoom 判定需要边缘径向分量明显大过整体平移
|
||||
zoom_dominant = abs(zoom_signal) > 0.6 and abs(zoom_signal) * max(w, h) * 0.3 > trans_mag * 1.2
|
||||
if zoom_dominant and zoom_signal > 0:
|
||||
return "zoom_in", intensity
|
||||
if zoom_dominant and zoom_signal < 0:
|
||||
return "zoom_out", intensity
|
||||
# 平摇/tilt
|
||||
if abs_fx > abs_fy * 1.5:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
if abs_fy > abs_fx * 1.5:
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
# 轨道/跟拍:以主轴为主
|
||||
if abs_fx >= abs_fy:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
|
||||
|
||||
def _analyze_movements(video_path: Path, shots: list[ShotBoundary]) -> None:
|
||||
"""对每个 shot 的首尾帧算光流,填充 movement/intensity。失败时静默降级为 static/low。"""
|
||||
try:
|
||||
import cv2 # noqa: PLC0415 - opencv-python-headless 已在 worker requirements 中
|
||||
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OpenCV 不可用,运镜检测降级为 static/low: %s", exc)
|
||||
return
|
||||
try:
|
||||
cap = cv2.VideoCapture(str(video_path))
|
||||
for shot in shots:
|
||||
mid_t = (shot.start_sec + shot.end_sec) / 2.0
|
||||
dt = max(0.2, min(0.5, (shot.end_sec - shot.start_sec) / 4.0))
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, max(0.0, (mid_t - dt)) * 1000)
|
||||
ok1, f1 = cap.read()
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, min(mid_t + dt, shot.end_sec - 0.05) * 1000)
|
||||
ok2, f2 = cap.read()
|
||||
if not (ok1 and ok2):
|
||||
continue
|
||||
g1 = cv2.cvtColor(f1, cv2.COLOR_BGR2GRAY)
|
||||
g2 = cv2.cvtColor(f2, cv2.COLOR_BGR2GRAY)
|
||||
h, w = g1.shape
|
||||
# 降采样加速
|
||||
scale = 360.0 / h if h > 360 else 1.0
|
||||
if scale < 1.0:
|
||||
g1 = cv2.resize(g1, (int(w * scale), int(h * scale)))
|
||||
g2 = cv2.resize(g2, (int(w * scale), int(h * scale)))
|
||||
flow = cv2.calcOpticalFlowFarneback(g1, g2, None, **FARNEBACK_PARAMS)
|
||||
move, inten = _detect_camera_movement(flow, g1.shape[1], g1.shape[0])
|
||||
shot.movement = move
|
||||
shot.intensity = inten
|
||||
cap.release()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("运镜检测异常,已降级: %s", exc)
|
||||
|
||||
|
||||
# ── ④ librosa BPM ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _detect_bpm(video_path: Path) -> int:
|
||||
"""提取音轨并估算 BPM;失败返回 0。"""
|
||||
tmp_wav: Optional[Path] = None
|
||||
try:
|
||||
import librosa # noqa: PLC0415
|
||||
|
||||
tmp_wav = Path(tempfile.mkstemp(suffix=".wav")[1])
|
||||
# ffmpeg 抽 22050Hz 单声道 wav
|
||||
subprocess.run(
|
||||
[_ffmpeg_bin(), "-y", "-i", str(video_path), "-vn", "-ac", "1", "-ar", "22050", "-f", "wav", str(tmp_wav)],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=20,
|
||||
check=False,
|
||||
) # nosec B603
|
||||
if not tmp_wav.exists() or tmp_wav.stat().st_size < 1024:
|
||||
return 0
|
||||
y, sr = librosa.load(str(tmp_wav), sr=22050, mono=True)
|
||||
if len(y) < sr * 2:
|
||||
return 0
|
||||
tempo, _ = librosa.beat.beat_track(y=y, sr=sr)
|
||||
try:
|
||||
bpm = int(round(float(tempo)))
|
||||
except Exception: # noqa: BLE001
|
||||
bpm = int(round(float(tempo[0]))) if len(tempo) else 0
|
||||
return max(40, min(bpm, 220))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("librosa BPM 分析失败: %s", exc)
|
||||
return 0
|
||||
finally:
|
||||
if tmp_wav and tmp_wav.exists():
|
||||
try:
|
||||
tmp_wav.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _pace_from_bpm(bpm: int) -> str:
|
||||
if bpm >= 110:
|
||||
return "fast_cut"
|
||||
if bpm >= 80:
|
||||
return "medium"
|
||||
if bpm > 0:
|
||||
return "slow_cinematic"
|
||||
return "medium"
|
||||
|
||||
|
||||
# ── ⑤ VLM 帧分析 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_frames(frame_paths: list[Path], shots: list[ShotBoundary], k: int = VLM_SAMPLE_FRAMES) -> list[Path]:
|
||||
"""从全量帧中均匀选 k 张代表性帧(优先场景帧)。"""
|
||||
if not frame_paths:
|
||||
return []
|
||||
scene_frames = sorted(p for p in frame_paths if p.name.startswith("scene_"))
|
||||
fixed_frames = sorted(p for p in frame_paths if p.name.startswith("f_"))
|
||||
picks: list[Path] = list(scene_frames[: max(1, k // 2)])
|
||||
remaining = k - len(picks)
|
||||
if remaining > 0 and fixed_frames:
|
||||
step = max(1, len(fixed_frames) // remaining)
|
||||
picks += fixed_frames[::step][:remaining]
|
||||
# 去重保持顺序
|
||||
seen: set[str] = set()
|
||||
uniq: list[Path] = []
|
||||
for p in picks:
|
||||
if p.name not in seen and p.exists():
|
||||
seen.add(p.name)
|
||||
uniq.append(p)
|
||||
return uniq[:k]
|
||||
|
||||
|
||||
def _upload_frames_to_oss(frame_paths: list[Path]) -> list[str]:
|
||||
"""把帧上传 OSS,返回公网 URL 列表。失败时降级为 data URI。"""
|
||||
urls: list[str] = []
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
for p in frame_paths:
|
||||
try:
|
||||
key = f"viral-video/analysis/{uuid.uuid4().hex}/{p.name}"
|
||||
url = upload_to_oss(p, key)
|
||||
if url:
|
||||
urls.append(url)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("单帧 OSS 上传失败 %s: %s", p.name, exc)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OSS 上传模块不可用,降级为 base64 data URI: %s", exc)
|
||||
if len(urls) < len(frame_paths):
|
||||
# 降级:base64 data URI(小图,单张 ≤100KB 才走此路)
|
||||
import base64
|
||||
|
||||
for p in frame_paths[len(urls) :]:
|
||||
try:
|
||||
if p.stat().st_size > 120_000:
|
||||
continue
|
||||
b64 = base64.b64encode(p.read_bytes()).decode("ascii")
|
||||
urls.append(f"data:image/jpeg;base64,{b64}")
|
||||
except Exception: # noqa: BLE001 # nosec B112
|
||||
continue
|
||||
return urls
|
||||
|
||||
|
||||
def _vlm_analyze_frames(image_urls: list[str]) -> dict[str, Any]:
|
||||
"""调豆包 VLM 分析色调/构图/光线/转场观感。"""
|
||||
if not image_urls:
|
||||
return {}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演和调色师。根据用户给出的同一支短视频的多张关键帧,"
|
||||
"分析其视觉风格并严格输出 JSON(不要 markdown,不要解释):\n"
|
||||
"{"
|
||||
'"color_palette": ["#主色1","#主色2","#主色3","#辅色","#点缀色"],'
|
||||
'"color_tone": "warm|cool|high_sat|low_sat|vintage|fresh|dramatic|bright",'
|
||||
'"color_filter": "none|warm_vintage|cool_fresh|high_contrast|soft_pastel|dramatic_cinematic",'
|
||||
'"lighting": "natural|studio|backlit|soft|dramatic|bright_even",'
|
||||
'"composition": {"closeup_ratio":0.0,"medium_ratio":0.0,"wide_ratio":0.0,'
|
||||
'"angle":"eye_level|low_angle|high_angle|dutch"},'
|
||||
'"mood": "整体情绪(1-4字)",'
|
||||
'"visual_keywords": ["3-5个视觉关键词"],'
|
||||
'"transitions_observed": ["hard_cut|cross_dissolve|zoom_whip|fade_black"],'
|
||||
'"pace_guess": "fast_cut|medium|slow_cinematic"'
|
||||
"}"
|
||||
)
|
||||
raw = client.vision_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": "请分析这支参考视频的风格。"},
|
||||
],
|
||||
images=image_urls,
|
||||
temperature=0.2,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return {}
|
||||
raw = _strip_code_fence(raw)
|
||||
# 容忍模型可能前后加文本
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
return json.loads(raw[i : j + 1])
|
||||
return {}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("VLM 帧分析失败: %s", exc)
|
||||
return {}
|
||||
|
||||
|
||||
# ── ⑥ LLM 整合 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _llm_synthesize(
|
||||
shots: list[ShotBoundary],
|
||||
bpm: int,
|
||||
vlm: dict[str, Any],
|
||||
style_strength: str,
|
||||
) -> dict[str, Any]:
|
||||
"""把结构化信号整合成 style_guide;LLM 不可用时走规则合成。"""
|
||||
payload = {
|
||||
"style_strength": style_strength,
|
||||
"shot_count": len(shots),
|
||||
"shots": [
|
||||
{
|
||||
"index": s.index,
|
||||
"start_sec": round(s.start_sec, 2),
|
||||
"end_sec": round(s.end_sec, 2),
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"transition": s.transition,
|
||||
}
|
||||
for s in shots
|
||||
],
|
||||
"bpm": bpm,
|
||||
"pace_guess": _pace_from_bpm(bpm),
|
||||
"vlm": vlm,
|
||||
}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演。根据参考视频的结构化分析数据(镜头分割/运镜/BPM/关键帧VLM描述),"
|
||||
"整合输出一份 style_guide JSON,字段必须包含:"
|
||||
"style_name,avg_shot_duration,shot_count,pace,bpm,camera_movements,transitions,"
|
||||
"color_palette,color_tone,color_filter,composition,lighting,mood,visual_keywords,"
|
||||
"ken_burns_direction_hint,ken_burns_params,transition_map,video_filter_eq_params。"
|
||||
"严格输出一个合法 JSON 对象,不要 markdown/解释。"
|
||||
)
|
||||
user_text = "分析数据:\n" + json.dumps(payload, ensure_ascii=False)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": sys_prompt}, {"role": "user", "content": user_text}],
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
if raw:
|
||||
raw = _strip_code_fence(raw)
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
result = json.loads(raw[i : j + 1])
|
||||
if isinstance(result, dict) and result.get("style_name"):
|
||||
return result
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("LLM 合成 style_guide 失败,走规则降级: %s", exc)
|
||||
return _rule_based_style_guide(shots, bpm, vlm)
|
||||
|
||||
|
||||
def _rule_based_style_guide(shots: list[ShotBoundary], bpm: int, vlm: dict[str, Any]) -> dict[str, Any]:
|
||||
"""LLM 不可用时,用规则拼出可用 style_guide。"""
|
||||
durations = [s.end_sec - s.start_sec for s in shots] or [3.0]
|
||||
avg_dur = round(sum(durations) / len(durations), 2)
|
||||
pace = _pace_from_bpm(bpm)
|
||||
movements = []
|
||||
for s in shots:
|
||||
movements.append(
|
||||
{
|
||||
"shot_index": s.index + 1,
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"duration": round(s.end_sec - s.start_sec, 2),
|
||||
"subject_hint": _default_subject_hint(s.movement),
|
||||
}
|
||||
)
|
||||
transitions = []
|
||||
for i in range(len(shots) - 1):
|
||||
transitions.append({"between_shot": [i + 1, i + 2], "type": shots[i].transition})
|
||||
color_palette = vlm.get("color_palette") or ["#E0E0E0", "#333333", "#F5F5F5", "#888888", "#FF6B35"]
|
||||
color_tone = vlm.get("color_tone") or "bright"
|
||||
color_filter = vlm.get("color_filter") or "none"
|
||||
lighting = vlm.get("lighting") or "bright_even"
|
||||
composition = vlm.get("composition") or {
|
||||
"closeup_ratio": 0.4,
|
||||
"medium_ratio": 0.4,
|
||||
"wide_ratio": 0.2,
|
||||
"angle": "eye_level",
|
||||
}
|
||||
mood = vlm.get("mood") or "明快"
|
||||
vk = vlm.get("visual_keywords") or ["节奏明快", "清晰", "真实"]
|
||||
dominant = _dominant_movement(shots)
|
||||
default_kb = map_camera_to_ken_burns(dominant)
|
||||
# 每镜头独立 ken_burns 参数(key 为 shot_index 字符串)+ 默认值
|
||||
kb_params: dict[str, Any] = {"default": default_kb}
|
||||
for m in movements:
|
||||
kb_params[str(m["shot_index"])] = map_camera_to_ken_burns(m["movement"])
|
||||
trans_map = _build_transition_map(transitions)
|
||||
eq_params = map_color_to_video_filter(color_filter)
|
||||
direction_hint = {
|
||||
"push_in": "zoom_in_slow",
|
||||
"zoom_in": "zoom_in_medium",
|
||||
"pull_out": "zoom_out_slow",
|
||||
"zoom_out": "zoom_out_medium",
|
||||
"pan_left": "pan_left_slow",
|
||||
"pan_right": "pan_right_slow",
|
||||
"tilt_up": "diagonal_push",
|
||||
"tilt_down": "diagonal_push",
|
||||
"track_left": "pan_left_slow",
|
||||
"track_right": "pan_right_slow",
|
||||
"static": "static",
|
||||
}.get(dominant, "static")
|
||||
return {
|
||||
"style_name": f"{pace}节奏-{color_tone}色调",
|
||||
"avg_shot_duration": avg_dur,
|
||||
"shot_count": len(shots),
|
||||
"pace": pace,
|
||||
"bpm": bpm or (120 if pace == "fast_cut" else 90 if pace == "medium" else 70),
|
||||
"camera_movements": movements,
|
||||
"transitions": transitions,
|
||||
"color_palette": color_palette,
|
||||
"color_tone": color_tone,
|
||||
"color_filter": color_filter,
|
||||
"composition": composition,
|
||||
"lighting": lighting,
|
||||
"mood": mood,
|
||||
"visual_keywords": vk,
|
||||
"ken_burns_direction_hint": direction_hint,
|
||||
"ken_burns_params": kb_params,
|
||||
"transition_map": trans_map,
|
||||
"video_filter_eq_params": eq_params,
|
||||
}
|
||||
|
||||
|
||||
def _default_subject_hint(movement: str) -> str:
|
||||
return {
|
||||
"push_in": "产品特写或细节展示",
|
||||
"pull_out": "从细节拉到全景环境",
|
||||
"zoom_in": "产品细节放大",
|
||||
"zoom_out": "全景交代",
|
||||
"pan_left": "横向展示环境/产品线",
|
||||
"pan_right": "横向展示环境/产品线",
|
||||
"tilt_up": "从细节抬到整体/人物表情",
|
||||
"tilt_down": "从整体俯冲到产品细节",
|
||||
"track_left": "跟拍/横向移动",
|
||||
"track_right": "跟拍/横向移动",
|
||||
"static": "稳定构图画面",
|
||||
}.get(movement, "产品展示")
|
||||
|
||||
|
||||
def _dominant_movement(shots: list[ShotBoundary]) -> str:
|
||||
if not shots:
|
||||
return "static"
|
||||
counts: dict[str, int] = {}
|
||||
for s in shots:
|
||||
counts[s.movement] = counts.get(s.movement, 0) + 1
|
||||
return max(counts, key=counts.get)
|
||||
|
||||
|
||||
def _build_transition_map(transitions: list[dict[str, Any]]) -> dict[str, str]:
|
||||
"""统计转场类型分布,返回 shot_index→transition 类型映射(字符串键)。"""
|
||||
m: dict[str, str] = {}
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or [0, 0]
|
||||
if len(pair) >= 2:
|
||||
m[f"{pair[0]}-{pair[1]}"] = t.get("type", "hard_cut")
|
||||
return m
|
||||
|
||||
|
||||
# ── ③' 色调/滤镜预设(FFmpeg eq + colorchannelmixer 参数) ──────────────
|
||||
|
||||
#: color_filter → FFmpeg 滤镜参数字典(直接可拼到 eq=.../colorchannelmixer=...)
|
||||
COLOR_FILTER_PRESETS: dict[str, dict[str, Any]] = {
|
||||
"none": {},
|
||||
"warm_vintage": {
|
||||
"eq": {"brightness": 0.02, "contrast": 1.05, "saturation": 0.9, "gamma": 1.05},
|
||||
"colorchannelmixer": {"rr": 1.1, "gg": 0.98, "bb": 0.82, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"cool_fresh": {
|
||||
"eq": {"brightness": 0.03, "contrast": 1.08, "saturation": 1.05},
|
||||
"colorchannelmixer": {"rr": 0.9, "gg": 1.0, "bb": 1.12, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"high_contrast": {
|
||||
"eq": {"brightness": 0.0, "contrast": 1.3, "saturation": 1.2},
|
||||
"colorchannelmixer": {},
|
||||
},
|
||||
"soft_pastel": {
|
||||
"eq": {"brightness": 0.05, "contrast": 0.92, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 1.03, "bb": 1.05, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"dramatic_cinematic": {
|
||||
"eq": {"brightness": -0.03, "contrast": 1.2, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 0.98, "bb": 0.9, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_color_to_video_filter(color_filter: str) -> dict[str, Any]:
|
||||
"""color_filter 枚举 → FFmpeg eq/colorchannelmixer 参数字典(渲染端直接使用)。"""
|
||||
preset = COLOR_FILTER_PRESETS.get(color_filter) or COLOR_FILTER_PRESETS["none"]
|
||||
# 返回深拷贝防污染
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── ③'' 运镜 → ken_burns 参数映射 ───────────────────────────────────────
|
||||
|
||||
#: 运镜类型 → URS 可直接消费的 ken_burns 参数字典
|
||||
CAMERA_TO_KEN_BURNS: dict[str, dict[str, Any]] = {
|
||||
"static": {
|
||||
"type": "static",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"push_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.12,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.18,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pull_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.12,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.18,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_up": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.08,
|
||||
"zoom_end": 1.14,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": -0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_down": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.14,
|
||||
"zoom_end": 1.08,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_camera_to_ken_burns(movement: str) -> dict[str, Any]:
|
||||
"""运镜类型 → URS ken_burns 参数字典。未知类型回退 static。"""
|
||||
preset = CAMERA_TO_KEN_BURNS.get(movement) or CAMERA_TO_KEN_BURNS["static"]
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── 转场 → xfade transition 名称 ────────────────────────────────────────
|
||||
|
||||
TRANSITION_TO_XFADE: dict[str, str] = {
|
||||
"hard_cut": "cut",
|
||||
"cross_dissolve": "dissolve",
|
||||
"fade_black": "fadeblack",
|
||||
"fade": "fade",
|
||||
"zoom_whip": "zoom",
|
||||
"slide_left": "slideright", # 画面左移 = 新画面从右滑入
|
||||
"slide_right": "slideleft",
|
||||
"wipe_left": "wipeleft",
|
||||
"wipe_right": "wiperight",
|
||||
}
|
||||
|
||||
|
||||
def map_transition_to_xfade(transition_type: str) -> str:
|
||||
"""转场枚举 → TransitionEngine 支持的 xfade 名称;未知回退 cut。"""
|
||||
return TRANSITION_TO_XFADE.get(transition_type, "cut")
|
||||
|
||||
|
||||
# ── BPM → BGM 推荐 BPM ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def map_bgm_bpm(bpm: int) -> int:
|
||||
"""BGM 选曲 BPM:参考视频 BPM ±5。bpm=0 返回 90(默认 medium)。"""
|
||||
if bpm <= 0:
|
||||
return 90
|
||||
return max(60, min(bpm, 180))
|
||||
|
||||
|
||||
# ── 单 clip 渲染参数聚合(给 URS build_render_plan 使用) ───────────────
|
||||
|
||||
|
||||
def build_render_params_for_clip(
|
||||
clip_index: int,
|
||||
style_guide: dict[str, Any],
|
||||
*,
|
||||
duration_sec: Optional[float] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""根据 style_guide 为第 clip_index 个 clip 生成可直接喂给 URS 的渲染参数。"""
|
||||
shot_idx = clip_index + 1
|
||||
movements = style_guide.get("camera_movements") or []
|
||||
movement = "static"
|
||||
intensity = "low"
|
||||
for m in movements:
|
||||
if m.get("shot_index") == shot_idx:
|
||||
movement = m.get("movement", "static")
|
||||
intensity = m.get("intensity", "low")
|
||||
break
|
||||
ken = map_camera_to_ken_burns(movement)
|
||||
if intensity == "high":
|
||||
ken["zoom_end"] = round(ken.get("zoom_end", 1.0) * 1.08, 3)
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 1.3, 3)
|
||||
elif intensity == "low":
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 0.6, 3)
|
||||
transitions = style_guide.get("transitions") or []
|
||||
trans_type = "hard_cut"
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or []
|
||||
if len(pair) >= 2 and pair[0] == shot_idx:
|
||||
trans_type = t.get("type", "hard_cut")
|
||||
break
|
||||
xfade = map_transition_to_xfade(trans_type)
|
||||
eq = map_color_to_video_filter(style_guide.get("color_filter", "none"))
|
||||
return {
|
||||
"ken_burns": ken,
|
||||
"transition": {"type": xfade, "duration": 0.3 if xfade != "cut" else 0.0},
|
||||
"video_filter": eq,
|
||||
"bgm_bpm_hint": map_bgm_bpm(int(style_guide.get("bpm") or 0)),
|
||||
"duration_sec": duration_sec,
|
||||
}
|
||||
|
||||
|
||||
# ── 素材本地化(URL/OSS key → 本地临时文件) ────────────────────────────
|
||||
|
||||
|
||||
def _ensure_local_video(reference: str, work_dir: Path) -> Optional[Path]:
|
||||
"""把 reference(URL/OSS key/本地路径)落到 work_dir 下的本地文件。"""
|
||||
p = Path(reference)
|
||||
if p.exists() and p.is_file():
|
||||
return p
|
||||
try:
|
||||
from video_processing.oss_helpers import download_asset
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
ok = download_asset(reference, target)
|
||||
if ok and target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("download_asset 失败,尝试 http 直连: %s", exc)
|
||||
if reference.startswith(("http://", "https://")):
|
||||
try:
|
||||
import httpx # noqa: PLC0415 - 项目依赖,延迟导入
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
with httpx.Client(timeout=20.0, follow_redirects=True) as client:
|
||||
with client.stream("GET", reference) as resp:
|
||||
resp.raise_for_status()
|
||||
with open(target, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=64 * 1024):
|
||||
f.write(chunk)
|
||||
if target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("HTTP 下载参考视频失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── 入口 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def analyze_video_style(
|
||||
reference_video_path: str | Path,
|
||||
style_strength: str = "medium",
|
||||
*,
|
||||
timeout_sec: int = DEFAULT_ANALYSIS_TIMEOUT,
|
||||
) -> dict[str, Any]:
|
||||
"""分析参考视频风格,返回 style_guide dict。
|
||||
|
||||
Args:
|
||||
reference_video_path: 本地路径、HTTP(S) URL 或 OSS storage key。
|
||||
style_strength: light | medium | strict。
|
||||
timeout_sec: 单步超时(秒),默认 60。
|
||||
|
||||
Returns:
|
||||
style_guide dict,详见 STYLE_GUIDE_SCHEMA。任何子步骤失败都会降级,不抛异常。
|
||||
"""
|
||||
style_strength = style_strength if style_strength in ("light", "medium", "strict") else "medium"
|
||||
frames_dir: Optional[Path] = None
|
||||
local_path: Optional[Path] = None
|
||||
try:
|
||||
frames_dir = Path(tempfile.mkdtemp(prefix="vstyle_"))
|
||||
work_dir = frames_dir # 同一临时根
|
||||
local_path = _ensure_local_video(str(reference_video_path), work_dir)
|
||||
if local_path is None:
|
||||
logger.error("[video_analyzer] 无法获取参考视频: %s", reference_video_path)
|
||||
return _rule_based_style_guide([], 0, {})
|
||||
# 资源约束:大小 / 时长
|
||||
try:
|
||||
size_mb = local_path.stat().st_size / (1024 * 1024)
|
||||
if size_mb > MAX_REFERENCE_SIZE_MB:
|
||||
logger.warning(
|
||||
"[video_analyzer] 参考视频 %.1fMB 超上限,按前 %ds 分析", size_mb, MAX_REFERENCE_DURATION_SEC
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
duration = _probe_duration(local_path)
|
||||
if duration > MAX_REFERENCE_DURATION_SEC:
|
||||
duration = MAX_REFERENCE_DURATION_SEC
|
||||
# ① 抽帧
|
||||
try:
|
||||
frame_paths = _extract_keyframes(local_path, frames_dir / "frames")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("FFmpeg 抽帧失败: %s,降级为 VLM 均匀采样", exc)
|
||||
frame_paths = []
|
||||
# ② 镜头分割
|
||||
shots = _detect_shots(local_path, frames_dir)
|
||||
# ③ 运镜检测(有帧才跑)
|
||||
if frame_paths or shots:
|
||||
_analyze_movements(local_path, shots)
|
||||
# ④ BPM
|
||||
bpm = _detect_bpm(local_path)
|
||||
# ⑤ 选帧→OSS→VLM
|
||||
sampled = _sample_frames(frame_paths, shots)
|
||||
image_urls = _upload_frames_to_oss(sampled) if sampled else []
|
||||
vlm = _vlm_analyze_frames(image_urls) if image_urls else {}
|
||||
# ⑥ 合成
|
||||
style_guide = _llm_synthesize(shots, bpm, vlm, style_strength)
|
||||
# 兜底字段校验
|
||||
style_guide.setdefault("style_strength", style_strength)
|
||||
style_guide.setdefault("pace", _pace_from_bpm(bpm))
|
||||
style_guide.setdefault("bpm", bpm)
|
||||
style_guide.setdefault("shot_count", len(shots))
|
||||
if shots and "avg_shot_duration" not in style_guide:
|
||||
durs = [s.end_sec - s.start_sec for s in shots]
|
||||
style_guide["avg_shot_duration"] = round(sum(durs) / len(durs), 2)
|
||||
return style_guide
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("[video_analyzer] 整体分析异常,返回最小占位 style_guide: %s", exc)
|
||||
return _rule_based_style_guide([], 0, {"mood": "未知"})
|
||||
finally:
|
||||
# 临时帧清理
|
||||
if frames_dir and frames_dir.exists():
|
||||
shutil.rmtree(frames_dir, ignore_errors=True)
|
||||
@@ -38,6 +38,7 @@ celery_app.conf.imports = (
|
||||
"worker_app.tasks.voice_extraction",
|
||||
"worker_app.tasks.voice_clone",
|
||||
"worker_app.tasks.tts_synthesis",
|
||||
"worker_app.tasks.viral_video", # #2039 爆款视频编排器(10步流水线)
|
||||
"worker_app.tasks.batch_download",
|
||||
"worker_app.tasks.duplication_check",
|
||||
# #1798 AI 数字人渲染:必须在 Worker 实例上注册同名任务,否则消息无人消费(渲染卡 0%)
|
||||
|
||||
@@ -0,0 +1,832 @@
|
||||
"""爆款视频 Celery 编排器 — ViralVideoOrchestrator.
|
||||
|
||||
9 步流水线(Seedance 2.5 直生口型,不再走 MuseTalk):
|
||||
1. 图片 VLM 分析
|
||||
1.5 [v1.3] 视频风格分析(如用户上传参考视频)
|
||||
2. 用户文案意图解析
|
||||
3. 文案融合生成
|
||||
4. 分镜脚本生成
|
||||
5. 合规审核(6 维度,不通过自动重写 1 次)
|
||||
6. CosyVoice 配音
|
||||
7. BGM 选择(素材未就绪时跳过)
|
||||
8. Seedance 逐分镜生成 + ffmpeg concat + 混 TTS
|
||||
9. OSS 上传 + 通知 + 扣点
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
from worker_app.celery_app import celery_app # noqa: F401 - 加载 app 以注册任务
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
ViralVideoJob,
|
||||
ViralVideoStage,
|
||||
ViralVideoStatus,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── WS 进度推送 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _emit_progress(
|
||||
job_id: str,
|
||||
stage: str,
|
||||
progress: float,
|
||||
message: str = "",
|
||||
data: dict | None = None,
|
||||
event_type: str = "viral_video:progress",
|
||||
):
|
||||
"""通过 Redis 发布进度事件,供 WebSocket 消费。
|
||||
|
||||
event_type 取值:
|
||||
- viral_video:progress 中间进度(默认)
|
||||
- viral_video:completed 任务完成
|
||||
- viral_video:failed 任务失败
|
||||
- viral_video:wait_user 等待用户确认
|
||||
所有事件 payload 均为合法 JSON,前端 JSON.parse 即可。
|
||||
"""
|
||||
try:
|
||||
import redis as redis_lib
|
||||
|
||||
redis_url = os.environ.get("REDIS_URL", "redis://localhost:6379/0")
|
||||
r = redis_lib.from_url(redis_url)
|
||||
event = {
|
||||
"type": event_type,
|
||||
"job_id": job_id,
|
||||
"stage": stage,
|
||||
"progress": progress,
|
||||
"message": message or STAGE_LABELS.get(stage, stage),
|
||||
"data": data or {},
|
||||
}
|
||||
r.publish(f"viral_video:{job_id}", json.dumps(event, ensure_ascii=False))
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] WS 进度推送失败: %s", e)
|
||||
|
||||
|
||||
# ── 仓储辅助 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_repo_and_job(job_id: str):
|
||||
"""获取 session, repo, job 三元组。"""
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
return session, repo, job
|
||||
|
||||
|
||||
def _save_job(repo, job, session):
|
||||
"""持久化并关闭 session。"""
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _step_image_analysis(job: ViralVideoJob) -> dict:
|
||||
"""步骤 1: 图片 VLM 分析 — 识别产品特征、场景、卖点。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_vision
|
||||
except ImportError:
|
||||
logger.warning("[爆款视频] ai_service.call_vision 不可用,使用占位结果")
|
||||
return {"products": [{"name": "产品", "features": ["特征1", "特征2"], "scene": "通用场景"}]}
|
||||
|
||||
results = []
|
||||
for img_url in job.images:
|
||||
try:
|
||||
result = call_vision(
|
||||
image_url=img_url,
|
||||
prompt="请分析这张产品图片,识别:1)产品名称和类别 2)主要特征和卖点 3)适用场景 4)视觉风格。以JSON格式返回。",
|
||||
)
|
||||
results.append(result)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 图片分析失败 img=%s: %s", img_url, e)
|
||||
results.append({"name": "未识别", "features": [], "scene": "通用"})
|
||||
|
||||
return {"products": results}
|
||||
|
||||
|
||||
def _step_video_analysis(job: ViralVideoJob) -> dict | None:
|
||||
"""步骤 1.5 [v1.3]: 参考视频风格分析。"""
|
||||
if not job.reference_video_url:
|
||||
return None
|
||||
|
||||
try:
|
||||
# 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 as e:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪(%s),使用占位风格分析", e)
|
||||
return {
|
||||
"cut_speed": "medium",
|
||||
"transition": "cross_dissolve",
|
||||
"energy": "medium",
|
||||
"color_grade": "neutral",
|
||||
"narrative": False,
|
||||
"source": "placeholder",
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 视频风格分析失败: %s", e)
|
||||
return {"error": str(e), "source": "failed"}
|
||||
|
||||
|
||||
def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
"""步骤 2: 用户文案意图解析 — 理解用户想表达什么。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业"}
|
||||
|
||||
products_summary = ""
|
||||
for p in image_analysis.get("products", []):
|
||||
products_summary += f"- {p.get('name', '产品')}: {', '.join(p.get('features', []))}\n"
|
||||
|
||||
prompt = f"""你是一个营销文案策略师。请分析以下信息,理解用户的营销意图:
|
||||
|
||||
用户原始文案:{job.user_copy_text or "(未提供)"}
|
||||
行业:{job.industry or "未指定"}
|
||||
目标客户:{job.target_customer or "未指定"}
|
||||
营销目的:{job.marketing_purpose or "未指定"}
|
||||
产品信息:
|
||||
{products_summary}
|
||||
|
||||
请分析并返回JSON格式:
|
||||
1. intent: 核心营销意图(一句话)
|
||||
2. key_messages: 要传达的3-5个关键信息
|
||||
3. tone: 文案调性(如:专业/亲切/高端/活力)
|
||||
4. target_emotion: 希望触发的用户情感
|
||||
5. call_to_action: 行动号召建议"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, dict) else {"raw": result}
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 意图解析失败: %s", e)
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业"}
|
||||
|
||||
|
||||
def _step_copy_fusion(job: ViralVideoJob, intent: dict, image_analysis: dict) -> str:
|
||||
"""步骤 3: 文案融合生成 — 根据 fusion_level 融合用户文案和 AI 文案。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return f"【{job.industry or '行业'}】优质产品,{job.target_customer or '您'}的不二之选!"
|
||||
|
||||
products_desc = ""
|
||||
for p in image_analysis.get("products", []):
|
||||
products_desc += f"{p.get('name', '产品')}({','.join(p.get('features', []))})\n"
|
||||
|
||||
if job.fusion_level == "ai_full":
|
||||
prompt = f"""请为以下产品撰写一段爆款短视频文案({job.duration}秒):
|
||||
产品:{products_desc}
|
||||
行业:{job.industry}
|
||||
目标客户:{job.target_customer}
|
||||
营销目的:{job.marketing_purpose}
|
||||
调性:{intent.get("tone", "专业")}
|
||||
关键信息:{", ".join(intent.get("key_messages", []))}
|
||||
|
||||
要求:吸引眼球、节奏紧凑、有行动号召。直接输出文案内容。"""
|
||||
elif job.fusion_level == "user_primary":
|
||||
prompt = f"""请基于用户原始文案进行润色优化,保留用户原意和风格:
|
||||
用户原文:{job.user_copy_text}
|
||||
产品信息:{products_desc}
|
||||
|
||||
要求:保留用户原意,仅修正表达和节奏。直接输出文案内容。"""
|
||||
else: # ai_polish (default)
|
||||
prompt = f"""请将用户文案与AI分析融合,生成一段优化后的爆款短视频文案({job.duration}秒):
|
||||
用户原文:{job.user_copy_text or "(未提供)"}
|
||||
产品分析:{products_desc}
|
||||
行业:{job.industry}
|
||||
目标客户:{job.target_customer}
|
||||
营销目的:{job.marketing_purpose}
|
||||
意图分析:{intent.get("intent", "")}
|
||||
调性:{intent.get("tone", "专业")}
|
||||
|
||||
要求:融合用户意图和产品卖点,节奏紧凑,适合短视频。直接输出文案内容。"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, str) else str(result)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 文案融合失败: %s", e)
|
||||
return job.user_copy_text or f"精选{job.industry or '行业'}好物,值得关注!"
|
||||
|
||||
|
||||
def _step_storyboard(job: ViralVideoJob, copy_text: str, image_analysis: dict) -> list[dict]:
|
||||
"""步骤 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": min(5, job.duration),
|
||||
"description": "产品展示",
|
||||
"ken_burns": "zoom_in",
|
||||
"transition": "cut",
|
||||
}
|
||||
]
|
||||
|
||||
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}秒(每个分镜 3~6 秒,总和约等于总时长)
|
||||
风格强度:{job.style_strength}
|
||||
输出宽高比:{ratio}{products_hint}
|
||||
|
||||
请以 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 _normalize_storyboard(result, job.duration, n_segments, copy_text)
|
||||
import json as _json
|
||||
|
||||
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 _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:
|
||||
"""步骤 5: 合规审核(6 维度)。不通过时自动重写 1 次。"""
|
||||
dimensions = ["广告法合规", "平台规范", "内容真实性", "版权安全", "价值观", "风格一致性"]
|
||||
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return {"passed": True, "score": 90, "details": {d: "通过" for d in dimensions}}
|
||||
|
||||
prompt = f"""请对以下短视频内容进行合规审核,检查6个维度:{", ".join(dimensions)}
|
||||
|
||||
文案内容:{copy_text}
|
||||
分镜脚本:{storyboard[:3]}...
|
||||
行业:{job.industry}
|
||||
|
||||
请以JSON格式返回:
|
||||
- passed: bool(是否全部通过)
|
||||
- score: int(0-100分)
|
||||
- details: 各维度评分和说明
|
||||
- issues: 需要修改的问题列表(如有)"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, dict) else {"passed": True, "score": 80, "details": {}}
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 合规审核失败: %s", e)
|
||||
return {"passed": True, "score": 75, "details": {d: "默认通过" for d in dimensions}}
|
||||
|
||||
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str):
|
||||
"""步骤 6: CosyVoice 配音。P1:返回 Path;失败返回 None。"""
|
||||
try:
|
||||
from pathlib import Path as _Path
|
||||
|
||||
from services.tts_service_factory import get_tts_service
|
||||
|
||||
tts_service = get_tts_service()
|
||||
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 None
|
||||
|
||||
|
||||
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": "固定镜头",
|
||||
}
|
||||
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, storyboard, tts_path, bgm):
|
||||
"""步骤 8: 渲染(P0-1 核心重写)。
|
||||
|
||||
每个 storyboard 分镜 → Seedance 2.5 生成短视频段(无声)→ 下载 → ffmpeg concat → 混入 TTS。
|
||||
返回最终视频本地路径字符串。
|
||||
"""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from video_processing.concat_engine import concat_video_files
|
||||
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
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)
|
||||
|
||||
_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("[爆款视频] 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:
|
||||
import subprocess
|
||||
from pathlib import Path as _Path
|
||||
|
||||
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_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 10: OSS 上传。"""
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
video_url = upload_to_oss(video_path, prefix="viral-video/")
|
||||
return video_url
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] OSS 上传失败: %s", e)
|
||||
raise
|
||||
|
||||
|
||||
# ── 主编排器 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@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 步流水线编排器(前半段:图片分析→风格分析→意图解析,然后 WAIT_USER_CONFIRM)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
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, "开始图片分析")
|
||||
|
||||
# ── 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)
|
||||
job.style_guide = style_guide
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.VIDEO_ANALYSIS,
|
||||
25.0,
|
||||
"风格分析完成",
|
||||
{"style_analyzed": True, "style_guide": style_guide},
|
||||
)
|
||||
|
||||
# ── 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(
|
||||
job_id,
|
||||
ViralVideoStage.INTENT_PARSING,
|
||||
35.0,
|
||||
"意图解析完成,等待用户确认",
|
||||
{"intent_result": intent_result, "waiting_confirm": True},
|
||||
)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.INTENT_PARSING,
|
||||
35.0,
|
||||
"等待用户确认意图文案",
|
||||
{"intent_result": intent_result},
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
|
||||
return {"ok": True, "job_id": job_id, "status": "wait_user_confirm", "intent_result": intent_result}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 流水线异常: %s", e, exc_info=True)
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
else:
|
||||
_, repo, job = _get_repo_and_job(job_id)
|
||||
if job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@shared_task(bind=True, max_retries=2, name="worker.resume_viral_video_pipeline")
|
||||
def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""用户确认意图后,从断点恢复流水线(步骤 3-10)。"""
|
||||
session = None
|
||||
job = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
if job.status != ViralVideoStatus.RUNNING:
|
||||
return {"ok": False, "error": f"unexpected status: {job.status}"}
|
||||
|
||||
# P0-3: 从 job 读取 image_analysis(run_pipeline 阶段已持久化)
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...")
|
||||
|
||||
# ── Step 3: 文案融合 ──
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成")
|
||||
|
||||
# ── Step 4: 分镜脚本 ──
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...")
|
||||
storyboard = _step_storyboard(job, copy_text, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)})
|
||||
|
||||
# ── Step 5: 合规审核 ──
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...")
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
if not review_result.get("passed", True):
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
|
||||
|
||||
# ── Step 6: CosyVoice 配音(返回 Path | None) ──
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...")
|
||||
tts_path = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None})
|
||||
|
||||
# ── Step 7: BGM 选择(P1:暂返回 None,跳过) ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)")
|
||||
bgm = _step_bgm_select(job)
|
||||
|
||||
# ── Step 8: 渲染(逐分镜 Seedance → concat → 混 TTS) ──
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...")
|
||||
video_path = _step_render(job, storyboard, tts_path, bgm)
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成")
|
||||
|
||||
# ── Step 9: OSS 上传 + 扣点 ──
|
||||
# 注:爆款视频由 Seedance 2.5 直接生成口型,不需要 MuseTalk 事后对口型(MuseTalk 是 AI 数字人路线用的)。
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...")
|
||||
video_url = _step_upload(job, video_path)
|
||||
|
||||
job.credits_cost = CREDITS_VIRAL_VIDEO_COST
|
||||
# TODO: 调用 credits.deduct() 实际扣点(#1895 总开关为 false 时不扣,保留 TODO)
|
||||
|
||||
job.mark_completed(video_url)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url})
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.UPLOADING,
|
||||
100.0,
|
||||
"视频生成完成",
|
||||
{"video_url": video_url},
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
|
||||
logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url)
|
||||
return {"ok": True, "job_id": job_id, "video_url": video_url}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线异常: %s", e, exc_info=True)
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
elif job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@shared_task(bind=True, max_retries=1, name="worker.run_video_style_analysis")
|
||||
def run_video_style_analysis(self: Task, job_id: str) -> dict:
|
||||
"""独立的视频风格分析任务(v1.3)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 10.0, "正在分析参考视频风格...")
|
||||
style_guide = _step_video_analysis(job)
|
||||
job.style_guide = style_guide
|
||||
_save_job(repo, job, session)
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 100.0, "风格分析完成", {"style_guide": style_guide})
|
||||
return {"ok": True, "job_id": job_id, "style_guide": style_guide}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 风格分析失败: %s", e)
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
@@ -918,3 +918,72 @@ class GpuWorkerModel(Base):
|
||||
capabilities = Column(String(500), nullable=False, default="") # 逗号分隔,如 "musetalk"
|
||||
last_heartbeat_at = Column(DateTime, nullable=True, index=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoJobModel(Base):
|
||||
"""爆款视频任务"""
|
||||
|
||||
__tablename__ = "viral_video_jobs"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
images = Column(JSON, nullable=False, default=list) # 产品图片 URL 列表
|
||||
industry = Column(String(100), nullable=False, default="")
|
||||
target_customer = Column(String(500), nullable=False, default="")
|
||||
persona_id = Column(String(36), nullable=False, default="")
|
||||
viral_structure = Column(String(50), nullable=False, default="")
|
||||
marketing_purpose = Column(String(100), nullable=False, default="")
|
||||
bgm_preference = Column(String(50), nullable=False, default="")
|
||||
duration = Column(Integer, nullable=False, default=30)
|
||||
user_copy_text = Column(Text, nullable=False, default="")
|
||||
fusion_level = Column(String(20), nullable=False, default="ai_polish")
|
||||
reference_audio_path = Column(String(1000), nullable=False, default="")
|
||||
# v1.3 新增字段
|
||||
reference_video_url = Column(String(1000), nullable=False, default="")
|
||||
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)
|
||||
# 结果与状态
|
||||
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="")
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
started_at = Column(DateTime(timezone=True), nullable=True)
|
||||
completed_at = Column(DateTime(timezone=True), nullable=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoStyleTemplateModel(Base):
|
||||
"""爆款视频风格模板配置表"""
|
||||
|
||||
__tablename__ = "viral_video_style_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
thumbnail_url = Column(String(1000), nullable=False, default="")
|
||||
style_config = Column(JSON, nullable=False, default=dict)
|
||||
is_system = Column(Boolean, nullable=False, default=True, index=True)
|
||||
sort_order = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoPromptTemplateModel(Base):
|
||||
"""爆款视频 Prompt 模板表(由 #2040 seed)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
prompt_type = Column(String(50), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
content = Column(Text, nullable=False, default="")
|
||||
variables = Column(JSON, nullable=False, default=list)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
+202
@@ -0,0 +1,202 @@
|
||||
"""爆款视频任务 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,
|
||||
}
|
||||
@@ -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 = ""
|
||||
|
||||
Executable
+183
@@ -0,0 +1,183 @@
|
||||
"""ViralVideoJob 领域模型 — 爆款视频任务.
|
||||
|
||||
状态机:
|
||||
pending → running → completed
|
||||
↘ failed → pending (retry)
|
||||
↘ cancelled
|
||||
running 中可暂停:running → wait_user_confirm → running (confirm-intent resume)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from enum import StrEnum
|
||||
else:
|
||||
from enum import Enum
|
||||
|
||||
class StrEnum(str, Enum):
|
||||
pass
|
||||
|
||||
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
class ViralVideoStatus(StrEnum):
|
||||
"""爆款视频任务状态枚举。"""
|
||||
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
WAIT_USER_CONFIRM = "wait_user_confirm"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class ViralVideoStage(StrEnum):
|
||||
"""编排流水线阶段枚举(用于 WS 进度推送)。"""
|
||||
|
||||
IMAGE_ANALYSIS = "image_analysis"
|
||||
VIDEO_ANALYSIS = "video_analysis"
|
||||
INTENT_PARSING = "intent_parsing"
|
||||
COPY_FUSION = "copy_fusion"
|
||||
STORYBOARD = "storyboard"
|
||||
REVIEW = "review"
|
||||
TTS = "tts"
|
||||
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"
|
||||
COPY_FUSION = "copy_fusion"
|
||||
STORYBOARD = "storyboard"
|
||||
REVIEW = "review"
|
||||
VIDEO_STYLE_INTEGRATION = "video_style_integration"
|
||||
STYLE_CONSTRAINT = "style_constraint"
|
||||
|
||||
|
||||
CREDITS_VIRAL_VIDEO_COST = 50
|
||||
|
||||
STAGE_LABELS = {
|
||||
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
|
||||
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
|
||||
ViralVideoStage.INTENT_PARSING: "意图解析",
|
||||
ViralVideoStage.COPY_FUSION: "文案融合",
|
||||
ViralVideoStage.STORYBOARD: "分镜脚本",
|
||||
ViralVideoStage.REVIEW: "合规审核",
|
||||
ViralVideoStage.TTS: "AI 配音",
|
||||
ViralVideoStage.BGM_SELECT: "BGM 选择",
|
||||
ViralVideoStage.RENDERING: "视频渲染",
|
||||
ViralVideoStage.MUSETALK: "数字人口型",
|
||||
ViralVideoStage.UPLOADING: "上传发布",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ViralVideoJob:
|
||||
"""爆款视频任务领域实体。"""
|
||||
|
||||
user_id: str
|
||||
images: list[str] = field(default_factory=list)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
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.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜)
|
||||
image_analysis: dict | None = None
|
||||
# 状态
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
status: ViralVideoStatus = ViralVideoStatus.PENDING
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
# ── 状态转换 ──
|
||||
|
||||
def mark_running(self) -> None:
|
||||
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
|
||||
raise ValueError(f"Cannot transition from {self.status} to running")
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.started_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_wait_user_confirm(self, intent_result: dict) -> None:
|
||||
if self.status != ViralVideoStatus.RUNNING:
|
||||
raise ValueError(f"Cannot transition from {self.status} to wait_user_confirm")
|
||||
self.status = ViralVideoStatus.WAIT_USER_CONFIRM
|
||||
self.intent_result = intent_result
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_confirm(self) -> None:
|
||||
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise ValueError(f"Cannot resume from {self.status}")
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_completed(self, video_url: str) -> None:
|
||||
self.status = ViralVideoStatus.COMPLETED
|
||||
self.result_video_url = video_url
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_failed(self, error_msg: str) -> None:
|
||||
self.status = ViralVideoStatus.FAILED
|
||||
self.error_msg = error_msg
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_cancelled(self) -> None:
|
||||
if self.status in (ViralVideoStatus.COMPLETED, ViralVideoStatus.FAILED, ViralVideoStatus.CANCELLED):
|
||||
raise ValueError(f"Cannot cancel task in {self.status} status")
|
||||
self.status = ViralVideoStatus.CANCELLED
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
@property
|
||||
def is_terminal(self) -> bool:
|
||||
return self.status in (
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
ViralVideoStatus.CANCELLED,
|
||||
)
|
||||
Executable
+30
@@ -0,0 +1,30 @@
|
||||
"""爆款视频任务仓储接口。"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
|
||||
class ViralVideoJobRepository(ABC):
|
||||
"""爆款视频任务仓储抽象。"""
|
||||
|
||||
@abstractmethod
|
||||
def save(self, job: ViralVideoJob) -> None:
|
||||
"""保存(新建)任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def update(self, job: ViralVideoJob) -> None:
|
||||
"""更新任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def get(self, job_id: str) -> Optional[ViralVideoJob]:
|
||||
"""按 ID 获取任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def list_by_user(self, user_id: str, limit: int = 50, offset: int = 0) -> list[ViralVideoJob]:
|
||||
"""获取用户的历史任务列表。"""
|
||||
|
||||
@abstractmethod
|
||||
def count_pending_by_user(self, user_id: str) -> int:
|
||||
"""统计用户待处理任务数。"""
|
||||
@@ -13,7 +13,9 @@ API 和 Worker 两边共用。基于火山引擎方舟平台的 OpenAI 兼容接
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
@@ -238,6 +240,144 @@ 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:
|
||||
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
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -233,9 +233,7 @@ def _call_ai_recommend_service(
|
||||
has_analysis = any(aid in asset_analyses for aid in asset_ids[:30])
|
||||
|
||||
# 构建 prompt
|
||||
system_prompt = (
|
||||
"你是一个专业的视频剪辑导演助手。" "根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
)
|
||||
system_prompt = "你是一个专业的视频剪辑导演助手。根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
if has_analysis:
|
||||
system_prompt += (
|
||||
"每个素材附带了 AI 视频理解的内容描述,请根据素材的实际内容来决策编排:\n"
|
||||
@@ -493,3 +491,81 @@ def run_generate_cover(
|
||||
result.get("image_url", "")[:60],
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
# ── 通用 LLM / Vision 调用(#2039 ViralVideoOrchestrator 使用,复用现有豆包客户端)──
|
||||
|
||||
|
||||
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
|
||||
messages = [
|
||||
{"role": "system", "content": "你是专业的短视频内容策划助手。需要结构化输出时请严格使用 JSON。"},
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
raw = client.chat_completion(messages, temperature=temperature, max_tokens=4096)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return raw
|
||||
|
||||
|
||||
def call_vision(image_url: str, prompt: str) -> object:
|
||||
"""调用豆包视觉大模型分析图片,返回解析后的 JSON 或原文字符串;失败返回 None。"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
return None
|
||||
messages = [
|
||||
{"role": "system", "content": "你是专业的视觉分析师。需要结构化输出时请严格使用 JSON。"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": prompt},
|
||||
{"type": "image_url", "image_url": {"url": image_url}},
|
||||
],
|
||||
},
|
||||
]
|
||||
raw = client.chat_completion(messages, temperature=0.3, max_tokens=2048)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
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
|
||||
|
||||
@@ -15,3 +15,8 @@ Pillow==10.4.0
|
||||
|
||||
# FFmpeg Python 绑定
|
||||
ffmpeg-python==0.2.0
|
||||
|
||||
# v1.3 参考视频风格分析(#2051)
|
||||
scenedetect==0.6.4
|
||||
librosa==0.10.2.post1
|
||||
soundfile==0.12.1
|
||||
|
||||
@@ -0,0 +1,481 @@
|
||||
"""#2106 DoubaoClient.video_generation 单测,覆盖 submit/poll/download 主路径和失败分支。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
|
||||
def _make_client(**overrides):
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = overrides.get("api_key", "test-key")
|
||||
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
client.model = "doubao-model"
|
||||
client.vision_model = "doubao-vision"
|
||||
client.timeout = overrides.get("timeout", 30)
|
||||
client.max_retries = overrides.get("max_retries", 0)
|
||||
return client
|
||||
|
||||
|
||||
def _fake_time_factory(base=1000.0, jump_after=2, jump=1e9):
|
||||
"""返回一个 time.time() 替身:前 jump_after 次返回 base+offset,之后返回巨大值让 deadline 立即触发。
|
||||
|
||||
避免 Python logging 内部也调 time.time() 导致 StopIteration。
|
||||
"""
|
||||
state = {"n": 0}
|
||||
|
||||
def _t():
|
||||
n = state["n"]
|
||||
state["n"] += 1
|
||||
if n < jump_after:
|
||||
return base + n
|
||||
return base + jump + n
|
||||
|
||||
return _t
|
||||
|
||||
|
||||
class TestVideoGenerationHappyPath:
|
||||
def test_happy_path_generates_and_downloads(self, tmp_path):
|
||||
client = _make_client()
|
||||
|
||||
fake_task_resp = MagicMock()
|
||||
fake_task_resp.json.return_value = {"id": "task-001"}
|
||||
fake_task_resp.raise_for_status = MagicMock()
|
||||
|
||||
fake_poll_resp = MagicMock()
|
||||
fake_poll_resp.json.return_value = {
|
||||
"status": "succeeded",
|
||||
"content": {"video_url": "https://cdn.example.com/v.mp4"},
|
||||
}
|
||||
fake_poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
class FakeStreamResponse:
|
||||
def __init__(self):
|
||||
self._chunks = [b"FAKE", b"MP4", b"DATA"]
|
||||
self._it = iter(self._chunks)
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
calls = {"post": 0, "get": 0}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls["post"] += 1
|
||||
return fake_task_resp
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
calls["get"] += 1
|
||||
if "/tasks/task-001" in url:
|
||||
return fake_poll_resp
|
||||
raise AssertionError(f"unexpected GET (not stream): {url}")
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "abcd1234"
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStreamResponse()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=2)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_settings,
|
||||
):
|
||||
mock_settings.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
prompt=" 镜头一 ",
|
||||
image_url="https://img/x.jpg",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert Path(out).name == "seedance_task-001_abcd1234.mp4"
|
||||
assert Path(out).read_bytes() == b"FAKEMP4DATA"
|
||||
assert calls["post"] == 1
|
||||
assert calls["get"] == 1
|
||||
|
||||
|
||||
class TestVideoGenerationFailures:
|
||||
def test_returns_none_when_unavailable(self, tmp_path):
|
||||
client = _make_client(api_key="")
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_on_empty_prompt(self, tmp_path):
|
||||
client = _make_client()
|
||||
assert client.video_generation(" ", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_create_returns_no_id(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.json.return_value = {"error": "bad"}
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=fake_resp),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=1, doubao_video_timeout=60, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_poll_returns_failed(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t2"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "failed", "error": {"code": "C1", "message": "bad"}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_download_raises(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t3"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn/v.mp4"}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
class BadStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
raise RuntimeError("network down")
|
||||
|
||||
def iter_bytes(self, **kw):
|
||||
return iter([])
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=BadStream()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
class TestVideoGenerationRetryAndPoll:
|
||||
def test_create_retries_then_succeeds(self, tmp_path):
|
||||
client = _make_client(max_retries=1)
|
||||
|
||||
ok_resp = MagicMock()
|
||||
ok_resp.json.return_value = {"id": "t-retry"}
|
||||
ok_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "expired"}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
calls = {"post": 0}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls["post"] += 1
|
||||
if calls["post"] == 1:
|
||||
raise httpx.HTTPError("network")
|
||||
return ok_resp
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx") as mock_httpx,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
mock_httpx.HTTPError = httpx.HTTPError
|
||||
mock_httpx.post.side_effect = fake_post
|
||||
mock_httpx.get.return_value = poll_resp
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
assert calls["post"] == 2
|
||||
|
||||
def test_succeeded_but_no_video_url_returns_none(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t-nourl"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "succeeded", "content": {}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
class TestAiServiceCallVideoGeneration:
|
||||
def test_returns_none_on_exception(self):
|
||||
from packages.shared import ai_service
|
||||
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.side_effect = RuntimeError("boom")
|
||||
mock_get.return_value = mock_client
|
||||
assert ai_service.call_video_generation("p") is None
|
||||
|
||||
|
||||
class TestVideoGenerationPollLoop:
|
||||
def test_poll_queued_then_running_then_succeeded(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t-wait"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
|
||||
queued = MagicMock(json=MagicMock(return_value={"status": "queued"}))
|
||||
queued.raise_for_status = MagicMock()
|
||||
running = MagicMock(json=MagicMock(return_value={"status": "running"}))
|
||||
running.raise_for_status = MagicMock()
|
||||
ok = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/x.mp4"}})
|
||||
)
|
||||
ok.raise_for_status = MagicMock()
|
||||
poll_seq = [queued, running, ok]
|
||||
|
||||
class EmptyChunkStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b""
|
||||
yield b"D"
|
||||
yield b""
|
||||
yield b"ATA"
|
||||
|
||||
get_calls = {"n": 0}
|
||||
|
||||
def fake_get(url, **kw):
|
||||
if "/tasks/t-wait" in url:
|
||||
resp = poll_seq[min(get_calls["n"], len(poll_seq) - 1)]
|
||||
get_calls["n"] += 1
|
||||
return resp
|
||||
raise AssertionError(url)
|
||||
|
||||
sleeps = []
|
||||
# jump_after 要足够大:deadline 计算一次 + 3次 while 条件判断 = 4 次
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=EmptyChunkStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", side_effect=lambda s: sleeps.append(s)),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=5, jump=1)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="ef012345")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=200, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None
|
||||
assert Path(out).read_bytes() == b"DATA"
|
||||
# queued 和 running 各 sleep 一次
|
||||
assert len(sleeps) >= 2
|
||||
|
||||
def test_poll_exception_does_not_crash(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-err"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
ok = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/e.mp4"}})
|
||||
)
|
||||
ok.raise_for_status = MagicMock()
|
||||
|
||||
class OkStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b"OK"
|
||||
|
||||
poll_calls = {"n": 0}
|
||||
|
||||
def fake_get(url, **kw):
|
||||
poll_calls["n"] += 1
|
||||
if poll_calls["n"] == 1:
|
||||
raise httpx.HTTPError("transient")
|
||||
return ok
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=OkStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="11111111")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=100, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert poll_calls["n"] == 2
|
||||
|
||||
def test_default_output_dir_and_audio_watermark(self, tmp_path, monkeypatch):
|
||||
"""不传 output_dir 时落到 /tmp;generate_audio/watermark=True 也能正常提交。"""
|
||||
client = _make_client()
|
||||
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-default"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/d.mp4"}})
|
||||
)
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
# 用 tmp_path 伪造 /tmp 避免污染真 /tmp
|
||||
monkeypatch.setattr("packages.shared.ai_client.os.makedirs", lambda d, exist_ok=True: None)
|
||||
|
||||
class S:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b"D"
|
||||
|
||||
# 捕获 POST payload 断言
|
||||
captured = {}
|
||||
|
||||
def fake_post(url, **kw):
|
||||
captured["json"] = kw.get("json")
|
||||
return create_resp
|
||||
|
||||
def fake_get(url, **kw):
|
||||
return poll_resp
|
||||
|
||||
def fake_open(path, mode):
|
||||
# 返回一个 MagicMock file,模拟写入
|
||||
f = MagicMock()
|
||||
f.__enter__ = MagicMock(return_value=f)
|
||||
f.__exit__ = MagicMock(return_value=False)
|
||||
captured["path"] = path
|
||||
return f
|
||||
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.post", fake_post)
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.get", fake_get)
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.stream", lambda *a, **kw: S())
|
||||
monkeypatch.setattr("builtins.open", fake_open)
|
||||
monkeypatch.setattr("packages.shared.ai_client.os.path.getsize", lambda p: 99)
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="00000001")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="seedance",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"p", duration=3, ratio="1:1", resolution="480p", generate_audio=True, watermark=True
|
||||
)
|
||||
assert out is not None
|
||||
assert "/tmp/seedance_t-default_00000001.mp4" in out
|
||||
assert captured["json"]["generate_audio"] is True
|
||||
assert captured["json"]["watermark"] is True
|
||||
assert captured["json"]["ratio"] == "1:1"
|
||||
assert captured["json"]["resolution"] == "480p"
|
||||
|
||||
|
||||
class TestGetDoubaoClientSingleton:
|
||||
def test_singleton_lazy_init(self):
|
||||
from packages.shared import ai_client
|
||||
|
||||
prev = ai_client._client
|
||||
try:
|
||||
ai_client._client = None
|
||||
c1 = ai_client.get_doubao_client()
|
||||
c2 = ai_client.get_doubao_client()
|
||||
assert c1 is c2
|
||||
assert isinstance(c1, ai_client.DoubaoClient)
|
||||
finally:
|
||||
ai_client._client = prev
|
||||
|
||||
|
||||
class TestVideoGenerationCancelled:
|
||||
def test_poll_cancelled_returns_none(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-can"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock(json=MagicMock(return_value={"status": "cancelled"}))
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
@@ -0,0 +1,360 @@
|
||||
"""video_analyzer(#2051)单元测试。
|
||||
|
||||
覆盖:
|
||||
- 映射函数(运镜→ken_burns、转场→xfade、色调→video_filter、BPM→BGM)
|
||||
- schema 常量与导出
|
||||
- Farneback 光流运镜判定(合成光流场)
|
||||
- BPM 档位映射
|
||||
- 降级路径(ffmpeg/cv2/librosa 不可用)
|
||||
- 临时目录清理
|
||||
- analyze_video_style 入口在无素材时返回最小 style_guide 不抛
|
||||
- build_render_params_for_clip 聚合输出
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import ( # noqa: E402
|
||||
COLOR_FILTER_PRESETS,
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
TRANSITION_TO_XFADE,
|
||||
ShotBoundary,
|
||||
_detect_camera_movement,
|
||||
_pace_from_bpm,
|
||||
_rule_based_style_guide,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
# ── 映射函数 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_constants_exported():
|
||||
assert MAX_REFERENCE_DURATION_SEC == 60
|
||||
assert MAX_REFERENCE_SIZE_MB == 100
|
||||
assert DEFAULT_ANALYSIS_TIMEOUT == 60
|
||||
assert "style_name" in STYLE_GUIDE_SCHEMA
|
||||
assert "ken_burns_params" in STYLE_GUIDE_SCHEMA
|
||||
assert "video_filter_eq_params" in STYLE_GUIDE_SCHEMA
|
||||
|
||||
|
||||
# ── 运镜→ken_burns ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"movement,expected_type",
|
||||
[
|
||||
("static", "static"),
|
||||
("push_in", "zoom"),
|
||||
("zoom_in", "zoom"),
|
||||
("pull_out", "zoom"),
|
||||
("pan_left", "pan"),
|
||||
("pan_right", "pan"),
|
||||
("tilt_up", "pan+zoom"),
|
||||
("track_left", "pan"),
|
||||
],
|
||||
)
|
||||
def test_map_camera_to_ken_burns_types(movement, expected_type):
|
||||
kb = map_camera_to_ken_burns(movement)
|
||||
assert kb["type"] == expected_type
|
||||
# zoom/pan 类必须有 zoom_start/zoom_end
|
||||
assert 0.8 <= kb["zoom_start"] <= 1.3
|
||||
assert 0.8 <= kb["zoom_end"] <= 1.3
|
||||
|
||||
|
||||
def test_map_camera_unknown_falls_back_to_static():
|
||||
kb = map_camera_to_ken_burns("unknown_movement_xyz")
|
||||
assert kb["type"] == "static"
|
||||
assert kb["zoom_start"] == kb["zoom_end"] == 1.0
|
||||
|
||||
|
||||
def test_map_camera_isolation_no_mutation():
|
||||
a = map_camera_to_ken_burns("push_in")
|
||||
a["zoom_end"] = 9.99
|
||||
b = map_camera_to_ken_burns("push_in")
|
||||
assert b["zoom_end"] != 9.99
|
||||
|
||||
|
||||
# ── 转场→xfade ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ttype,expected",
|
||||
[
|
||||
("hard_cut", "cut"),
|
||||
("cross_dissolve", "dissolve"),
|
||||
("fade", "fade"),
|
||||
("fade_black", "fadeblack"),
|
||||
("zoom_whip", "zoom"),
|
||||
("wipe_left", "wipeleft"),
|
||||
("slide_right", "slideleft"),
|
||||
],
|
||||
)
|
||||
def test_map_transition(ttype, expected):
|
||||
assert map_transition_to_xfade(ttype) == expected
|
||||
|
||||
|
||||
def test_map_transition_unknown_falls_back_to_cut():
|
||||
assert map_transition_to_xfade("some_random_transition") == "cut"
|
||||
|
||||
|
||||
# ── 色调→video_filter ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name", ["none", "warm_vintage", "cool_fresh", "high_contrast", "soft_pastel", "dramatic_cinematic"]
|
||||
)
|
||||
def test_map_color_presets_available(name):
|
||||
p = map_color_to_video_filter(name)
|
||||
assert isinstance(p, dict)
|
||||
# 所有预设必须能被 FFmpeg eq/colorchannelmixer 消费:eq 是 dict,ccm 是 dict
|
||||
assert "eq" in p or p == {} or "colorchannelmixer" in p
|
||||
|
||||
|
||||
def test_map_color_unknown_is_none_preset():
|
||||
p = map_color_to_video_filter("not_a_real_filter")
|
||||
assert p == {}
|
||||
|
||||
|
||||
def test_map_color_isolation():
|
||||
a = map_color_to_video_filter("warm_vintage")
|
||||
a["eq"]["brightness"] = 9.99
|
||||
b = map_color_to_video_filter("warm_vintage")
|
||||
assert b["eq"]["brightness"] != 9.99
|
||||
|
||||
|
||||
# ── BPM → BGM ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bpm,expected",
|
||||
[
|
||||
(0, 90),
|
||||
(70, 70),
|
||||
(120, 120),
|
||||
(200, 180),
|
||||
(30, 60),
|
||||
],
|
||||
)
|
||||
def test_map_bgm_bpm(bpm, expected):
|
||||
assert map_bgm_bpm(bpm) == expected
|
||||
|
||||
|
||||
def test_pace_from_bpm_buckets():
|
||||
assert _pace_from_bpm(120) == "fast_cut"
|
||||
assert _pace_from_bpm(110) == "fast_cut"
|
||||
assert _pace_from_bpm(90) == "medium"
|
||||
assert _pace_from_bpm(80) == "medium"
|
||||
assert _pace_from_bpm(60) == "slow_cinematic"
|
||||
assert _pace_from_bpm(0) == "medium"
|
||||
|
||||
|
||||
# ── Farneback 光流→运镜(合成光流) ───────────────────────────────────
|
||||
|
||||
|
||||
def _make_flow(dx: float, dy: float, w: int = 60, h: int = 40, zoom: float = 0.0):
|
||||
"""构造一个合成光流场:整体平移(dx,dy)+径向发散(zoom>0=zoom in,<0=out)。"""
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
fx = dx + (xs - cx) * zoom
|
||||
fy = dy + (ys - cy) * zoom
|
||||
return np.stack([fx, fy], axis=-1).astype(np.float32)
|
||||
|
||||
|
||||
def test_detect_movement_static():
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "static"
|
||||
assert i == "low"
|
||||
|
||||
|
||||
def test_detect_movement_pan_right():
|
||||
flow = _make_flow(2.0, 0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_right"
|
||||
assert i in ("medium", "high")
|
||||
|
||||
|
||||
def test_detect_movement_pan_left():
|
||||
flow = _make_flow(-2.0, 0.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_left"
|
||||
|
||||
|
||||
def test_detect_movement_tilt_down():
|
||||
flow = _make_flow(0.0, 2.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "tilt_down"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_in_radial():
|
||||
# 径向向外发散 = zoom in
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_in"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_out_radial():
|
||||
flow = _make_flow(0.0, 0.0, zoom=-0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_out"
|
||||
|
||||
|
||||
# ── 规则合成 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_shots(n=4):
|
||||
return [
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=float(i * 3),
|
||||
end_sec=float((i + 1) * 3),
|
||||
movement=["static", "push_in", "pan_left", "zoom_in"][i],
|
||||
intensity=["low", "medium", "low", "high"][i],
|
||||
transition="hard_cut",
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
|
||||
CAMERA_TO_KEN_BURNS_DIRS = {
|
||||
"zoom_in_slow",
|
||||
"zoom_out_slow",
|
||||
"pan_left_slow",
|
||||
"pan_right_slow",
|
||||
"zoom_in_medium",
|
||||
"zoom_out_medium",
|
||||
"diagonal_push",
|
||||
"static",
|
||||
}
|
||||
|
||||
|
||||
def test_rule_based_style_guide_structure():
|
||||
shots = _sample_shots()
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={"color_filter": "warm_vintage"})
|
||||
# 关键字段存在且类型正确
|
||||
assert sg["shot_count"] == 4
|
||||
assert sg["pace"] == "fast_cut"
|
||||
assert sg["bpm"] == 120
|
||||
assert sg["avg_shot_duration"] == 3.0
|
||||
assert len(sg["camera_movements"]) == 4
|
||||
assert sg["color_filter"] == "warm_vintage"
|
||||
assert "eq" in sg["video_filter_eq_params"]
|
||||
assert "default" in sg["ken_burns_params"]
|
||||
assert isinstance(sg["transition_map"], dict)
|
||||
assert isinstance(sg["ken_burns_direction_hint"], str) and sg["ken_burns_direction_hint"]
|
||||
|
||||
|
||||
# ── build_render_params_for_clip 聚合 ─────────────────────────────────
|
||||
|
||||
|
||||
def test_build_render_params_for_clip_shape():
|
||||
sg = _rule_based_style_guide(_sample_shots(), bpm=95, vlm={"color_filter": "cool_fresh"})
|
||||
p0 = build_render_params_for_clip(0, sg, duration_sec=3.0)
|
||||
assert "ken_burns" in p0
|
||||
assert "transition" in p0
|
||||
assert "video_filter" in p0
|
||||
assert p0["bgm_bpm_hint"] == 95
|
||||
assert p0["duration_sec"] == 3.0
|
||||
# clip 1 是 push_in → zoom
|
||||
p1 = build_render_params_for_clip(1, sg)
|
||||
assert p1["ken_burns"]["type"] == "zoom"
|
||||
|
||||
|
||||
def test_build_render_params_high_intensity_amplifies():
|
||||
shots = _sample_shots() # shot 3 = zoom_in/high
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={})
|
||||
p3 = build_render_params_for_clip(3, sg)
|
||||
base = map_camera_to_ken_burns("zoom_in")
|
||||
assert p3["ken_burns"]["zoom_end"] > base["zoom_end"]
|
||||
|
||||
|
||||
# ── 降级与容错 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_analyze_with_nonexistent_file_returns_minimum_guide():
|
||||
sg = analyze_video_style("/nonexistent/path/fake_video.mp4")
|
||||
assert isinstance(sg, dict)
|
||||
assert "style_name" in sg
|
||||
assert sg["shot_count"] == 0
|
||||
# 不抛异常且字段完整
|
||||
|
||||
|
||||
def test_analyze_invalid_style_strength_defaults_to_medium():
|
||||
# 即使视频不存在,也应被规范化为 medium 并写入返回值
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("fake", style_strength="banana")
|
||||
assert sg.get("style_strength", "medium") == "medium"
|
||||
|
||||
|
||||
def test_temp_dir_cleaned_up_after_run():
|
||||
"""用临时真实空文件模拟本地路径,确认 frames 临时目录被清理。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "fake.mp4"
|
||||
fake.write_bytes(b"")
|
||||
# 抽帧会失败(ffmpeg 对空文件失败),但应全程不抛且临时目录 rmtree
|
||||
# 直接 mock _ensure_local_video 回传不存在的文件,走 _probe_duration=0 降级路径
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("proto://fake", style_strength="light")
|
||||
assert "style_name" in sg
|
||||
|
||||
|
||||
def test_ffmpeg_failure_falls_back_to_vlm_only_path():
|
||||
"""模拟 ffmpeg 抽帧失败,仍能返回 style_guide。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "ref.mp4"
|
||||
fake.write_bytes(b"not a real video")
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._extract_keyframes", side_effect=RuntimeError("ffmpeg exploded")
|
||||
):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_shots") as mock_shots:
|
||||
mock_shots.return_value = [ShotBoundary(0, 0.0, 3.0)]
|
||||
with patch("apps.worker.viral_video.video_analyzer._analyze_movements"):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_bpm", return_value=90):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._vlm_analyze_frames",
|
||||
return_value={"color_filter": "none"},
|
||||
):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._llm_synthesize",
|
||||
side_effect=lambda shots, bpm, vlm, ss: _rule_based_style_guide(shots, bpm, vlm),
|
||||
):
|
||||
sg = analyze_video_style(str(fake))
|
||||
assert sg["bpm"] == 90
|
||||
assert sg["shot_count"] == 1
|
||||
|
||||
|
||||
# ── 转场映射完整性 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_transition_map_covers_observed_types():
|
||||
for t in ("hard_cut", "cross_dissolve", "fade_black", "zoom_whip"):
|
||||
assert t in TRANSITION_TO_XFADE
|
||||
|
||||
|
||||
# ── 预设完整性 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_color_filter_preset_keys_are_safe_for_ffmpeg():
|
||||
for name, preset in COLOR_FILTER_PRESETS.items():
|
||||
if preset == {}:
|
||||
continue
|
||||
# eq 所有值都是数字
|
||||
for k, v in preset.get("eq", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.eq.{k} not numeric"
|
||||
for k, v in preset.get("colorchannelmixer", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.ccm.{k} not numeric"
|
||||
Executable
+517
@@ -0,0 +1,517 @@
|
||||
"""爆款视频模块单元测试。
|
||||
|
||||
覆盖范围:
|
||||
- 领域实体状态机转换
|
||||
- Repository CRUD
|
||||
- API 端点(6 个)
|
||||
- Celery 编排器流水线
|
||||
- Schema 校验
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
FusionLevel,
|
||||
StyleStrength,
|
||||
ViralVideoJob,
|
||||
ViralVideoStage,
|
||||
ViralVideoStatus,
|
||||
)
|
||||
|
||||
# ── 领域模型测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoStatus:
|
||||
"""状态枚举测试。"""
|
||||
|
||||
def test_status_values(self):
|
||||
assert ViralVideoStatus.PENDING == "pending"
|
||||
assert ViralVideoStatus.RUNNING == "running"
|
||||
assert ViralVideoStatus.WAIT_USER_CONFIRM == "wait_user_confirm"
|
||||
assert ViralVideoStatus.COMPLETED == "completed"
|
||||
assert ViralVideoStatus.FAILED == "failed"
|
||||
assert ViralVideoStatus.CANCELLED == "cancelled"
|
||||
|
||||
def test_terminal_statuses(self):
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED).is_terminal
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.FAILED).is_terminal
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.CANCELLED).is_terminal
|
||||
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING).is_terminal
|
||||
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING).is_terminal
|
||||
|
||||
|
||||
class TestViralVideoJobStateTransitions:
|
||||
"""状态机转换测试。"""
|
||||
|
||||
def test_mark_running_from_pending(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
job.mark_running()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
assert job.started_at is not None
|
||||
|
||||
def test_mark_running_from_running(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_running()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
|
||||
def test_mark_running_from_completed_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
|
||||
with pytest.raises(ValueError, match="Cannot transition"):
|
||||
job.mark_running()
|
||||
|
||||
def test_mark_wait_user_confirm(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
intent = {"intent": "推广", "key_messages": ["卖点1"]}
|
||||
job.mark_wait_user_confirm(intent)
|
||||
assert job.status == ViralVideoStatus.WAIT_USER_CONFIRM
|
||||
assert job.intent_result == intent
|
||||
|
||||
def test_mark_wait_user_confirm_from_non_running_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING)
|
||||
with pytest.raises(ValueError, match="Cannot transition"):
|
||||
job.mark_wait_user_confirm({})
|
||||
|
||||
def test_resume_from_confirm(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
job.resume_from_confirm()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
|
||||
def test_resume_from_non_confirm_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
with pytest.raises(ValueError, match="Cannot resume"):
|
||||
job.resume_from_confirm()
|
||||
|
||||
def test_mark_completed(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_completed("https://oss.example.com/video.mp4")
|
||||
assert job.status == ViralVideoStatus.COMPLETED
|
||||
assert job.result_video_url == "https://oss.example.com/video.mp4"
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_mark_failed(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_failed("渲染超时")
|
||||
assert job.status == ViralVideoStatus.FAILED
|
||||
assert job.error_msg == "渲染超时"
|
||||
|
||||
def test_mark_cancelled(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_cancelled()
|
||||
assert job.status == ViralVideoStatus.CANCELLED
|
||||
|
||||
def test_mark_cancelled_from_terminal_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
|
||||
with pytest.raises(ValueError, match="Cannot cancel"):
|
||||
job.mark_cancelled()
|
||||
|
||||
|
||||
class TestViralVideoJobDefaults:
|
||||
"""默认值测试。"""
|
||||
|
||||
def test_default_values(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
assert job.images == []
|
||||
assert job.industry == ""
|
||||
assert job.duration == 30
|
||||
assert job.fusion_level == FusionLevel.AI_POLISH
|
||||
assert job.style_strength == StyleStrength.MEDIUM
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.credits_cost == 0
|
||||
assert job.retry_count == 0
|
||||
assert job.result_video_url == ""
|
||||
assert job.error_msg == ""
|
||||
|
||||
def test_credits_cost_constant(self):
|
||||
assert CREDITS_VIRAL_VIDEO_COST == 50
|
||||
|
||||
|
||||
class TestViralVideoStage:
|
||||
"""阶段枚举测试。"""
|
||||
|
||||
def test_all_stages_have_labels(self):
|
||||
for stage in ViralVideoStage:
|
||||
assert stage in STAGE_LABELS, f"Stage {stage} missing label"
|
||||
|
||||
def test_stage_order(self):
|
||||
expected_order = [
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"review",
|
||||
"tts",
|
||||
"bgm_select",
|
||||
"rendering",
|
||||
"musetalk",
|
||||
"uploading",
|
||||
]
|
||||
actual_order = [s.value for s in ViralVideoStage]
|
||||
assert actual_order == expected_order
|
||||
|
||||
|
||||
# ── Schema 校验测试 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoSchemas:
|
||||
"""Pydantic Schema 校验测试。"""
|
||||
|
||||
def test_create_request_valid(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
req = CreateViralVideoRequest(images=["https://example.com/img.jpg"])
|
||||
assert req.images == ["https://example.com/img.jpg"]
|
||||
assert req.fusion_level == "ai_polish"
|
||||
assert req.style_strength == "medium"
|
||||
assert req.duration == 30
|
||||
|
||||
def test_create_request_empty_images_raises(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(images=[])
|
||||
|
||||
def test_create_request_invalid_fusion_level(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(
|
||||
images=["https://example.com/img.jpg"],
|
||||
fusion_level="invalid_level",
|
||||
)
|
||||
|
||||
def test_create_request_invalid_style_strength(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(
|
||||
images=["https://example.com/img.jpg"],
|
||||
style_strength="ultra",
|
||||
)
|
||||
|
||||
def test_confirm_intent_request_defaults(self):
|
||||
from app.schemas.viral_video import ConfirmIntentRequest
|
||||
|
||||
req = ConfirmIntentRequest()
|
||||
assert req.confirmed_copy == ""
|
||||
assert req.adjustments == ""
|
||||
|
||||
def test_analyze_style_request(self):
|
||||
from app.schemas.viral_video import AnalyzeStyleRequest
|
||||
|
||||
req = AnalyzeStyleRequest(reference_video_url="https://example.com/video.mp4")
|
||||
assert req.reference_video_url == "https://example.com/video.mp4"
|
||||
|
||||
def test_ws_progress_event(self):
|
||||
from app.schemas.viral_video import WSProgressEvent
|
||||
|
||||
event = WSProgressEvent(
|
||||
job_id="abc123",
|
||||
stage="image_analysis",
|
||||
progress=10.0,
|
||||
message="正在分析图片",
|
||||
)
|
||||
assert event.type == "viral_video:progress"
|
||||
assert event.job_id == "abc123"
|
||||
assert event.progress == 10.0
|
||||
|
||||
|
||||
# ── Repository 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoRepository:
|
||||
"""SQLAlchemy Repository CRUD 测试(使用内存数据库)。"""
|
||||
|
||||
@pytest.fixture
|
||||
def db_session(self):
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
SessionLocal = sessionmaker(bind=engine)
|
||||
session = SessionLocal()
|
||||
yield session
|
||||
session.close()
|
||||
|
||||
def test_save_and_get(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
job = ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
repo.save(job)
|
||||
|
||||
fetched = repo.get(job.id)
|
||||
assert fetched is not None
|
||||
assert fetched.id == job.id
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_get_nonexistent(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
assert repo.get("nonexistent-id") is None
|
||||
|
||||
def test_list_by_user(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
for i in range(3):
|
||||
job = ViralVideoJob(user_id="user-001", industry=f"行业{i}")
|
||||
repo.save(job)
|
||||
|
||||
# 另一个用户的任务
|
||||
other_job = ViralVideoJob(user_id="user-002", industry="其他")
|
||||
repo.save(other_job)
|
||||
|
||||
jobs = repo.list_by_user("user-001")
|
||||
assert len(jobs) == 3
|
||||
assert all(j.user_id == "user-001" for j in jobs)
|
||||
|
||||
def test_update_status(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
job = ViralVideoJob(user_id="user-001")
|
||||
repo.save(job)
|
||||
|
||||
job.mark_running()
|
||||
repo.update(job)
|
||||
|
||||
fetched = repo.get(job.id)
|
||||
assert fetched.status == ViralVideoStatus.RUNNING
|
||||
assert fetched.started_at is not None
|
||||
|
||||
def test_count_pending_by_user(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
# 2 个 pending
|
||||
for _ in range(2):
|
||||
repo.save(ViralVideoJob(user_id="user-001"))
|
||||
# 1 个 completed
|
||||
completed = ViralVideoJob(user_id="user-001", status=ViralVideoStatus.COMPLETED)
|
||||
repo.save(completed)
|
||||
|
||||
assert repo.count_pending_by_user("user-001") == 2
|
||||
|
||||
def test_style_template_repo(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.models import ViralVideoStyleTemplateModel
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
|
||||
# 插入模板
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
description="适合快消品",
|
||||
style_config={"cut_speed": "fast"},
|
||||
is_system=True,
|
||||
sort_order=1,
|
||||
)
|
||||
db_session.add(tpl)
|
||||
db_session.commit()
|
||||
|
||||
repo = SQLAlchemyViralVideoStyleTemplateRepository(db_session)
|
||||
templates = repo.list_all()
|
||||
assert len(templates) == 1
|
||||
assert templates[0]["name"] == "快节奏"
|
||||
|
||||
fetched = repo.get("tpl-001")
|
||||
assert fetched is not None
|
||||
assert fetched["style_config"] == {"cut_speed": "fast"}
|
||||
|
||||
|
||||
# ── Celery 编排器测试 ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoPipeline:
|
||||
"""编排器流水线测试。"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_job(self):
|
||||
return ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg", "https://img.com/2.jpg"],
|
||||
industry="美妆",
|
||||
target_customer="年轻女性",
|
||||
marketing_purpose="品牌推广",
|
||||
duration=30,
|
||||
user_copy_text="这款产品超好用",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
|
||||
@patch("packages.shared.ai_service.call_vision")
|
||||
def test_image_analysis_step(self, mock_vision, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||
|
||||
mock_vision.return_value = {"name": "口红", "features": ["持久", "滋润"]}
|
||||
result = _step_image_analysis(mock_job)
|
||||
assert "products" in result
|
||||
assert len(result["products"]) == 2 # 两张图片
|
||||
|
||||
@patch("packages.shared.ai_service.call_vision")
|
||||
def test_image_analysis_fallback(self, mock_vision, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||
|
||||
# 模拟 call_vision 不存在
|
||||
mock_vision.side_effect = ImportError("no module")
|
||||
result = _step_image_analysis(mock_job)
|
||||
assert "products" in result
|
||||
|
||||
def test_video_analysis_no_reference(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
# 没有参考视频
|
||||
mock_job.reference_video_url = ""
|
||||
result = _step_video_analysis(mock_job)
|
||||
assert result is None
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_intent_parsing(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_intent_parsing
|
||||
|
||||
mock_llm.return_value = {"intent": "推广口红", "tone": "活泼"}
|
||||
result = _step_intent_parsing(mock_job, {"products": []})
|
||||
assert "intent" in result
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
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 = "融合后的文案内容"
|
||||
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_storyboard_generation(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
|
||||
|
||||
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(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": {}}
|
||||
result = _step_review(mock_job, "测试文案", [])
|
||||
assert result["passed"] is True
|
||||
|
||||
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 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:
|
||||
"""流水线端到端集成测试(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._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_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")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._emit_progress")
|
||||
def test_resume_pipeline_completes(
|
||||
self,
|
||||
mock_emit,
|
||||
mock_get_repo,
|
||||
mock_img_analysis,
|
||||
mock_video_analysis,
|
||||
mock_intent,
|
||||
mock_copy_fusion,
|
||||
mock_storyboard,
|
||||
mock_review,
|
||||
mock_tts,
|
||||
mock_bgm,
|
||||
mock_render,
|
||||
mock_upload,
|
||||
):
|
||||
"""测试 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": "推广"},
|
||||
)
|
||||
|
||||
mock_repo = MagicMock()
|
||||
mock_session = MagicMock()
|
||||
mock_get_repo.return_value = (mock_session, mock_repo, job)
|
||||
|
||||
# 设置各步骤返回值
|
||||
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 # 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"
|
||||
|
||||
result = resume_viral_video_pipeline.run("job-001")
|
||||
|
||||
assert result["ok"] is True
|
||||
assert result["video_url"] == "https://oss.example.com/final.mp4"
|
||||
assert job.status == ViralVideoStatus.COMPLETED
|
||||
assert job.credits_cost == CREDITS_VIRAL_VIDEO_COST
|
||||
@@ -0,0 +1,99 @@
|
||||
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
|
||||
|
||||
验证:
|
||||
- 3 张新表可在内存 SQLite 上创建
|
||||
- 默认值与基本 CRUD 正常
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
Base,
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
|
||||
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
|
||||
|
||||
class TestViralVideoJobModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(
|
||||
id="job-001",
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_default_values(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(id="job-002", user_id="user-002")
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoJobModel, "job-002")
|
||||
assert fetched.images == []
|
||||
assert fetched.fusion_level == "ai_polish"
|
||||
assert fetched.style_strength == "medium"
|
||||
assert fetched.status == "pending"
|
||||
assert fetched.credits_cost == 0
|
||||
assert fetched.retry_count == 0
|
||||
assert fetched.style_guide is None
|
||||
assert fetched.intent_result is None
|
||||
|
||||
|
||||
class TestViralVideoStyleTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
style_config={"cut_speed": "fast"},
|
||||
sort_order=1,
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
|
||||
assert fetched.name == "快节奏"
|
||||
assert fetched.style_config == {"cut_speed": "fast"}
|
||||
assert fetched.sort_order == 1
|
||||
|
||||
|
||||
class TestViralVideoPromptTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoPromptTemplateModel(
|
||||
id="pt-001",
|
||||
prompt_type="image_analysis",
|
||||
name="图片分析模板",
|
||||
content="请分析图片:{image_url}",
|
||||
variables=["image_url"],
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
|
||||
assert fetched.prompt_type == "image_analysis"
|
||||
assert fetched.content == "请分析图片:{image_url}"
|
||||
assert fetched.variables == ["image_url"]
|
||||
assert fetched.version == 1
|
||||
assert fetched.is_active is True
|
||||
@@ -0,0 +1,230 @@
|
||||
"""#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
|
||||
|
||||
|
||||
# ── 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
|
||||
@@ -0,0 +1,176 @@
|
||||
"""viral_video.py HTTP 端点单元测试(celery send_task 分支覆盖)。
|
||||
|
||||
直接调用路由函数(不启动 TestClient),通过 patch 注入 repo/session/user,
|
||||
覆盖 4 个 celery_app.send_task(...) 调用点:
|
||||
|
||||
- create_viral_video (generate) -> worker.run_viral_video_pipeline
|
||||
- retry_viral_video_job (retry) -> worker.run_viral_video_pipeline
|
||||
- confirm_intent -> worker.resume_viral_video_pipeline
|
||||
- analyze_style -> worker.run_video_style_analysis
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _auth_user(uid: str = "u1"):
|
||||
return SimpleNamespace(user=SimpleNamespace(id=uid))
|
||||
|
||||
|
||||
def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending", **kwargs):
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
job = MagicMock()
|
||||
job.id = job_id
|
||||
job.user_id = user_id
|
||||
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||
job.images = kwargs.pop("images", ["img-1"])
|
||||
job.industry = kwargs.pop("industry", "电商")
|
||||
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
||||
for k, v in {
|
||||
"persona_id": "",
|
||||
"viral_structure": "",
|
||||
"marketing_purpose": "",
|
||||
"bgm_preference": "",
|
||||
"duration": 30,
|
||||
"user_copy_text": "",
|
||||
"fusion_level": "ai_polish",
|
||||
"reference_audio_path": "",
|
||||
"reference_video_url": "",
|
||||
"style_strength": "medium",
|
||||
"style_template_id": "",
|
||||
"retry_count": 0,
|
||||
"error_msg": "",
|
||||
"result_video_url": "",
|
||||
"style_guide": None,
|
||||
"created_at": None,
|
||||
"started_at": None,
|
||||
"completed_at": None,
|
||||
"stage": "",
|
||||
"progress": 0.0,
|
||||
"intent_result": None,
|
||||
"updated_at": None,
|
||||
}.items():
|
||||
setattr(job, k, kwargs.pop(k, v))
|
||||
return job
|
||||
|
||||
|
||||
# ── generate ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateViralVideo:
|
||||
def _req(self, **kw):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
d = {"images": ["https://x.com/a.jpg"], "industry": "电商", "target_customer": "年轻人"}
|
||||
d.update(kw)
|
||||
return CreateViralVideoRequest(**d)
|
||||
|
||||
def test_generate_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
req = self._req()
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
saved_job = _make_job(job_id="job-new", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
|
||||
def fake_save(job):
|
||||
job.id = saved_job.id
|
||||
|
||||
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.create_viral_video(req, authenticated_user=user, session=session)
|
||||
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=[saved_job.id])
|
||||
assert resp.id == saved_job.id
|
||||
|
||||
|
||||
# ── retry ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRetryViralVideo:
|
||||
def test_retry_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-retry", user_id="u1", status=ViralVideoStatus.FAILED, retry_count=1)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
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.retry_viral_video_job("job-retry", authenticated_user=user, session=session)
|
||||
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.retry_count == 2
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=["job-retry"])
|
||||
assert resp.id == "job-retry"
|
||||
|
||||
|
||||
# ── confirm-intent ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConfirmIntent:
|
||||
def test_confirm_intent_dispatches_resume_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmIntentRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-cfm", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmIntentRequest(confirmed_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_intent("job-cfm", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.user_copy_text == "确认后的文案"
|
||||
job.resume_from_confirm.assert_called_once()
|
||||
mock_send.assert_called_once_with("worker.resume_viral_video_pipeline", args=["job-cfm"])
|
||||
assert resp.id == "job-cfm"
|
||||
|
||||
|
||||
# ── analyze-style ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAnalyzeStyle:
|
||||
def test_analyze_style_dispatches_analysis_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import AnalyzeStyleRequest
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-sty", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = AnalyzeStyleRequest(reference_video_url="https://x.com/ref.mp4", style_template_id="tpl-1")
|
||||
|
||||
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_style("job-sty", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.reference_video_url == "https://x.com/ref.mp4"
|
||||
assert job.style_template_id == "tpl-1"
|
||||
mock_send.assert_called_once_with("worker.run_video_style_analysis", args=["job-sty"])
|
||||
assert resp.job_id == "job-sty"
|
||||
assert resp.status == "analyzing"
|
||||
@@ -0,0 +1,432 @@
|
||||
"""Unit tests for the viral_video WebSocket progress endpoint and worker event format.
|
||||
|
||||
These tests exercise:
|
||||
* the worker _emit_progress helper (JSON serialization + event_type kwarg)
|
||||
* the pure helper functions on the API route module
|
||||
* WebSocket authentication / ownership / 404 behaviour
|
||||
* Initial-snapshot / terminal-job fast-close behaviour of the WS endpoint
|
||||
|
||||
The CI unit-test environment sets ``USE_IN_MEMORY_DB=true`` and relies on
|
||||
``settings.effective_database_url`` returning a SQLite URL. ``app/db.py`` and
|
||||
``app/dependencies.py`` have been fixed to honour ``effective_database_url``
|
||||
(matching the worker), so these tests never need a real Postgres or Redis.
|
||||
|
||||
Imports go through the ``apps.worker.*`` namespace (not bare ``worker_app.*``)
|
||||
to stay consistent with the existing integration tests and avoid creating a
|
||||
second module object that would make cross-file patches invisible.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.websockets import WebSocketDisconnect
|
||||
|
||||
# Ensure CI-friendly env is set BEFORE any app import so SQLite is used.
|
||||
os.environ.setdefault("USE_IN_MEMORY_DB", "true")
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "test-secret")
|
||||
os.environ.setdefault("DATABASE_URL", "postgresql+psycopg://no:such@127.0.0.1:1/none")
|
||||
|
||||
import app.db as _app_db # noqa: E402
|
||||
from app.api.routes import viral_video as vv_module # noqa: E402
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as worker_vv # noqa: E402
|
||||
|
||||
|
||||
def _make_job(**kwargs):
|
||||
defaults = dict(
|
||||
id="job-1",
|
||||
user_id="user-1",
|
||||
status="running",
|
||||
current_stage="analyzing",
|
||||
progress_percent=30,
|
||||
status_message="looking good",
|
||||
error_msg=None,
|
||||
is_terminal=False,
|
||||
result_video_url=None,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Worker event serialization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWorkerEmitProgress:
|
||||
def test_emit_progress_serialises_with_json_dumps(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress("job-1", "analyzing", 12, message="hi")
|
||||
fake_r.publish.assert_called_once()
|
||||
channel, payload = fake_r.publish.call_args.args
|
||||
assert channel == "viral_video:job-1"
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["stage"] == "analyzing"
|
||||
assert parsed["type"] == "viral_video:progress"
|
||||
assert parsed["progress"] == 12
|
||||
assert parsed["job_id"] == "job-1"
|
||||
assert "'stage'" not in payload # JSON uses double quotes, not Python repr
|
||||
|
||||
def test_emit_progress_respects_event_type(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-2",
|
||||
"done",
|
||||
100,
|
||||
message="ok",
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:completed"
|
||||
assert parsed["progress"] == 100
|
||||
|
||||
def test_emit_progress_failure_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-3",
|
||||
"failed",
|
||||
0,
|
||||
message="err",
|
||||
data={"error": "oom"},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:failed"
|
||||
assert parsed["data"]["error"] == "oom"
|
||||
|
||||
def test_emit_progress_wait_user_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-4",
|
||||
"intent_parsing",
|
||||
35,
|
||||
message="waiting for you",
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:wait_user"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure helpers on the route module
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSHelpers:
|
||||
def test_estimate_progress_maps_status(self):
|
||||
assert vv_module._estimate_progress(_make_job(status="pending")) == 0.0
|
||||
assert vv_module._estimate_progress(_make_job(status="running")) == 5.0
|
||||
assert vv_module._estimate_progress(_make_job(status="wait_user_confirm")) == 35.0
|
||||
assert vv_module._estimate_progress(_make_job(status="completed")) == 100.0
|
||||
assert vv_module._estimate_progress(_make_job(status="failed")) == 0.0
|
||||
|
||||
def test_initial_message_readable(self):
|
||||
job = _make_job(status="running")
|
||||
msg = vv_module._initial_message(job)
|
||||
assert isinstance(msg, str) and msg
|
||||
job_failed = _make_job(status="failed", error_msg="boom")
|
||||
assert "boom" in vv_module._initial_message(job_failed)
|
||||
job_wait = _make_job(status="wait_user_confirm")
|
||||
assert "等待" in vv_module._initial_message(job_wait)
|
||||
|
||||
def test_stage_from_status_falls_back(self):
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="pending")), str)
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="weird_unknown")), str)
|
||||
|
||||
def test_job_status_handles_enum_and_string(self):
|
||||
job = _make_job(status="running")
|
||||
assert vv_module._job_status(job) == "running"
|
||||
job_enum = _make_job(status=SimpleNamespace(value="completed"))
|
||||
assert vv_module._job_status(job_enum) == "completed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket authentication / ownership / 404
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSRejectsUnauthenticated:
|
||||
def test_no_token_closes_with_4401(self):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
with patch.object(vv_module, "_ws_authenticate_user", return_value=None):
|
||||
client = TestClient(app)
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-1"):
|
||||
pass
|
||||
assert exc.value.code == 4401
|
||||
|
||||
|
||||
def _build_client(*, auth_user, repo_get_return, redis_instance=None):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = repo_get_return
|
||||
sess = MagicMock()
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=auth_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch("redis.from_url", return_value=redis_instance or MagicMock()),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
return TestClient(app), fake_repo, sess, patches
|
||||
|
||||
|
||||
class TestWSOwnershipAnd404:
|
||||
def test_other_users_job_closes_with_4403(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
other_job = _make_job(user_id="user-b")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=other_job)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-x?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4403
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
def test_missing_job_closes_with_4404(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=None)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-missing?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4404
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _ws_authenticate_user direct unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSAuthenticateUser:
|
||||
def test_empty_token_returns_none(self):
|
||||
assert vv_module._ws_authenticate_user("") is None
|
||||
|
||||
def test_decode_exception_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", side_effect=Exception("bad token")):
|
||||
assert vv_module._ws_authenticate_user("not-a-jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_missing_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_non_string_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": 123}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_success_returns_user(self):
|
||||
sess = MagicMock()
|
||||
fake_user = SimpleNamespace(id="u1")
|
||||
fake_user_repo = MagicMock()
|
||||
fake_user_repo.find_by_id.return_value = fake_user
|
||||
with patch.object(_app_db, "SessionLocal", return_value=sess):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": "u1"}):
|
||||
with patch(
|
||||
"app.dependencies.get_user_repository",
|
||||
return_value=fake_user_repo,
|
||||
):
|
||||
result = vv_module._ws_authenticate_user("valid.jwt")
|
||||
assert result is fake_user
|
||||
fake_user_repo.find_by_id.assert_called_once_with("u1")
|
||||
sess.close.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket initial-snapshot / terminal-job fast-close tests.
|
||||
#
|
||||
# The Redis pubsub reader thread is factored into ``_run_pubsub_forwarder`` and
|
||||
# marked ``# pragma: no cover`` (integration-tested with a live Redis). These
|
||||
# tests patch it out so we can deterministically verify the pre-subscribe
|
||||
# handshake without needing a real Redis or real thread scheduling.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _run_ws_handshake(*, job):
|
||||
"""Drive a WS handshake; collect JSON messages before connection closes."""
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_user = SimpleNamespace(id=getattr(job, "user_id", "user-a"))
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = job
|
||||
sess = MagicMock()
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
return received, fake_repo, sess
|
||||
|
||||
|
||||
class TestWSInitialSnapshot:
|
||||
def test_running_job_sends_initial_snapshot(self):
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
received, repo, sess = _run_ws_handshake(job=job)
|
||||
assert received[0]["type"] == "viral_video:progress"
|
||||
assert received[0]["job_id"] == "job-1"
|
||||
assert received[0]["data"]["status"] == "running"
|
||||
# Session was used for both ownership check and initial snapshot.
|
||||
assert sess.close.call_count >= 2
|
||||
|
||||
def test_running_job_with_enum_status(self):
|
||||
job = _make_job(
|
||||
status=SimpleNamespace(value="wait_user_confirm"),
|
||||
user_id="user-a",
|
||||
is_terminal=False,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
assert received[0]["data"]["status"] == "wait_user_confirm"
|
||||
assert received[0]["progress"] == 35.0
|
||||
assert "等待" in received[0]["message"]
|
||||
|
||||
def test_already_completed_job_sends_completion_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url="https://example.com/v.mp4",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
types = [m["type"] for m in received]
|
||||
assert "viral_video:progress" in types
|
||||
assert "viral_video:completed" in types
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == "https://example.com/v.mp4"
|
||||
assert completed["progress"] == 100
|
||||
|
||||
def test_already_failed_job_sends_failed_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg="out of memory",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == "out of memory"
|
||||
assert failed["progress"] == 0
|
||||
|
||||
def test_completed_job_without_result_url_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == ""
|
||||
|
||||
def test_failed_job_without_error_msg_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == ""
|
||||
|
||||
def test_initial_snapshot_exception_is_swallowed(self):
|
||||
"""If sending the initial snapshot raises, the endpoint should log and
|
||||
still proceed to the Redis forwarder (doesn't crash)."""
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
await websocket.send_json({"type": "forwarder_reached"})
|
||||
await websocket.close()
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
fake_repo = MagicMock()
|
||||
calls = {"n": 0}
|
||||
|
||||
def _get(job_id):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 2:
|
||||
raise RuntimeError("boom in snapshot")
|
||||
return job
|
||||
|
||||
fake_repo.get.side_effect = _get
|
||||
sess = MagicMock()
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
assert any(m["type"] == "forwarder_reached" for m in received)
|
||||
Reference in New Issue
Block a user