Compare commits

..

1 Commits

Author SHA1 Message Date
CI Bot b127c115b6 style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 4s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 3s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 28s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m37s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 36s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 5m59s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 5m48s
AI Code Review / AI Code Review (pull_request) Successful in 7m38s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 6m35s
PR Automation / Auto Approve on CI Green (pull_request) Failing after 11m1s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Successful in 12m30s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 12m11s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Failing after 1s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
ACR Cleanup / ACR Image Cleanup (pull_request_target) Has been cancelled
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 31s
2026-09-30 00:19:17 +00:00
66 changed files with 847 additions and 11587 deletions
@@ -1,25 +0,0 @@
"""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")
@@ -1,51 +0,0 @@
"""viral video add copy_result + voice/video columns
Revision ID: 088_viral_video_copy_result
Revises: 087_viral_video_image_analysis
Create Date: 2026-10-01
v1.6 爆款视频字段补齐:
- copy_result JSON: 编导分镜脚本完整结构(overview/scene_and_lighting/shots/hard_constraints/negative_prompts/voiceover_script)
- voice_id/voice_source: TTS 音色参数
- video_ratio/video_model: Seedance 视频比例/模型
注意:线上启动也有幂等 ADD COLUMN 补列逻辑 (_ensure_viral_video_columns),本 migration 提供标准 Alembic 路径,
两套机制互不冲突(IF NOT EXISTS 等价行为)。
"""
import sqlalchemy as sa
from alembic import op
revision = "088_viral_video_copy_result"
down_revision = "087_viral_video_image_analysis"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 幂等添加列(通过单独执行 + 异常忽略兼容已由 backfill 补上的环境)
cols = [
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
conn = op.get_bind()
for name, ddl in cols:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
except Exception:
# 不支持 IF NOT EXISTS 的库(如老版本 SQLite)直接尝试 ADD COLUMN,失败则忽略
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
except Exception:
pass
def downgrade() -> None:
for name in ("copy_result", "video_model", "video_ratio", "voice_source", "voice_id"):
try:
op.drop_column("viral_video_jobs", name)
except Exception:
pass
-62
View File
@@ -1,62 +0,0 @@
"""viral video add storyboard + generated_copy_text (complement 088)
Revision ID: 089_viral_video_cols
Revises: 088_viral_video_copy_result
Create Date: 2026-10-01
#2129 兜底迁移:补齐 _VIRAL_VIDEO_BACKFILL_COLS 中所有列,覆盖
# watchtower 自动部署未跑历史 migration、且 AUTO_CREATE_SCHEMA=false 时
# _ensure_viral_video_columns 未执行的场景。
# 幂等 ADD COLUMN IF NOT EXISTS,已存在则跳过。
"""
import sqlalchemy as sa
from alembic import op
revision = "089_viral_video_cols"
down_revision = "088_viral_video_copy_result"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 扩展 alembic_version.version_num 字段长度(原来 VARCHAR(32) 装不下长 revision id)
conn = op.get_bind()
try:
conn.execute(sa.text("ALTER TABLE alembic_version ALTER COLUMN version_num TYPE VARCHAR(256)"))
except Exception:
pass
cols = [
("storyboard", "JSON"),
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
for name, ddl in cols:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
except Exception:
try:
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
except Exception:
pass
def downgrade() -> None:
for name in (
"copy_result",
"video_model",
"video_ratio",
"voice_source",
"voice_id",
"generated_copy_text",
"storyboard",
):
try:
op.drop_column("viral_video_jobs", name)
except Exception:
pass
+2 -9
View File
@@ -9,11 +9,6 @@ from fastapi.responses import JSONResponse
router = APIRouter(tags=["Health"])
def _pg_url(url: str) -> str:
"""Convert SQLAlchemy URL (postgresql+psycopg://...) to libpq connection string."""
return url.replace("postgresql+psycopg://", "postgresql://", 1).replace("postgresql+psycopg2://", "postgresql://", 1)
@router.get("/health", status_code=status.HTTP_200_OK)
async def health_check():
return {
@@ -54,7 +49,7 @@ async def _check_database() -> dict:
"message": "Using in-memory database",
}
try:
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
@@ -129,7 +124,7 @@ async def _check_migrations() -> dict:
"message": "Using in-memory database, no migrations needed",
}
try:
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("""
SELECT COUNT(*) FROM information_schema.tables
@@ -142,5 +137,3 @@ async def _check_migrations() -> dict:
return {"status": "unhealthy", "message": f"Missing tables, found {count}/5"}
except Exception as error:
return {"status": "unhealthy", "message": f"Migration check failed: {error}"}
-19
View File
@@ -191,23 +191,6 @@ def _find_duplicate_asset(
return None
def _get_existing_asset_url(existing: Any, storage_service: Any) -> str:
"""安全获取已存在素材的公网 URL,兼容 domain Asset(无 file_url 字段)和 ORM model。"""
# Domain Asset 只有 storage_key 字段;ORM model 有 file_url 但存的也是 storage_key
key = ""
for attr in ("storage_key", "file_url"):
v = getattr(existing, attr, None)
if v:
key = v
break
if not key:
return ""
try:
return storage_service.get_url(key) or ""
except Exception:
return ""
def _create_pending_asset(
asset_repository,
project_id,
@@ -407,7 +390,6 @@ async def prepare_direct_upload(
duplicated=True,
skip_transfer=True,
asset_id=existing.id,
url=_get_existing_asset_url(existing, storage_service),
)
file_id = uuid4().hex[:8]
@@ -461,7 +443,6 @@ async def prepare_direct_upload(
duplicated=False,
skip_transfer=False,
asset_id=pending_asset_id,
url="",
)
+23 -490
View File
@@ -1,21 +1,13 @@
"""爆款视频 API 路由。
v1.6 三步分步流水线端点(单次 Seedance 出片版):
POST /api/v1/viral-video/analyze-images 阶段1:创建任务 + 仅做图片/视频分析,暂停在 image_analyzed
POST /api/v1/viral-video/{job_id}/generate-copy 阶段2:用户填完参数后跑意图+文案+分镜+审核,暂停在 copy_generated
POST /api/v1/viral-video/{job_id}/confirm-copy 阶段3:用户确认/编辑文案后跑渲染,直到完成
旧端点(兼容保留,旧前端/一键生成模式):
POST /api/v1/viral-video/generate 一键入队,前半段跑到 wait_user_confirm
POST /api/v1/viral-video/{job_id}/confirm-intent 旧的意图确认后继续渲染
通用:
GET /api/v1/viral-video/{job_id} 查询任务状态(含 image_analysis/copy_result 编导脚本)
GET /api/v1/viral-video/history 历史记录
POST /api/v1/viral-video/{job_id}/retry 重试失败任务
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
GET /api/v1/viral-video/style-templates 风格模板列表
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送
端点:
POST /api/v1/viral-video/generate 创建爆款视频任务
GET /api/v1/viral-video/{job_id} 查询任务状态
GET /api/v1/viral-video/history 历史记录
POST /api/v1/viral-video/{job_id}/retry 重试失败任务
POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
GET /api/v1/viral-video/style-templates 获取风格模板列表
"""
from __future__ import annotations
@@ -23,22 +15,18 @@ 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 (
AnalyzeImagesRequest,
AnalyzeStyleRequest,
AnalyzeStyleResponse,
ConfirmCopyRequest,
ConfirmIntentRequest,
CreateViralVideoRequest,
GenerateCopyRequest,
StyleTemplateListResponse,
StyleTemplateResponse,
ViralVideoHistoryResponse,
ViralVideoJobResponse,
)
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
@@ -55,59 +43,6 @@ router = APIRouter()
# ── Helpers ──────────────────────────────────────────────────────────────
def _build_copy_result(job) -> dict | None:
"""v1.6: 返回编导分镜脚本 CopyResult 结构(给前端/Seedance 使用)。
- 若 job.copy_result 已持久化(v1.6 worker 生成),直接返回(补 final_copy 兜底)。
- 否则从老字段(generated_copy_text=口播, storyboard=分镜列表, intent_result)拼装兼容结构。
"""
cr = getattr(job, "copy_result", None)
if isinstance(cr, dict) and cr:
out = dict(cr)
# 向后兼容字段
voiceover = out.get("voiceover_script", "") or ""
out.setdefault("final_copy", voiceover)
out.setdefault("suggested_copy", voiceover)
out.setdefault("title", "")
return out
# 兼容 v1.5 老数据:storyboard 是老格式 [{order,type,description,text,duration,...}]
copy_text = getattr(job, "generated_copy_text", "") or ""
sb = getattr(job, "storyboard", None) or []
intent = getattr(job, "intent_result", None) or {}
if not copy_text and not sb:
return None
title = ""
if isinstance(intent, dict):
title = intent.get("suggested_title") or intent.get("intent", "") or ""
shots = []
for seg in sb:
if isinstance(seg, dict):
shots.append(
{
"time_range": "",
"shot_type_angle_movement": seg.get("ken_burns", ""),
"scene_and_dialogue": (seg.get("text") or "")
+ (" " + seg.get("description", "") if seg.get("description") else ""),
"action_details": "",
"audio_bgm": "",
"transition": seg.get("transition", "硬切"),
"reference_image_index": None,
}
)
ratio = getattr(job, "video_ratio", None) or "9:16"
return {
"overview": {"theme": title, "total_duration": getattr(job, "duration", 15), "aspect_ratio": ratio},
"scene_and_lighting": "",
"shots": shots,
"hard_constraints": ["无字幕", "无水印", "人物一致性"],
"negative_prompts": ["字幕", "水印", "错误文字", "五官崩坏"],
"voiceover_script": copy_text,
"final_copy": copy_text,
"suggested_copy": copy_text,
"title": title,
}
def _to_response(job) -> ViralVideoJobResponse:
return ViralVideoJobResponse(
id=job.id,
@@ -119,7 +54,7 @@ def _to_response(job) -> ViralVideoJobResponse:
viral_structure=job.viral_structure,
marketing_purpose=job.marketing_purpose,
bgm_preference=job.bgm_preference,
duration=job.duration or 15,
duration=job.duration,
user_copy_text=job.user_copy_text,
fusion_level=job.fusion_level,
reference_audio_path=job.reference_audio_path,
@@ -128,14 +63,6 @@ def _to_response(job) -> ViralVideoJobResponse:
style_guide=job.style_guide,
style_template_id=job.style_template_id,
status=job.status,
image_analysis=getattr(job, "image_analysis", None),
storyboard=getattr(job, "storyboard", None),
generated_copy_text=getattr(job, "generated_copy_text", "") or "",
copy_result=_build_copy_result(job),
voice_id=getattr(job, "voice_id", "") or "",
voice_source=getattr(job, "voice_source", "") or "",
video_ratio=getattr(job, "video_ratio", "9:16") or "9:16",
video_model=getattr(job, "video_model", "") or "",
intent_result=job.intent_result,
result_video_url=job.result_video_url,
credits_cost=job.credits_cost,
@@ -180,18 +107,13 @@ def create_viral_video(
viral_structure=request.viral_structure,
marketing_purpose=request.marketing_purpose,
bgm_preference=request.bgm_preference,
duration=request.duration or 15,
duration=request.duration,
user_copy_text=request.user_copy_text,
fusion_level=request.fusion_level,
reference_audio_path=request.reference_audio_path,
reference_video_url=request.reference_video_url,
style_strength=request.style_strength,
style_template_id=request.style_template_id,
voice_id=getattr(request, "voice_id", "") or "",
voice_source=getattr(request, "voice_source", "") or "",
video_ratio=getattr(request, "video_ratio", "9:16") or "9:16",
video_model=getattr(request, "video_model", "") or "",
copy_result=None,
)
# 持久化
@@ -199,7 +121,9 @@ def create_viral_video(
# 入队 Celery 任务
try:
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
from worker_app.tasks.viral_video import run_viral_video_pipeline
run_viral_video_pipeline.delay(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)
@@ -209,139 +133,6 @@ def create_viral_video(
return _to_response(job)
@router.post("/analyze-images", response_model=ViralVideoJobResponse)
def analyze_images(
request: AnalyzeImagesRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
session: Session = Depends(get_db_session),
) -> ViralVideoJobResponse:
"""v1.5 阶段1:创建任务并仅做图片/视频 VLM 分析,跑完后状态=image_analyzed。
前端拿到 image_analysis(商品名/品牌/特征/颜色/材质等结构化结果)展示给用户;
用户填完营销参数后再调 /{id}/generate-copy 进入阶段2。
"""
from packages.domain.viral_video import ViralVideoJob
repo = _get_job_repo(session)
job = ViralVideoJob(
user_id=authenticated_user.user.id,
images=list(request.images),
reference_video_url=request.reference_video_url or "",
style_template_id=request.style_template_id or "",
style_strength=request.style_strength or "medium",
voice_id=request.voice_id or "",
voice_source=request.voice_source or "",
video_ratio=request.video_ratio or "9:16",
video_model=request.video_model or "",
duration=request.duration or 15,
)
repo.save(job)
try:
celery_app.send_task("worker.run_viral_video_analyze", args=[job.id])
logger.info("[爆款视频][阶段1] analyze-images 入队: job_id=%s", job.id)
except Exception as e:
logger.error("[爆款视频][阶段1] analyze-images 入队失败: %s", e, exc_info=True)
job.mark_failed(f"任务入队失败: {e}")
repo.update(job)
return _to_response(job)
@router.post("/{job_id}/generate-copy", response_model=ViralVideoJobResponse)
def generate_copy(
job_id: str,
request: GenerateCopyRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
session: Session = Depends(get_db_session),
) -> ViralVideoJobResponse:
"""v1.6 阶段2:用户填完营销参数后,跑 意图解析 → 编导分镜脚本生成 → 合规审核。
跑完后状态=copy_generated,响应 copy_result(含 overview/scene_and_lighting/shots/
hard_constraints/negative_prompts/voiceover_script),前端展示脚本与口播供用户编辑;
确认/编辑后调 /{id}/confirm-copy 进入阶段3(TTS + 单次 Seedance 出片)。
"""
repo = _get_job_repo(session)
job = repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="无权操作此任务")
if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED):
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案")
# 允许失败任务重试:重置
if job.status == ViralVideoStatus.FAILED:
job.retry_count += 1
job.error_msg = ""
# 把用户填的营销参数写到 job 上
job.industry = request.industry or job.industry
job.target_customer = request.target_customer or job.target_customer
job.persona_id = request.persona_id or job.persona_id
job.viral_structure = request.viral_structure or job.viral_structure
job.marketing_purpose = request.marketing_purpose or job.marketing_purpose
job.bgm_preference = request.bgm_preference or job.bgm_preference
if request.duration:
job.duration = max(5, min(30, int(request.duration)))
job.user_copy_text = request.user_copy_text if request.user_copy_text else job.user_copy_text
job.fusion_level = request.fusion_level or job.fusion_level
job.reference_audio_path = request.reference_audio_path or job.reference_audio_path
job.reference_video_url = request.reference_video_url or job.reference_video_url
job.style_strength = request.style_strength or job.style_strength
job.style_template_id = request.style_template_id or job.style_template_id
if request.style_guide is not None:
job.style_guide = request.style_guide
job.voice_id = request.voice_id or job.voice_id
job.voice_source = request.voice_source or job.voice_source
job.video_ratio = request.video_ratio or job.video_ratio or "9:16"
job.video_model = request.video_model or job.video_model or ""
job.resume_from_image_analyzed()
repo.update(job)
try:
celery_app.send_task("worker.run_viral_video_generate_copy", args=[job.id])
logger.info("[爆款视频][阶段2] generate-copy 入队: job_id=%s", job.id)
except Exception as e:
logger.error("[爆款视频][阶段2] generate-copy 入队失败: %s", e, exc_info=True)
job.mark_failed(f"任务入队失败: {e}")
repo.update(job)
return _to_response(job)
@router.post("/{job_id}/confirm-copy", response_model=ViralVideoJobResponse)
def confirm_copy(
job_id: str,
request: ConfirmCopyRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
session: Session = Depends(get_db_session),
) -> ViralVideoJobResponse:
"""v1.6 阶段3:用户确认/编辑口播后开始 TTS + 单次 Seedance 生成 + 上传。"""
repo = _get_job_repo(session)
job = repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="无权操作此任务")
if job.status != ViralVideoStatus.COPY_GENERATED:
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
job.resume_from_copy_generated(edited_copy=request.edited_copy or None)
repo.update(job)
try:
celery_app.send_task("worker.run_viral_video_render", args=[job.id])
logger.info("[爆款视频][阶段3] confirm-copy 入队: job_id=%s", job.id)
except Exception as e:
logger.error("[爆款视频][阶段3] confirm-copy 入队失败: %s", e, exc_info=True)
job.mark_failed(f"任务入队失败: {e}")
repo.update(job)
return _to_response(job)
@router.get("/history", response_model=ViralVideoHistoryResponse)
def list_viral_video_history(
limit: int = 50,
@@ -418,7 +209,9 @@ def retry_viral_video_job(
# 重新入队
try:
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
from worker_app.tasks.viral_video import run_viral_video_pipeline
run_viral_video_pipeline.delay(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)
@@ -455,7 +248,9 @@ def confirm_intent(
# 从断点恢复 Celery 任务
try:
celery_app.send_task("worker.resume_viral_video_pipeline", args=[job.id])
from worker_app.tasks.viral_video import resume_viral_video_pipeline
resume_viral_video_pipeline.delay(job.id)
logger.info("[爆款视频] 意图确认,恢复流水线: job_id=%s", job.id)
except Exception as e:
logger.error("[爆款视频] 恢复流水线失败: %s", e, exc_info=True)
@@ -488,7 +283,9 @@ def analyze_style(
# 入队风格分析任务
try:
celery_app.send_task("worker.run_video_style_analysis", args=[job.id])
from worker_app.tasks.viral_video import run_video_style_analysis
run_video_style_analysis.delay(job.id)
logger.info("[爆款视频] 风格分析入队: job_id=%s", job.id)
except Exception as e:
logger.error("[爆款视频] 风格分析入队失败: %s", e, exc_info=True)
@@ -498,267 +295,3 @@ def analyze_style(
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": "",
"image_analyzed": "image_analysis",
"copy_generated": "review",
"wait_user_confirm": "intent_parsing",
"completed": "uploading",
"failed": "",
"cancelled": "",
}
_STATUS_PROGRESS = {
"pending": 0.0,
"running": 5.0,
"image_analyzed": 15.0,
"copy_generated": 70.0,
"wait_user_confirm": 35.0,
"completed": 100.0,
"failed": 0.0,
"cancelled": 0.0,
}
_STATUS_MESSAGE = {
"pending": "任务已创建,等待执行",
"running": "任务执行中",
"image_analyzed": "图片分析完成,等待填写营销参数",
"copy_generated": "文案与分镜已生成,等待确认文案",
"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
View File
@@ -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.effective_database_url)
ensure_database_exists(settings.DATABASE_URL)
engine, SessionLocal = build_session_factory(
settings.effective_database_url,
settings.DATABASE_URL,
pool_size=settings.DATABASE_POOL_SIZE,
max_overflow=settings.DATABASE_MAX_OVERFLOW,
pool_timeout=settings.DATABASE_POOL_TIMEOUT,
+1 -1
View File
@@ -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.effective_database_url)
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
def get_db_session() -> Generator[Session, None, None]:
-2
View File
@@ -29,8 +29,6 @@ class DirectUploadPrepareResponse(BaseModel):
duplicated: bool = False
skip_transfer: bool = False
asset_id: str = ""
# duplicated=true 时填充已存在素材的公网 URL,前端可直接用而不必再调 complete
url: str = Field(default="", description="duplicated=true 时已存在素材的公网 URL")
class DirectUploadCompleteRequest(BaseModel):
+48 -150
View File
@@ -1,4 +1,4 @@
"""爆款视频 API schemas (v1.6 单次 Seedance 出片版)。"""
"""爆款视频 API schemas。"""
from __future__ import annotations
@@ -6,182 +6,81 @@ from datetime import datetime
from pydantic import BaseModel, Field, field_validator
# -- 枚举常量 --
# ── 枚举常量 ─────────────────────────────────────────────────────────────
VALID_FUSION_LEVELS = ("ai_full", "full_ai", "ai_polish", "user_primary")
VALID_FUSION_LEVELS = ("ai_full", "ai_polish", "user_primary")
VALID_STYLE_STRENGTHS = ("light", "medium", "strict")
VALID_STAGES = (
"image_analysis",
"video_analysis",
"intent_parsing",
"script_generation",
"copy_fusion",
"storyboard",
"review",
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
)
VALID_VIDEO_RATIOS = ("9:16", "16:9", "1:1", "4:3", "3:4", "21:9")
VALID_DURATIONS = (5, 10, 15, 20, 25, 30)
# -- 编导脚本结构(v1.6) --
class ShotScript(BaseModel):
"""逐镜头分镜。"""
time_range: str = Field(default="", description="时间区间,如 0-3秒")
shot_type_angle_movement: str = Field(default="", description="景别/角度/运镜,如『近景俯拍45度,缓慢推镜』")
scene_and_dialogue: str = Field(default="", description="场景描述+口播台词")
action_details: str = Field(default="", description="人物动作、表情、物品操作细节")
audio_bgm: str = Field(default="", description="环境音+BGM提示")
transition: str = Field(default="硬切", description="转场方式:硬切/淡入淡出/叠化")
reference_image_index: int | None = Field(
default=None, description="参考图片索引(0-based,对应上传的第几张产品图)"
)
class CopyResultOverview(BaseModel):
theme: str = ""
total_duration: int = 15
aspect_ratio: str = "9:16"
class CopyResult(BaseModel):
"""v1.6 编导分镜脚本结构(给前端 + Seedance 用)。"""
overview: CopyResultOverview = Field(default_factory=CopyResultOverview)
scene_and_lighting: str = ""
shots: list[ShotScript] = Field(default_factory=list)
hard_constraints: list[str] = Field(default_factory=list)
negative_prompts: list[str] = Field(default_factory=list)
voiceover_script: str = Field(
default="", description="纯口播对白,从各镜 scene_and_dialogue 的对白部分拼接,供 TTS 使用"
)
# 向后兼容:final_copy = voiceover_script
final_copy: str = ""
suggested_copy: str = ""
title: str = ""
# -- Request Schemas --
# ── Request Schemas ────────────────────────────────────────────────────────
class CreateViralVideoRequest(BaseModel):
"""旧接口:一键创建(保留兼容)。"""
"""创建爆款视频任务请求。"""
images: list[str] = Field(..., min_length=1, max_length=20)
industry: str = ""
target_customer: str = ""
persona_id: str = ""
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = Field(default=15, ge=5, le=30, description="视频时长(秒),5-30")
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
reference_video_url: str = ""
style_strength: str = "medium"
style_template_id: str = ""
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
images: list[str] = Field(..., min_length=1, max_length=20, description="产品图片 URL 列表")
industry: str = Field(default="", description="行业")
target_customer: str = Field(default="", description="目标客户描述")
persona_id: str = Field(default="", description="人设 ID")
viral_structure: str = Field(default="", description="爆款结构类型")
marketing_purpose: str = Field(default="", description="营销目的")
bgm_preference: str = Field(default="", description="BGM 偏好")
duration: int = Field(default=30, ge=5, le=180, description="视频时长(秒)")
user_copy_text: str = Field(default="", description="用户原始文案(我说你写)")
fusion_level: str = Field(default="ai_polish", description="文案融合级别: ai_full/ai_polish/user_primary")
reference_audio_path: str = Field(default="", description="参考音频路径")
# v1.3 新增
reference_video_url: str = Field(default="", description="参考爆款视频 URL")
style_strength: str = Field(default="medium", description="风格强度: light/medium/strict")
style_template_id: str = Field(default="", description="风格模板 ID")
@field_validator("fusion_level")
@classmethod
def _v_fl(cls, v: str) -> str:
if v == "full_ai":
return "ai_full"
def _validate_fusion_level(cls, v: str) -> str:
if v not in VALID_FUSION_LEVELS:
raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}")
raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一")
return v
@field_validator("style_strength")
@classmethod
def _v_ss(cls, v: str) -> str:
def _validate_style_strength(cls, v: str) -> str:
if v not in VALID_STYLE_STRENGTHS:
raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}")
raise ValueError(f"style_strength 必须是 {VALID_STYLE_STRENGTHS} 之一")
return v
class AnalyzeImagesRequest(BaseModel):
"""v1.5+ 阶段1:创建任务 + 图片/视频分析。"""
images: list[str] = Field(..., min_length=1, max_length=30)
reference_video_url: str = ""
style_template_id: str = ""
style_strength: str = "medium"
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
duration: int = Field(default=15, ge=5, le=30)
class GenerateCopyRequest(BaseModel):
"""v1.5+ 阶段2:填完营销参数,生成编导脚本。"""
industry: str = ""
target_customer: str = ""
persona_id: str = ""
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = Field(default=15, ge=5, le=30)
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
reference_video_url: str = ""
style_strength: str = "medium"
style_template_id: str = ""
style_guide: dict | None = None
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
@field_validator("fusion_level")
@classmethod
def _v_fl(cls, v: str) -> str:
if v == "full_ai":
return "ai_full"
if v not in VALID_FUSION_LEVELS:
raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}")
return v
@field_validator("style_strength")
@classmethod
def _v_ss(cls, v: str) -> str:
if v not in VALID_STYLE_STRENGTHS:
raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}")
return v
class ConfirmCopyRequest(BaseModel):
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
class ConfirmIntentRequest(BaseModel):
"""旧 confirm-intent(兼容)。"""
"""确认意图请求(confirm-intent)。"""
confirmed_copy: str = ""
adjustments: str = ""
confirmed_copy: str = Field(default="", description="用户确认/修改后的文案,为空表示使用 AI 生成的文案")
adjustments: str = Field(default="", description="用户对 AI 文案的调整意见")
class AnalyzeStyleRequest(BaseModel):
"""触发参考视频风格分析请求。"""
reference_video_url: str = Field(..., description="参考视频 URL")
style_template_id: str = ""
style_template_id: str = Field(default="", description="风格模板 ID(可选覆盖)")
# -- Response Schemas --
# ── Response Schemas ───────────────────────────────────────────────────────
class ViralVideoJobResponse(BaseModel):
"""爆款视频任务响应(v1.6 包含 copy_result 编导脚本结构)。"""
"""爆款视频任务响应。"""
id: str
user_id: str
@@ -192,7 +91,7 @@ class ViralVideoJobResponse(BaseModel):
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = 15
duration: int = 30
user_copy_text: str = ""
fusion_level: str = "ai_polish"
reference_audio_path: str = ""
@@ -201,17 +100,6 @@ class ViralVideoJobResponse(BaseModel):
style_guide: dict | None = None
style_template_id: str = ""
status: str
image_analysis: dict | None = None
# v1.6 编导脚本(推荐前端使用)
copy_result: dict | None = None
# v1.5 兼容字段
storyboard: list | None = None
generated_copy_text: str = ""
# 音色/视频参数
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
intent_result: dict | None = None
result_video_url: str = ""
credits_cost: int = 0
@@ -224,11 +112,15 @@ class ViralVideoJobResponse(BaseModel):
class ViralVideoHistoryResponse(BaseModel):
"""历史记录列表响应。"""
items: list[ViralVideoJobResponse]
total: int
class StyleTemplateResponse(BaseModel):
"""风格模板响应。"""
id: str
name: str
description: str = ""
@@ -237,19 +129,25 @@ class StyleTemplateResponse(BaseModel):
class StyleTemplateListResponse(BaseModel):
"""风格模板列表响应。"""
items: list[StyleTemplateResponse]
class AnalyzeStyleResponse(BaseModel):
"""风格分析结果响应。"""
job_id: str
status: str
style_guide: dict | None = None
# -- WebSocket 事件 Schema --
# ── WebSocket 事件 Schema ──────────────────────────────────────────────────
class WSProgressEvent(BaseModel):
"""WebSocket 进度推送事件。"""
type: str = "viral_video:progress"
job_id: str
stage: str
-2
View File
@@ -150,8 +150,6 @@ export interface DirectUploadPrepareResult {
* 两个字段是同一语义的别名(后端可能只返回其一),前端任意为 true 即视为命中去重。
*/
skip_transfer?: boolean
/** duplicated=true 时后端返回已存在素材的公网 URL,前端直接用而不必再调 complete */
url?: string
}
/** 直传完成确认返回 */
+6 -47
View File
@@ -3,24 +3,9 @@
*/
import apiClient from "../client"
import { getOrCreateDefaultProject } from "../projects"
import { ensureDefaultLibrary } from "./libraries"
import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types"
import { computeFileHash, makeClientUploadId } from "./uploadDedup"
/** 根据 File.type 推断素材库 kind(image/video/voice);无法推断时默认 image */
function inferKindFromFile(file: File): "image" | "video" | "voice" {
const t = (file.type || "").toLowerCase()
if (t.startsWith("image/")) return "image"
if (t.startsWith("video/")) return "video"
if (t.startsWith("audio/")) return "voice"
// 兜底:按扩展名再判一次
const name = file.name.toLowerCase()
if (/\.(png|jpe?g|gif|webp|bmp|svg|avif)$/.test(name)) return "image"
if (/\.(mp4|mov|webm|avi|mkv|flv|wmv|m4v)$/.test(name)) return "video"
if (/\.(mp3|wav|m4a|aac|ogg|flac|opus|webm)$/.test(name)) return "voice"
return "image"
}
/** 预签名直传准备 */
export const prepareDirectUpload = async (data: {
project_id: string
@@ -123,8 +108,6 @@ const putToOSS = (
/** 单个文件的上传阶段信息(供批量上传队列做状态绑定) */
export interface DirectUploadHandle {
/** 实际使用的素材库(内部解析出来,便于调用方做后续 UI/缓存操作) */
library: { id: string; kind: "image" | "video" | "voice" }
/** prepare 返回(含可能的预建 asset_id) */
prepared: DirectUploadPrepareResult
/** 直传 OSS(可重复调用用于重试) */
@@ -136,17 +119,10 @@ export interface DirectUploadHandle {
/**
* 准备一次直传:调 prepare 拿到签名表单(后端可能同时预建 uploading 态 asset),
* 返回分段执行的 handle,调用方自行控制 transfer/complete 时机(便于队列并发与重试)。
*
* 修复 P0 404:library_id 改为可选;未传时自动根据文件类型在默认项目下确保对应素材库存在,
* 避免调用方从「全部素材库列表」里挑一个 library_id、但与默认项目 project_id 不匹配,
* 导致后端返回 "Asset library not found" 404。
*/
export const prepareDirectUploadHandle = async (data: {
file: File
/** 素材库 ID;未传时按文件类型自动在默认项目下 ensure-default */
library_id?: string
/** 显式指定素材库 kind;未传时按 MIME/扩展名推断 */
kind?: "image" | "video" | "voice"
library_id: string
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
fileHash?: string
/** 本次逻辑上传的幂等 token,prepare/complete 一致、重试复用 */
@@ -162,17 +138,9 @@ export const prepareDirectUploadHandle = async (data: {
throw new Error(`初始化默认项目失败,无法开始上传:${reason}`)
}
// 解析 library_id:调用方传了就用,没传就按 kind 自动 ensure-default
let resolvedLibraryId = data.library_id
const resolvedKind = data.kind ?? inferKindFromFile(data.file)
if (!resolvedLibraryId) {
const lib = await ensureDefaultLibrary({ project_id: project.id, kind: resolvedKind })
resolvedLibraryId = lib.id
}
const prepared = await prepareDirectUpload({
project_id: project.id,
library_id: resolvedLibraryId,
library_id: data.library_id,
filename: data.file.name,
content_type: data.file.type || "application/octet-stream",
file_size: data.file.size,
@@ -181,13 +149,12 @@ export const prepareDirectUploadHandle = async (data: {
})
return {
library: { id: resolvedLibraryId, kind: resolvedKind },
prepared,
transfer: (onProgress) => putToOSS(prepared, data.file, onProgress),
complete: () =>
completeDirectUpload({
project_id: project.id,
library_id: resolvedLibraryId,
library_id: data.library_id,
storage_key: prepared.storage_key,
file_hash: data.fileHash,
client_upload_id: data.clientUploadId,
@@ -197,17 +164,10 @@ export const prepareDirectUploadHandle = async (data: {
}
}
/** 直传上传(大文件推荐),支持可选进度回调;一次性完成 prepare→transfer→complete
*
* P0 404 修复:library_id 可选;不传时内部按文件类型自动匹配正确项目下的素材库,
* 保证 project_id 与 library_id 必然一致。
*/
/** 直传上传(大文件推荐),支持可选进度回调;一次性完成 prepare→transfer→complete */
export const uploadAssetDirect = async (data: {
file: File
/** 素材库 ID;可选,不传按文件类型自动解析默认项目下的对应素材库(推荐用法) */
library_id?: string
/** 显式指定素材库 kind;未传时按文件 MIME/扩展名推断 */
kind?: "image" | "video" | "voice"
library_id: string
onProgress?: (percent: number) => void
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
fileHash?: string
@@ -220,7 +180,6 @@ export const uploadAssetDirect = async (data: {
const handle = await prepareDirectUploadHandle({
file: data.file,
library_id: data.library_id,
kind: data.kind,
fileHash,
clientUploadId,
})
@@ -229,7 +188,7 @@ export const uploadAssetDirect = async (data: {
return {
storage_key: handle.prepared.storage_key,
ingest_job_id: "",
url: handle.prepared.url || "",
url: "",
duplicated: true,
asset_id: handle.prepared.asset_id,
}
-131
View File
@@ -1,131 +0,0 @@
import apiClient from "@/api/client"
import type {
GenerateViralVideoRequest,
HistoryResponse,
StyleTemplate,
ViralVideoJob,
ImageAnalysisResult,
CopyResult,
AnalyzeImagesRequest,
GenerateCopyRequest,
ConfirmCopyRequest,
} 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)
}
/** 上传参考视频后触发风格分析 */
export function analyzeViralStyle(id: string) {
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
}
/** ── 三步拆分:前端 mock 辅助函数(后端新接口上线后可替换) ── */
/**
* 客户端图片分析 mock(后端未提供 analyze-only 端点前的占位方案):
* 基于已上传图片生成一份示例识别汇览,让 STEP1→STEP2 交互可走通。
* 后端上线后改为调用真实接口。
*/
export function mockImageAnalysis(images: { name: string }[]): Promise<ImageAnalysisResult> {
return new Promise((resolve) => {
setTimeout(() => {
const products = images.slice(0, 3).map((img, i) => {
const n = img.name.replace(/\.[^.]+$/, "")
return {
name: n || `商品 ${i + 1}`,
spec: i === 0 ? "500ml/瓶" : i === 1 ? "300g/盒" : undefined,
brand: i === 0 ? "示例品牌" : undefined,
features:
i === 0
? "瓶身透明、蓝色标签、白色瓶盖;标签上印有品牌Logo和产品名称;光线均匀,主体居中"
: i === 1
? "盒装包装、主色调为米白+暖黄;正面有产品实物图;文字清晰可辨"
: "产品主体清晰、背景干净、色彩鲜艳,突出核心卖点",
label_text: i === 0 ? "包装正面印有产品名称、净含量、品牌Logo" : undefined,
image_index: i,
}
})
resolve({ products })
}, 1800)
})
}
/**
* 客户端文案生成 mock(后端未提供 generate-copy 端点前的占位方案):
* 后端上线后改为调用真实接口。
*/
export function mockGenerateCopy(params: {
product: string
sellingPoints?: string[]
tone?: string
duration?: number
marketingPurpose?: string
industry?: string
targetCustomer?: string
}): Promise<CopyResult> {
return new Promise((resolve) => {
setTimeout(() => {
const product = params.product || "这款产品"
const tone = params.tone || "亲切务实"
const purpose = params.marketingPurpose || "品牌种草"
resolve({
title: `【${purpose}】${product},用过的人都说好!`,
final_copy: `你有没有发现,选对一款${params.industry || "好物"}真的能让生活省心很多?\n\n今天给大家推荐这款${product}。${tone.includes("亲切") ? "说实话," : ""}我自己用了一段时间,最直观的感受就是——好用、省心、值得回购。\n\n✅ 亮点一:品质到位,用料扎实,细节处见用心\n✅ 亮点二:使用体验舒服,日常高频场景都能打\n✅ 亮点三:性价比很能打,这个价位真的没什么可挑的\n\n如果你也在找一款靠谱的${params.industry || "日常好物"},真的建议试试${product},不会让你失望。点击左下角,直接入手!`,
suggested_copy: `你有没有发现,选对一款${params.industry || "好物"}真的能让生活省心很多?\n\n今天给大家推荐这款${product}。${tone.includes("亲切") ? "说实话," : ""}我自己用了一段时间,最直观的感受就是——好用、省心、值得回购。\n\n✅ 亮点一:品质到位,用料扎实,细节处见用心\n✅ 亮点二:使用体验舒服,日常高频场景都能打\n✅ 亮点三:性价比很能打,这个价位真的没什么可挑的\n\n如果你也在找一款靠谱的${params.industry || "日常好物"},真的建议试试${product},不会让你失望。点击左下角,直接入手!`,
})
}, 2200)
})
}
/** ── 三步拆分 v1.5 真实后端 API(PR #2117 合入后启用,前端可替换 mock 调用) ── */
/** 阶段1:上传图片后仅做 VLM 图片分析 + 可选参考视频风格分析,完成后状态=image_analyzed */
export function analyzeViralImages(payload: AnalyzeImagesRequest) {
return apiClient.post<ViralVideoJob>("/viral-video/analyze-images", payload).then((r) => r.data)
}
/** 阶段2:用户填完营销参数后生成文案+分镜+合规审核,完成后状态=copy_generated,返回 copy_result */
export function generateViralCopy(id: string, payload: GenerateCopyRequest) {
return apiClient
.post<ViralVideoJob>(`/viral-video/${id}/generate-copy`, payload)
.then((r) => r.data)
}
/** 阶段3:用户确认/编辑文案后开始 TTS→渲染→上传,完成后状态=completed */
export function confirmViralCopy(id: string, payload: ConfirmCopyRequest = {}) {
return apiClient
.post<ViralVideoJob>(`/viral-video/${id}/confirm-copy`, payload)
.then((r) => r.data)
}
-298
View File
@@ -1,298 +0,0 @@
export type FusionLevel = "ai_full" | "ai_polish" | "user_primary"
export const FUSION_LEVELS: { value: FusionLevel; label: string; desc: string }[] = [
{ value: "ai_full", label: "AI 全写", desc: "给我方向,全由AI创作" },
{ value: "ai_polish", label: "AI润色", desc: "我写草稿,AI帮我润色" },
{ value: "user_primary", 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: "像素级复刻" },
]
/** v1.6 前端时长下拉选项(5/10/15/20/25/30秒) */
export const VALID_DURATIONS = [5, 10, 15, 20, 25, 30] as const
export type VideoDuration = (typeof VALID_DURATIONS)[number]
/** v1.6 支持的画幅比例 */
export const VALID_RATIOS = ["9:16", "16:9", "1:1"] as const
export type VideoRatio = (typeof VALID_RATIOS)[number]
export type ViralVideoStatus =
| "pending"
| "running"
| "wait_user_confirm"
| "image_analyzed"
| "copy_generated"
| "completed"
| "failed"
| "cancelled"
/**
* v1.6 后端流水线阶段。单次 Seedance 出片版:
* image_analysis → video_analysis(可选) → intent_parsing → script_generation → review → tts → rendering → uploading
*/
export type ViralVideoStage =
| "image_analysis"
| "video_analysis"
| "intent_parsing"
| "script_generation"
| "review"
| "tts"
| "rendering"
| "uploading"
/** 图片+视频分析阶段:属于「分析图片」按钮的范围 */
const IMAGE_ANALYSIS_STAGES = new Set<ViralVideoStage>(["image_analysis", "video_analysis"])
/** 编导脚本阶段:属于「生成文案」按钮的范围 */
const COPY_STAGES = new Set<ViralVideoStage>(["intent_parsing", "script_generation", "review"])
/** 视频生成阶段:属于「开始生成视频」按钮的范围(v1.6: TTS+单次Seedance+上传) */
const VIDEO_STAGES = new Set<ViralVideoStage>(["tts", "rendering", "uploading"])
export function isImageAnalysisStage(stage: ViralVideoStage | undefined): boolean {
return !!stage && IMAGE_ANALYSIS_STAGES.has(stage)
}
export function isCopyStage(stage: ViralVideoStage | undefined): boolean {
return !!stage && COPY_STAGES.has(stage)
}
export function isVideoStage(stage: ViralVideoStage | undefined): boolean {
return !!stage && VIDEO_STAGES.has(stage)
}
/** 兼容旧调用:分析图片+生成文案 的所有前置阶段 */
export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
return isImageAnalysisStage(stage) || isCopyStage(stage)
}
/** 单张图片 VLM 识别出的商品信息 */
export interface ImageProductAnalysis {
name?: string
category?: string
brand?: string
colors?: string[]
material_or_texture?: string
key_features?: string[]
visual_style?: string
scene?: string
target_audience_hint?: string
text_on_image?: string
/** 旧字段兼容 */
spec?: string
features?: string[] | string
label_text?: string
selling_points?: string
image_index?: number
}
export interface ImageAnalysisResult {
products?: ImageProductAnalysis[]
}
/** v1.6 编导分镜脚本 - 单镜头 */
export interface ShotScript {
/** 时间区间,如 "0-3秒" */
time_range?: string
/** 景别/角度/运镜,如 "近景俯拍45度,缓慢推镜" */
shot_type_angle_movement?: string
/** 场景描述+对白 */
scene_and_dialogue?: string
/** 人物动作/表情/物品操作细节 */
action_details?: string
/** 环境音+BGM提示 */
audio_bgm?: string
/** 转场方式(硬切/淡入淡出/叠化/结束) */
transition?: string
/** 参考图片索引(0-based,对应上传产品图数组) */
reference_image_index?: number | null
}
/** v1.6 编导分镜脚本 - 总览 */
export interface CopyResultOverview {
theme?: string
total_duration?: number
aspect_ratio?: string
}
/** v1.6 编导分镜脚本(核心输出结构,给 Seedance 做 prompt,给 TTS 取 voiceover_script) */
export interface CopyResult {
overview?: CopyResultOverview
/** 整体场景+光线描述 */
scene_and_lighting?: string
/** 逐镜头时间轴 */
shots?: ShotScript[]
/** 硬性约束(禁止字幕/水印/变形等) */
hard_constraints?: string[]
/** 负面提示词 */
negative_prompts?: string[]
/** 完整口播稿(纯文本,用于 TTS 合成) */
voiceover_script?: string
/** 向后兼容:= voiceover_script */
final_copy?: string
/** 向后兼容:= voiceover_script */
suggested_copy?: string
title?: string
/** v1.5 旧字段兼容(老数据降级时可能出现) */
scenes?: Array<{ shot: string; narration: string; duration?: number }>
}
export interface StyleTemplate {
id: string
name: string
description?: string
thumbnail_url?: string
style_config?: Record<string, unknown>
tags?: string[]
}
export interface IntentResult {
intent?: string
key_messages?: string[]
tone?: string
target_emotion?: string
call_to_action?: string
suggested_title?: string
/** v1.5 旧字段兼容 */
product?: string
selling_points?: string[]
target_audience?: string
structure?: string
duration?: number
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 | Record<string, unknown>
user_copy_text?: string
/** v1.6: = copy_result.voiceover_script(从 copy_result 派生,向后兼容) */
final_copy_text?: string
generated_copy_text?: string
fusion_level?: FusionLevel
voice_id?: string
voice_mode?: "global" | "per_video"
voice_source?: "preset" | "library" | "clone" | "upload"
bgm_preference?: string
intent_result?: IntentResult
intent_text?: string
/** v1.6 编导分镜脚本(核心产物) */
copy_result?: CopyResult
/** 向后兼容:= copy_result.shots */
storyboard?: ShotScript[]
image_analysis?: ImageAnalysisResult
/** 视频比例:9:16 / 16:9 / 1:1,默认 9:16 */
video_ratio?: string
/** Seedance 模型 ID(空=后端默认) */
video_model?: string
/** 视频时长(秒,5-30,默认15) */
duration?: number
progress_stage?: ViralVideoStage
progress_percent?: number
progress_message?: string
output_url?: string
result_video_url?: string
error_message?: string
error_msg?: string
credits_cost?: number
created_at?: string
updated_at?: string
}
export interface GenerateViralVideoRequest {
images: string[]
reference_video_url?: string
douyin_url?: string
style_strength?: StyleStrength
style_template_id?: string
user_copy_text?: string
fusion_level?: FusionLevel
voice_id?: string
voice_source?: "preset" | "library" | "clone" | "upload"
bgm_preference?: string
industry?: string
target_customer?: string
language?: string
persona_id?: string
viral_structure?: string
marketing_purpose?: string
/** 视频时长(5-30秒,默认15) */
duration?: number
video_model?: string
video_ratio?: string
/** 三步拆分:step 控制后端执行到哪一步暂停 */
step?: "analyze" | "generate_copy" | "generate_video"
}
export interface HistoryResponse {
items: ViralVideoJob[]
total: number
page: number
page_size: number
}
/** v1.6 阶段1请求:图片/视频分析(POST /viral-video/analyze-images) */
export interface AnalyzeImagesRequest {
images: string[]
reference_video_url?: string
style_template_id?: string
style_strength?: StyleStrength
/** TTS 音色 ID(STEP1 已选音色时传) */
voice_id?: string
/** 音色来源:preset | library | clone | upload */
voice_source?: "preset" | "library" | "clone" | "upload"
/** Seedance 视频比例:9:16 | 16:9 | 1:1 */
video_ratio?: string
/** Seedance 模型 ID(空则使用服务端默认) */
video_model?: string
/** 视频时长(秒,5-30,默认15) */
duration?: number
}
/** v1.6 阶段2请求:填完营销参数后生成编导分镜脚本(POST /viral-video/{id}/generate-copy) */
export interface GenerateCopyRequest {
industry?: string
target_customer?: string
persona_id?: string
viral_structure?: string
marketing_purpose?: string
bgm_preference?: string
/** 视频时长(秒,5-30,默认15) */
duration?: number
user_copy_text?: string
fusion_level?: FusionLevel
reference_audio_path?: string
reference_video_url?: string
style_strength?: StyleStrength
style_template_id?: string
style_guide?: string | Record<string, unknown>
/** TTS 音色 ID(优先级高于 persona_id) */
voice_id?: string
/** 音色来源:preset | library | clone | upload */
voice_source?: "preset" | "library" | "clone" | "upload"
/** Seedance 视频比例(9:16/16:9/1:1 等) */
video_ratio?: string
/** Seedance 模型 ID(空则使用服务端默认) */
video_model?: string
}
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
export interface ConfirmCopyRequest {
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
edited_copy?: string
}
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
export interface StoryboardSegment {
order: number
type: string
description: string
text: string
duration: number
ken_burns?: string
transition?: string
}
@@ -2,7 +2,6 @@
export const ROUTE_TITLE_MAP: Record<string, string> = {
"/app/dashboard": "首页",
"/app/generate": "智能剪辑",
"/app/viral-video": "爆款视频",
"/app/assets": "视频库",
"/app/voices": "配音库",
"/app/products": "成片库",
-13
View File
@@ -18,7 +18,6 @@ import {
ThunderboltOutlined,
UnorderedListOutlined,
UserOutlined,
FireOutlined,
} from "@ant-design/icons"
/** 导航项类型 */
@@ -77,12 +76,6 @@ 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: "任务历史",
@@ -149,12 +142,6 @@ 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),
},
],
},
{
@@ -580,6 +580,20 @@ const AiAvatarPage: React.FC = () => {
return (
<div className="aa-page">
<div className="aa-page-header">
<h1>AI数字人</h1>
</div>
{/* 步骤切换导航条 */}
<div className="aa-step-nav">
<span className={`aa-step-nav__item${currentStep === 1 ? " active" : ""}`}>
1. 视频 / 配音 / 文案
</span>
<span className={`aa-step-nav__item${currentStep === 2 ? " active" : ""}`}>
2. 对口型 / 标题 / 封面 / 生成
</span>
</div>
<div className="aa-page-body">
{/* ════ 步骤 1:出镜视频 / 配音库 / 文案 ════ */}
{currentStep === 1 && (
@@ -13,6 +13,8 @@ import CloneModal from "@/components/voice/CloneModal"
import VoiceSelectModal from "./components/VoiceSelectModal"
import ScriptSelectModal from "./components/ScriptSelectModal"
import TtsVoiceModal from "./components/TtsVoiceModal"
import GenerateHeader from "./components/GenerateHeader"
import GenerateStepsBar from "./components/GenerateStepsBar"
import GenerateStepContent from "./components/GenerateStepContent"
import GenerateStepActions from "./components/GenerateStepActions"
import { useGenerateFormState } from "./hooks/useGenerateFormState"
@@ -86,6 +88,7 @@ const GeneratePage: React.FC = () => {
style,
autoSubtitles,
bgm,
editPlanId,
sourceEditPlanId,
previewTaskId,
setPreviewTaskId,
@@ -520,6 +523,10 @@ const GeneratePage: React.FC = () => {
return (
<div className="xx-generate-page">
<GenerateHeader fromEditPlan={!!editPlanId} />
<GenerateStepsBar currentStep={currentStep} onStepClick={setCurrentStep} />
<div className={layoutClassName}>
{/* ════ 步骤1~2 表单 / 步骤3 标题设置 / 步骤4 确认生成进度 / 步骤5 封面 ════ */}
<div className="xx-generate-form">
@@ -11,12 +11,7 @@
* 防止长标题在窄列里溢出导致与相邻卡片进度条视觉重叠。
*/
import React from "react"
import {
LoadingOutlined,
CheckCircleFilled,
CloseCircleOutlined,
ClockCircleOutlined,
} from "@ant-design/icons"
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
import type { GeneratedVideo } from "@/api/template-editor"
@@ -66,11 +61,6 @@ 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"
@@ -95,21 +85,6 @@ 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,7 +12,6 @@ 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"
@@ -152,12 +151,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
selectedCoverTemplate,
onSelectedCoverTemplateChange,
onConfirmGenerate,
selectedVoice,
onSelectedVoiceChange,
voiceModePerVideo,
onVoiceModePerVideoChange,
voiceLibraryIds,
onVoiceLibraryIdsChange,
} = props
switch (currentStep) {
@@ -191,46 +184,32 @@ 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
}
/>
{/* 批量配音选择:共用/独立切换(#2096) */}
<Step3VoiceWithMode
previewCount={previewCount}
selectedVoice={selectedVoice}
onSelectedVoiceChange={onSelectedVoiceChange}
voiceModePerVideo={voiceModePerVideo}
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
voiceLibraryIds={voiceLibraryIds}
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
/>
</>
<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
}
/>
)
case 4:
return (
@@ -157,7 +157,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
/>
</div>
) : (
/* ── 批量:N 个独立标题输入框(两列布局 #2096) ── */
/* ── 批量:N 个独立标题输入框 ── */
<div className="xx-batch-titles">
<div
style={{
@@ -170,25 +170,17 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
为每个视频输入独立标题。标题样式(字体/颜色/位置)全局统一。
</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>
{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>
)}
@@ -17,7 +17,7 @@ import CoverSettingsModal from "./cover-settings/CoverSettingsModal"
import CoverEditorModal from "./cover-settings/CoverEditorModal"
import { useSharedCover } from "@/components/cover/useSharedCover"
import { generateCover as apiGenerateCover } from "@/api/generation"
import { uploadAssetDirect } from "@/api/assets"
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
interface Step6CoverSettingsProps {
coverSettings: CoverConfig
@@ -176,8 +176,15 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
thumbnail_url: previewUrl,
mode: "upload",
})
// 后端自动在默认项目下确保图片素材库存在(P0 404 修复)
const result = await uploadAssetDirect({ file, kind: "image" })
// 查找图片素材库(复用批量封面的逻辑)
const libs = await getAssetLibraries()
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
if (!imageLib) {
hide()
message.error("未找到素材库,请先创建图片素材库")
return previewUrl
}
const result = await uploadAssetDirect({ file, library_id: imageLib.id })
const realUrl = result?.url || ""
if (!realUrl) {
hide()
@@ -10,7 +10,7 @@ export interface BatchTaskState {
taskId: string
/** 变体序号(0-based,与标题/封面数组对齐) */
variantIndex: number
status: "running" | "completed" | "awaiting_cover" | "failed" | "queued"
status: "running" | "completed" | "awaiting_cover" | "failed"
progress: number
error: string | null
/** 完成后的成片视频 */
@@ -374,42 +374,5 @@ export function useGenerationPolling({
}
}, [])
/**
* 批量队列模式:逐任务追加到轮询队列(支持串行提交、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 }
return { startPolling, startPollingBatch, retryTask, clearTimer }
}
@@ -12,7 +12,7 @@
import { useCallback, useState } from "react"
import { message } from "antd"
import { generateCover } from "@/api/generation"
import { uploadAssetDirect } from "@/api/assets"
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
import type { GeneratedVideo } from "@/api/template-editor"
/** onCoversChange 支持直接传值或函数式 updater(函数式用于串行回写避免闭包覆盖) */
@@ -182,9 +182,15 @@ export function useBatchCovers({
async (index: number, file: File) => {
addUploading(index)
try {
const libs = await getAssetLibraries()
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
if (!imageLib) {
message.error("未找到素材库,请先创建")
return
}
const result = await uploadAssetDirect({
file,
kind: "image",
library_id: imageLib.id,
})
const url = result?.url || ""
if (url) {
@@ -2,12 +2,10 @@
* 视频生成 Hook
* 封装视频生成的核心逻辑、状态管理、轮询等
*/
import { useState, useCallback, useEffect, useRef } from "react"
import { useState, useCallback, useEffect } 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"
@@ -17,26 +15,6 @@ 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
@@ -53,22 +31,6 @@ 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 || []
@@ -96,6 +58,7 @@ 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)
@@ -118,30 +81,21 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
[onGenerationSuccess],
)
const handleFailed = useCallback((errorMsg: string) => {
setGenerating(false)
setGenerateError(errorMsg)
}, [])
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
按变体索引排序,供步骤6封面按勾选顺序逐个取视频。
当全部任务都已结束(completed/awaiting_cover/failed)且无排队/渲染中任务时,关闭 generating。 */
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
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 === "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
if (
t.status === "completed" ||
(t.status === "awaiting_cover" && t.videos && t.videos.length > 0)
) {
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
}
})
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
@@ -151,185 +105,17 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
}
return ordered
})
if (allDone && !hasQueued && !hasRunning) {
setGenerating(false)
if (hasSuccess) {
setGenerated(true)
setCompletionStatus("awaiting_cover")
}
}
}, [batchTasks])
const { startPolling, pollBatchTaskQueued, retryTask, clearTimer } = useGenerationPolling({
const { startPolling, startPollingBatch, 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) {
@@ -337,8 +123,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
return false
}
cancelledRef.current = false
clearQueueTimers()
setGenerating(true)
setProgress(0)
setGenerated(false)
@@ -349,16 +133,25 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
setCurrentTaskId("")
clearTimer()
const basePayload = buildBasePayload()
const assetIds = basePayload.asset_ids
const isBatch = (props.previewCount || 1) > 1
try {
// from-assets 兜底:片段不存在则补一次
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。
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 {
@@ -366,185 +159,221 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
}
}
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,请稍后在任务列表查看")
}
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 {
setCurrentTaskId(taskIds[0])
startPolling(taskIds[0])
} catch (err) {
hide()
throw err
}
return true
} catch (err) {
hide()
throw err
}
/* ── 批量:支持任意数量视频,按队列容量串行提交,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
}
}, [
props,
clearTimer,
startPolling,
selectedTemplate,
buildBasePayload,
handleBatchTaskUpdate,
clearQueueTimers,
pollBatchTaskQueued,
])
return true
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
const retry = useCallback(() => {
setGenerateError(null)
@@ -554,10 +383,9 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
/** 第5步:单独重试某个失败任务 */
const retryBatchTask = useCallback(
(taskId: string) => {
handleBatchTaskUpdate(taskId, { status: "running", progress: 0, error: null, videos: [] })
retryTask(taskId)
},
[retryTask, handleBatchTaskUpdate],
[retryTask],
)
const dismissError = useCallback(() => {
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,264 +0,0 @@
/**
* 爆款视频素材选择弹窗(通用版,支持 image/video/voice)
* 基于 ai-avatar 的 ModalAssetPicker 改造:
* - kind 可传 "image" | "video" | "voice"
* - 多图场景 multiple=true 时底部"确认选择"
* - 单选场景点击即回调关闭
*/
import { useEffect, useState } from "react"
import { getAssets, getAssetLibraries, type AssetItem, type AssetLibraryItem } from "@/api/assets"
export interface AssetPickerModalProps {
open: boolean
kind: "image" | "video" | "voice"
multiple?: boolean
title?: string
onClose: () => void
onSelect: (assets: AssetItem[]) => void
}
const KIND_LABEL: Record<AssetPickerModalProps["kind"], string> = {
image: "图片",
video: "视频",
voice: "音频",
}
const MIME_KIND: Record<AssetPickerModalProps["kind"], string> = {
image: "image",
video: "video",
voice: "audio",
}
export default function AssetPickerModal({
open,
kind,
multiple = false,
title,
onClose,
onSelect,
}: AssetPickerModalProps) {
const [keyword, setKeyword] = useState("")
const [libraries, setLibraries] = useState<AssetLibraryItem[]>([])
const [libraryId, setLibraryId] = useState<string>("")
const [assets, setAssets] = useState<AssetItem[]>([])
const [picked, setPicked] = useState<Set<string>>(new Set())
const [loadingLibs, setLoadingLibs] = useState(false)
const [loadingAssets, setLoadingAssets] = useState(false)
const [error, setError] = useState("")
useEffect(() => {
if (!open) return
setKeyword("")
setLibraries([])
setLibraryId("")
setAssets([])
setError("")
setPicked(new Set())
}, [open])
useEffect(() => {
if (!open) return
let cancelled = false
setLoadingLibs(true)
getAssetLibraries(kind)
.then((libs) => {
if (cancelled) return
const list = Array.isArray(libs) ? libs : []
setLibraries(list)
if (list.length > 0) setLibraryId(list[0].id)
})
.catch(() => {
if (!cancelled) setError("素材库加载失败,请重试")
})
.finally(() => {
if (!cancelled) setLoadingLibs(false)
})
return () => {
cancelled = true
}
}, [open, kind])
useEffect(() => {
if (!open || !libraryId) return
let cancelled = false
setLoadingAssets(true)
const load = async () => {
try {
const { items } = await getAssets(libraryId, { page_size: 100 })
if (cancelled) return
let list = Array.isArray(items) ? items : []
const mimePrefix = MIME_KIND[kind]
list = list.filter((a) => !a.mime_type || a.mime_type.startsWith(mimePrefix))
const kw = keyword.trim()
if (kw) list = list.filter((a) => a.name?.includes(kw))
setAssets(list)
setError("")
} catch {
if (!cancelled) {
setError("素材加载失败,请重试")
setAssets([])
}
} finally {
if (!cancelled) setLoadingAssets(false)
}
}
const timer = window.setTimeout(load, 250)
return () => {
cancelled = true
window.clearTimeout(timer)
}
}, [open, libraryId, keyword, kind])
const thumbFor = (a: AssetItem) => {
if (kind === "image") return a.thumbnail_url || a.file_url
if (kind === "video") return a.thumbnail_url
return ""
}
const togglePick = (id: string) => {
if (multiple) {
setPicked((prev) => {
const n = new Set(prev)
if (n.has(id)) n.delete(id)
else n.add(id)
return n
})
} else {
const asset = assets.find((a) => a.id === id)
if (asset) {
onSelect([asset])
onClose()
}
}
}
const handleConfirm = () => {
const list = assets.filter((a) => picked.has(a.id))
if (list.length > 0) onSelect(list)
onClose()
}
if (!open) return null
return (
<div className="vv-modal-mask" onClick={onClose}>
<div className="vv-modal" onClick={(e) => e.stopPropagation()}>
<div className="vv-modal-head">
<span className="vv-modal-title">{title || `选择${KIND_LABEL[kind]}素材`}</span>
<button className="vv-modal-close" onClick={onClose} aria-label="关闭">
×
</button>
</div>
<div className="vv-modal-body">
<div className="vv-asset-search">
<select
className="vv-input"
style={{ width: 170, flex: "0 0 auto" }}
value={libraryId}
onChange={(e) => setLibraryId(e.target.value)}
disabled={loadingLibs || libraries.length === 0}
>
{libraries.length === 0 ? (
<option value="">
{loadingLibs ? "加载中…" : `暂无${KIND_LABEL[kind]}素材库`}
</option>
) : (
libraries.map((lib) => (
<option key={lib.id} value={lib.id}>
📁 {lib.name}
</option>
))
)}
</select>
<input
className="vv-input"
type="text"
placeholder={`搜索${KIND_LABEL[kind]}名称…`}
value={keyword}
onChange={(e) => setKeyword(e.target.value)}
/>
</div>
{libraries.length === 0 && !loadingLibs ? (
<div className="vv-modal-empty">
<div className="vv-empty-icon">📁</div>
暂无{KIND_LABEL[kind]}素材库,请先在「素材库」中创建并上传
</div>
) : loadingAssets ? (
<div className="vv-modal-empty">
<div className="vv-empty-icon">⏳</div>
素材加载中…
</div>
) : error ? (
<div className="vv-modal-empty">
<div className="vv-empty-icon">⚠️</div>
{error}
</div>
) : assets.length === 0 ? (
<div className="vv-modal-empty">
<div className="vv-empty-icon">
{kind === "image" ? "🖼️" : kind === "video" ? "🎬" : "🎵"}
</div>
{kind === "voice" ? (
<>
<div style={{ marginTop: 8, fontSize: 13 }}>暂无配音素材</div>
<div style={{ marginTop: 4, fontSize: 12, color: "#9ca3af" }}>
请先在「配音/我的音色」中上传音频文件,或在素材库管理中添加
</div>
</>
) : (
<>该素材库暂无{KIND_LABEL[kind]}素材</>
)}
</div>
) : (
<div className={`vv-asset-thumbs vv-asset-${kind}`}>
{assets.map((asset) => {
const active = picked.has(asset.id)
const thumb = thumbFor(asset)
return (
<div
key={asset.id}
className={`vv-thumb-card${active ? " selected" : ""}`}
onClick={() => togglePick(asset.id)}
>
{thumb ? (
<img src={thumb} alt={asset.name} />
) : kind === "video" ? (
<video src={asset.file_url} muted preload="metadata" />
) : (
<div className="vv-thumb-ph">{kind === "voice" ? "🎵" : "📄"}</div>
)}
{active && <div className="vv-thumb-check">✓</div>}
<div className="vv-thumb-name" title={asset.name}>
<span className="vv-thumb-name-txt">{asset.name}</span>
{kind === "voice" &&
typeof asset.duration === "number" &&
asset.duration > 0 && (
<span className="vv-thumb-dur">{Math.round(asset.duration)}s</span>
)}
</div>
</div>
)
})}
</div>
)}
</div>
{multiple && (
<div className="vv-modal-foot">
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
取消
</button>
<button
className="vv-btn vv-btn-primary"
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
onClick={handleConfirm}
disabled={picked.size === 0}
>
确认选择({picked.size})
</button>
</div>
)}
</div>
</div>
)
}
@@ -1,355 +0,0 @@
/**
* 内置音色选择弹窗(浅色紫调版)
* - 标题「选择音色」+ 搜索框 + 分类筛选 + 3列卡片网格 + 试听 + 选中 + 完成选择
*/
import React, { useEffect, useMemo, useRef, useState } from "react"
import {
CloseOutlined,
SearchOutlined,
PlayCircleOutlined,
PauseCircleOutlined,
UserOutlined,
} from "@ant-design/icons"
import { Select, Input } from "antd"
export interface PresetVoice {
id: string
name: string
gender?: "female" | "male" | "child" | "other"
gender_label?: string
category?: string
avatar_url?: string
sample_audio_url?: string
desc?: string
}
interface Props {
open: boolean
voices?: PresetVoice[]
loading?: boolean
selectedId?: string
onClose: () => void
onConfirm: (voice: PresetVoice) => void
}
/** 兜底 mock 音色(后端 /api/v1/tts/presets 返回字段不够时使用) */
const MOCK_VOICES: PresetVoice[] = [
// ⚠️ 兜底 mock,仅在 /voices/presets 接口不可达时使用;ID 必须与后端
// packages/domain/preset_voices.py PRESET_VOICES 的 voice_id 对齐(v3后缀)
{
id: "longxiaochun_v3",
name: "龙小淳",
gender: "female",
category: "女声",
desc: "知性积极女声,适合语音助手",
},
{
id: "longxiaoxia_v3",
name: "龙小夏",
gender: "female",
category: "女声",
desc: "沉稳权威女声,适合新闻播报",
},
{
id: "longsanshu_v3",
name: "龙三叔",
gender: "male",
category: "男声",
desc: "沉稳质感男声,适合有声书",
},
{
id: "longyue_v3",
name: "龙悦",
gender: "female",
category: "女声",
desc: "温暖磁性女声,适合广告配音",
},
{
id: "longshu_v3",
name: "龙书",
gender: "male",
category: "男声",
desc: "沉稳青年男声,适合教育讲解",
},
{
id: "longyingjing_v3",
name: "龙应静",
gender: "female",
category: "女声",
desc: "低调冷静女声,适合纪录片解说",
},
{
id: "longshuo_v3",
name: "龙硕",
gender: "male",
category: "男声",
desc: "博才干练男声,适合科技类内容",
},
{
id: "longtian_v3",
name: "龙甜",
gender: "female",
category: "女声",
desc: "活泼女声,适合短视频配音",
},
]
const CATEGORY_LABELS: Record<string, string> = {
all: "全部分类",
female: "女声",
male: "男声",
child: "童声",
dialect: "方言",
emotion: "情绪",
}
const GENDER_LABEL = (v: PresetVoice) => {
if (v.gender_label) return v.gender_label
const g = v.gender
if (g === "female") return "女声·女声"
if (g === "male") return "男声·男声"
if (g === "child") return "童声·童声"
return "性别未标注·其他"
}
const AVATAR_BG = (gender?: string) => {
if (gender === "female") return "#fce7f3"
if (gender === "male") return "#dbeafe"
if (gender === "child") return "#fef3c7"
return "#f3f0ff"
}
const AVATAR_COLOR = (gender?: string) => {
if (gender === "female") return "#be185d"
if (gender === "male") return "#1d4ed8"
if (gender === "child") return "#b45309"
return "#7c3aed"
}
const PresetVoicePickerModal: React.FC<Props> = ({
open,
voices,
loading,
selectedId,
onClose,
onConfirm,
}) => {
const [keyword, setKeyword] = useState("")
const [category, setCategory] = useState<string>("all")
const [pickedId, setPickedId] = useState<string | undefined>(selectedId)
const [playingId, setPlayingId] = useState<string | null>(null)
const audioRef = useRef<HTMLAudioElement | null>(null)
useEffect(() => {
if (open) {
setKeyword("")
setCategory("all")
setPickedId(selectedId)
setPlayingId(null)
}
}, [open, selectedId])
// 停止播放
useEffect(() => {
return () => {
audioRef.current?.pause()
audioRef.current = null
}
}, [])
// 合并真实数据和 mock:如果真实数据 gender/category 缺失,用 mock 兜底
const allVoices: PresetVoice[] = useMemo(() => {
// 真实 API 返回的 voice_id 以 API 为准(如 longxiaochun_v3),前端不做硬编码覆盖
const realList: PresetVoice[] = (voices || []).map((v) => {
// 按 id 精确匹配 mock 获取补充元信息(id 即 voice_id,唯一稳定键)
const mockMatch = MOCK_VOICES.find((m) => m.id === v.id)
return {
...v,
gender: v.gender || mockMatch?.gender,
category:
v.category ||
mockMatch?.category ||
(v.gender === "female" ? "女声" : v.gender === "male" ? "男声" : "其他"),
desc: v.desc || mockMatch?.desc,
sample_audio_url: v.sample_audio_url,
}
})
// 如果没有真实数据,使用兜底 mock(接口失败时)
return realList.length > 0 ? realList : MOCK_VOICES
}, [voices])
const categories = useMemo(() => {
const set = new Set<string>()
allVoices.forEach((v) => {
if (v.category) set.add(v.category)
})
return Array.from(set)
}, [allVoices])
const filtered = useMemo(() => {
const kw = keyword.trim().toLowerCase()
return allVoices.filter((v) => {
if (category !== "all") {
if (v.category !== category && category !== CATEGORY_LABELS[v.gender || ""]) {
// gender 兜底匹配
if (
!(category === "女声" && v.gender === "female") &&
!(category === "男声" && v.gender === "male") &&
!(category === "童声" && v.gender === "child") &&
!(category === "方言" && v.category === "方言") &&
!(category === "情绪" && v.category === "情绪")
) {
return false
}
}
}
if (!kw) return true
return (
v.name?.toLowerCase().includes(kw) ||
v.desc?.toLowerCase().includes(kw) ||
v.category?.toLowerCase().includes(kw)
)
})
}, [allVoices, keyword, category])
const handlePreview = (v: PresetVoice) => {
if (!v.sample_audio_url) {
// 无示例音频
return
}
if (playingId === v.id) {
audioRef.current?.pause()
setPlayingId(null)
return
}
audioRef.current?.pause()
const a = new Audio(v.sample_audio_url)
a.onended = () => setPlayingId(null)
a.onerror = () => setPlayingId(null)
a.play().catch(() => {})
audioRef.current = a
setPlayingId(v.id)
}
const handleConfirm = () => {
const picked = allVoices.find((v) => v.id === pickedId)
if (!picked) return
onConfirm(picked)
}
if (!open) return null
return (
<div className="vv-modal-mask" onClick={onClose}>
<div className="vv-modal vv-modal-lg" onClick={(e) => e.stopPropagation()}>
<div className="vv-modal-head">
<div className="vv-modal-title">选择音色</div>
<button className="vv-modal-close" onClick={onClose}>
<CloseOutlined />
</button>
</div>
<div className="vv-modal-body">
{/* 搜索 */}
<Input
className="vv-voice-search"
placeholder="搜索音色名称或风格"
prefix={<SearchOutlined style={{ color: "#9ca3af" }} />}
value={keyword}
onChange={(e) => setKeyword(e.target.value)}
allowClear
size="large"
/>
{/* 分类筛选 */}
<div className="vv-voice-cat-row">
<span className="vv-voice-cat-label">音色分类</span>
<Select
value={category}
onChange={setCategory}
style={{ width: 180 }}
options={[
{ value: "all", label: "全部分类" },
...[
"女声",
"男声",
"童声",
"方言",
"情绪",
...categories.filter(
(c) => !["女声", "男声", "童声", "方言", "情绪"].includes(c),
),
].map((c) => ({ value: c, label: c })),
]}
/>
</div>
{/* 卡片网格 */}
<div className="vv-voice-grid">
{loading && filtered.length === 0 ? (
<div className="vv-modal-empty">加载中…</div>
) : filtered.length === 0 ? (
<div className="vv-modal-empty">没有匹配的音色</div>
) : (
filtered.map((v) => {
const isPicked = pickedId === v.id
const isPlaying = playingId === v.id
return (
<div
key={v.id}
className={`vv-voice-card ${isPicked ? "selected" : ""}`}
onClick={() => setPickedId(v.id)}
>
<div
className="vv-voice-card-avatar"
style={{ background: AVATAR_BG(v.gender), color: AVATAR_COLOR(v.gender) }}
>
{v.avatar_url ? (
<img src={v.avatar_url} alt={v.name} />
) : (
<UserOutlined style={{ fontSize: 22 }} />
)}
</div>
<div className="vv-voice-card-name" title={v.name}>
{v.name}
</div>
<div className="vv-voice-card-gender">{GENDER_LABEL(v)}</div>
{v.desc && <div className="vv-voice-card-desc">{v.desc}</div>}
<div className="vv-voice-card-actions">
<button
className={`vv-voice-card-btn ${isPicked ? "picked" : ""}`}
onClick={(e) => {
e.stopPropagation()
setPickedId(v.id)
}}
>
{isPicked ? "✓ 已选择" : "选择"}
</button>
<button
className={`vv-voice-card-btn vv-voice-card-btn-preview ${isPlaying ? "playing" : ""} ${!v.sample_audio_url ? "disabled" : ""}`}
onClick={(e) => {
e.stopPropagation()
handlePreview(v)
}}
disabled={!v.sample_audio_url}
>
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
{isPlaying ? "停止" : "试听"}
</button>
</div>
</div>
)
})
)}
</div>
</div>
<div className="vv-modal-foot">
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
取消
</button>
<button className="vv-btn vv-btn-primary" onClick={handleConfirm} disabled={!pickedId}>
完成选择
</button>
</div>
</div>
</div>
)
}
export default PresetVoicePickerModal
@@ -1,74 +0,0 @@
import { useCallback, useEffect, useRef } from "react"
import { getViralVideoJob } from "@/api/viral-video"
import { isAnalysisStage, type ViralVideoJob, type ViralVideoStatus } from "@/api/viral-video/types"
const TERMINAL: ViralVideoStatus[] = ["completed", "failed", "cancelled"]
export interface UseViralVideoPollingOptions {
/** 轮询间隔(毫秒),默认 1500 */
intervalMs?: number
}
/**
* 爆款视频任务 HTTP 轮询 hook。
* 负责持续拉取任务状态并回调给上层;上层负责根据状态/阶段切换 UI 文案。
* 任务进入终态(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
// 视频渲染阶段(Seedance 多段视频生成较慢)拉长轮询间隔
const inRender = job.progress_stage === "rendering"
// 分析阶段走默认间隔即可
const isAnalyzing = isAnalysisStage(job.progress_stage)
const nextDelay = inRender ? 3000 : isAnalyzing ? 2000 : intervalMs
timerRef.current = setTimeout(() => pollOnce(id), nextDelay)
} catch (_err) {
failCountRef.current += 1
if (stoppedRef.current) return
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 }
}
@@ -1,26 +1,25 @@
import { useState, useCallback } from "react"
import { useMutation, useQueryClient } from "@tanstack/react-query"
import { message } from "antd"
import { uploadAssetDirect, getIngestJob, type AssetLibraryItem } from "@/api/assets"
import {
uploadAssetDirect,
getAssetLibraries,
getIngestJob,
type AssetLibraryItem,
} from "@/api/assets"
import { tagAsset } from "@/api/tags"
import { type VoiceGender, type VoiceMaterial } from "../../../types"
interface UseVoiceUploadOptions {
voiceLibrary?: { id: string; kind: string }
createLibMutation?: {
mutateAsync: () => Promise<AssetLibraryItem>
isPending: boolean
}
createLibMutation: { mutateAsync: () => Promise<AssetLibraryItem>; isPending: boolean }
}
/**
* 配音素材上传 Hook
* 封装上传流程:获取库 → 上传文件 → 获取时长 → 创建记录 → 打标签
*/
export function useVoiceUpload({
voiceLibrary,
createLibMutation: _createLibMutation,
}: UseVoiceUploadOptions) {
export function useVoiceUpload({ voiceLibrary, createLibMutation }: UseVoiceUploadOptions) {
const queryClient = useQueryClient()
const [uploadProgress, setUploadProgress] = useState<number | null>(null)
@@ -34,12 +33,24 @@ export function useVoiceUpload({
}) => {
setUploadProgress(0)
try {
// 1. 上传文件:后端自动在默认项目下确保配音库存在(P0 404 修复)
// 兼容 voiceLibrary 参数:若调用方已传入正确的库 ID 则直接复用,否则内部自动解析
// 1. 获取或等待 voice library
let lib = voiceLibrary
if (!lib) {
if (createLibMutation.isPending) {
await createLibMutation.mutateAsync()
}
const libs = await queryClient.fetchQuery({
queryKey: ["asset-libraries"],
queryFn: () => getAssetLibraries(),
})
lib = libs.find((l: AssetLibraryItem) => l.kind === "voice")
if (!lib) throw new Error("无法创建配音库")
}
// 2. 上传文件(带进度,后端自动创建 ingest job)
const complete = await uploadAssetDirect({
file: data.file,
library_id: voiceLibrary?.id,
kind: "voice",
library_id: lib.id,
onProgress: (p) => setUploadProgress(p),
})
@@ -1,6 +1,6 @@
import { useState, useCallback } from "react"
import { useMutation, useQueryClient } from "@tanstack/react-query"
import { uploadAssetDirect, getIngestJob } from "@/api/assets"
import { uploadAssetDirect, getAssetLibraries, getIngestJob } from "@/api/assets"
/**
* 配音上传 Hook
@@ -23,10 +23,18 @@ export function useVoiceUpload({ showToast }: UseVoiceUploadProps) {
mutationFn: async (data: { file: File; name: string; description: string }) => {
setUploadProgress(0)
try {
/* 直传文件(后端会自动在默认项目下确保配音库存在,P0 404 修复) */
/* 获取或创建默认配音库 */
const libs = await queryClient.fetchQuery({
queryKey: ["asset-libraries"],
queryFn: () => getAssetLibraries(),
})
const lib = libs.find((l) => l.kind === "voice")
if (!lib) throw new Error("配音库不存在,请先在配音库页面创建")
/* 直传文件(后端会自动创建 ingest job) */
const complete = await uploadAssetDirect({
file: data.file,
kind: "voice",
library_id: lib.id,
onProgress: (p) => setUploadProgress(p),
})
-4
View File
@@ -52,10 +52,6 @@ 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")),
-45
View File
@@ -1,45 +0,0 @@
import { describe, it, expect } from "vitest"
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
describe("api/errors", () => {
it("returns string error directly", () => {
expect(getErrorMessage("plain")).toBe("plain")
})
it("uses Error.message", () => {
expect(getErrorMessage(new Error("boom"))).toBe("boom")
})
it("returns fallback for empty/unknown", () => {
expect(getErrorMessage(null)).toBe("操作失败,请稍后重试")
expect(getErrorMessage(undefined, "f")).toBe("f")
})
it("reads axios-like response.data.detail", () => {
const err = { response: { data: { detail: "后端报错" } }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("后端报错")
})
it("reads axios-like response.data.message", () => {
const err = { response: { data: { message: "消息字段" } }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("消息字段")
})
it("HTTP 404 fallback", () => {
const err = { response: { status: 404, data: null }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("404")
})
it("HTTP 401 fallback", () => {
const err = { response: { status: 401, data: null }, isAxiosError: true }
expect(getErrorMessage(err)).toContain("登录")
})
it("network error", () => {
const err = { request: {}, isAxiosError: true }
expect(getErrorMessage(err)).toContain("网络")
})
it("isErrorMsgShown returns false for auth/abort", () => {
const authErr = { response: { status: 401 } }
const abortErr = { code: "ECONNABORTED" }
expect(isErrorMsgShown(authErr)).toBe(false)
expect(isErrorMsgShown(abortErr)).toBe(false)
const e: any = new Error("x")
e.__msgShown = true
expect(isErrorMsgShown(e)).toBe(true)
expect(isErrorMsgShown(new Error("x"))).toBe(false)
})
})
-226
View File
@@ -1,226 +0,0 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import {
generateViralVideo,
getViralVideoJob,
confirmViralVideoIntent,
retryViralVideo,
getViralVideoHistory,
getViralStyleTemplates,
analyzeViralStyle,
mockImageAnalysis,
mockGenerateCopy,
analyzeViralImages,
generateViralCopy,
confirmViralCopy,
} from "@/api/viral-video"
import {
VALID_DURATIONS,
VALID_RATIOS,
isVideoStage,
isImageAnalysisStage,
isCopyStage,
isAnalysisStage,
} from "@/api/viral-video/types"
const mockGet = vi.fn()
const mockPost = vi.fn()
vi.mock("@/api/client", () => ({
default: {
get: (...args: unknown[]) => mockGet(...args),
post: (...args: unknown[]) => mockPost(...args),
},
}))
vi.mock("antd", () => ({ message: { error: vi.fn(), success: vi.fn() } }))
// 让 setTimeout 同步执行,避免测试等待 1.8s/2.2s
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
mockGet.mockResolvedValue({ data: {} })
mockPost.mockResolvedValue({ data: {} })
})
afterEach(() => {
vi.useRealTimers()
})
describe("viral-video constants & stage helpers", () => {
afterEach(() => {
vi.useRealTimers()
})
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
})
it("VALID_DURATIONS/VALID_RATIOS", () => {
expect(VALID_DURATIONS).toEqual([5, 10, 15, 20, 25, 30])
expect(VALID_RATIOS).toEqual(expect.arrayContaining(["9:16", "16:9", "1:1"]))
})
it("isVideoStage", () => {
expect(isVideoStage("tts")).toBe(true)
expect(isVideoStage("rendering")).toBe(true)
expect(isVideoStage("uploading")).toBe(true)
expect(isVideoStage("script_generation")).toBe(false)
expect(isVideoStage("completed")).toBe(false)
expect(isVideoStage(undefined)).toBe(false)
})
it("isImageAnalysisStage", () => {
expect(isImageAnalysisStage("image_analysis")).toBe(true)
expect(isImageAnalysisStage("video_analysis")).toBe(true)
expect(isImageAnalysisStage("script_generation")).toBe(false)
expect(isImageAnalysisStage(undefined)).toBe(false)
})
it("isCopyStage", () => {
expect(isCopyStage("intent_parsing")).toBe(true)
expect(isCopyStage("script_generation")).toBe(true)
expect(isCopyStage("review")).toBe(true)
expect(isCopyStage("tts")).toBe(false)
})
it("isAnalysisStage is union", () => {
expect(isAnalysisStage("image_analysis")).toBe(true)
expect(isAnalysisStage("script_generation")).toBe(true)
expect(isAnalysisStage("tts")).toBe(false)
expect(isAnalysisStage(undefined)).toBe(false)
})
})
describe("viral-video API wrappers", () => {
afterEach(() => {
vi.useRealTimers()
})
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
mockGet.mockResolvedValue({ data: {} })
mockPost.mockResolvedValue({ data: {} })
})
it("generateViralVideo", async () => {
mockPost.mockResolvedValue({ data: { id: "j1" } })
const r = generateViralVideo({ images: ["img1"] } as never)
vi.runAllTimersAsync()
expect(await r).toEqual({ id: "j1" })
expect(mockPost).toHaveBeenCalledWith("/viral-video/generate", { images: ["img1"] })
})
it("getViralVideoJob", async () => {
mockGet.mockResolvedValue({ data: { id: "j2" } })
const r = getViralVideoJob("j2")
vi.runAllTimersAsync()
expect(await r).toEqual({ id: "j2" })
expect(mockGet).toHaveBeenCalledWith("/viral-video/j2")
})
it("confirmViralVideoIntent", async () => {
mockPost.mockResolvedValue({ data: { id: "j3" } })
const r = confirmViralVideoIntent("j3", { confirmed_copy: "hi" })
vi.runAllTimersAsync()
await r
expect(mockPost).toHaveBeenCalledWith("/viral-video/j3/confirm-intent", {
confirmed_copy: "hi",
})
})
it("retryViralVideo", async () => {
mockPost.mockResolvedValue({ data: { id: "j4" } })
await retryViralVideo("j4")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j4/retry")
})
it("getViralVideoHistory", async () => {
mockGet.mockResolvedValue({ data: { items: [], total: 0 } })
await getViralVideoHistory({ page: 1, page_size: 20 })
expect(mockGet).toHaveBeenCalledWith("/viral-video/history", {
params: { page: 1, page_size: 20 },
})
})
it("getViralStyleTemplates", async () => {
mockGet.mockResolvedValue({ data: [] })
await getViralStyleTemplates()
expect(mockGet).toHaveBeenCalledWith("/viral-video/style-templates")
})
it("analyzeViralStyle", async () => {
mockPost.mockResolvedValue({ data: { id: "j5" } })
await analyzeViralStyle("j5")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j5/analyze-style")
})
it("analyzeViralImages", async () => {
mockPost.mockResolvedValue({ data: { id: "j6" } })
await analyzeViralImages({ images: ["a.png"] } as never)
expect(mockPost).toHaveBeenCalledWith("/viral-video/analyze-images", { images: ["a.png"] })
})
it("generateViralCopy", async () => {
mockPost.mockResolvedValue({ data: { id: "j7" } })
await generateViralCopy("j7", { duration: 15 } as never)
expect(mockPost).toHaveBeenCalledWith("/viral-video/j7/generate-copy", { duration: 15 })
})
it("confirmViralCopy", async () => {
mockPost.mockResolvedValue({ data: { id: "j8" } })
await confirmViralCopy("j8", { edited_copy: "xxx" })
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", { edited_copy: "xxx" })
mockPost.mockClear()
await confirmViralCopy("j8")
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", {})
})
})
describe("viral-video client mocks", () => {
beforeEach(() => {
vi.useFakeTimers()
vi.clearAllMocks()
})
afterEach(() => {
vi.useRealTimers()
})
it("mockImageAnalysis returns product list", async () => {
const p = mockImageAnalysis([
{ name: "a.png" },
{ name: "b.jpg" },
{ name: "c.webp" },
{ name: "d.png" },
])
vi.advanceTimersByTime(2000)
const r = await p
expect(r.products).toHaveLength(3)
expect(r.products[0].image_index).toBe(0)
expect(r.products[0].brand).toBe("示例品牌")
expect(r.products[1].spec).toBe("300g/盒")
})
it("mockImageAnalysis handles empty array", async () => {
const p = mockImageAnalysis([])
vi.advanceTimersByTime(2000)
const r = await p
expect(r.products).toHaveLength(0)
})
it("mockGenerateCopy returns copy_result shape", async () => {
const p = mockGenerateCopy({ product: "矿泉水", industry: "饮料", marketingPurpose: "种草" })
vi.advanceTimersByTime(3000)
const r = await p
expect(r.title).toContain("种草")
expect(r.title).toContain("矿泉水")
expect(r.final_copy.length).toBeGreaterThan(50)
expect(r.suggested_copy).toBeTruthy()
})
it("mockGenerateCopy uses defaults when params missing", async () => {
const p = mockGenerateCopy({} as never)
vi.advanceTimersByTime(3000)
const r = await p
expect(r.title).toContain("品牌种草")
expect(r.final_copy).toContain("这款产品")
})
})
@@ -1,21 +0,0 @@
import { describe, it, expect } from "vitest"
import { getGenerationPhase } from "@/pages/generate/hooks/generate-video/phase"
describe("getGenerationPhase", () => {
it("returns 分析素材与配置 for p<20", () => {
expect(getGenerationPhase(0)).toEqual({ label: "分析素材与配置", icon: "🔍" })
expect(getGenerationPhase(19).label).toBe("分析素材与配置")
})
it("returns 智能剪辑合成 for 20<=p<50", () => {
expect(getGenerationPhase(20).label).toBe("智能剪辑合成")
expect(getGenerationPhase(49).label).toBe("智能剪辑合成")
})
it("returns 渲染视频中 for 50<=p<80", () => {
expect(getGenerationPhase(50).label).toBe("渲染视频中")
expect(getGenerationPhase(79).label).toBe("渲染视频中")
})
it("returns 即将完成 for p>=80", () => {
expect(getGenerationPhase(80)).toEqual({ label: "即将完成", icon: "✨" })
expect(getGenerationPhase(100).label).toBe("即将完成")
})
})
@@ -1,26 +0,0 @@
import { describe, it, expect, vi, afterEach } from "vitest"
import { formatDuration, formatFileSize, formatDate } from "@/pages/products/detailUtils"
describe("products/detailUtils", () => {
afterEach(() => {
vi.useRealTimers()
})
it("formatDuration", () => {
expect(formatDuration(0)).toBe("00:00")
expect(formatDuration(-1)).toBe("00:00")
expect(formatDuration(5)).toBe("00:05")
expect(formatDuration(65)).toBe("01:05")
expect(formatDuration(3600)).toBe("60:00")
})
it("formatFileSize MB/GB", () => {
expect(formatFileSize(0)).toBe("-")
expect(formatFileSize(-1)).toBe("-")
expect(formatFileSize(5.3)).toBe("5.3 MB")
expect(formatFileSize(2048)).toBe("2.00 GB")
})
it("formatDate returns zh-CN format", () => {
vi.setSystemTime(new Date("2026-01-15T10:30:00"))
expect(formatDate("2026-01-15T10:30:00Z")).toMatch(/2026/)
expect(formatDate("")).toBe("-")
})
})
@@ -1,54 +0,0 @@
import { describe, it, expect, beforeEach, vi, afterEach } from "vitest"
import { renderHook, act } from "@testing-library/react"
import { useViralVideoPolling } from "@/pages/viral-video/hooks/useViralVideoPolling"
const getViralVideoJobMock = vi.fn()
vi.mock("@/api/viral-video", () => ({
getViralVideoJob: (...args: unknown[]) => getViralVideoJobMock(...args),
}))
describe("useViralVideoPolling", () => {
beforeEach(() => {
vi.clearAllMocks()
vi.useFakeTimers()
})
afterEach(() => {
vi.useRealTimers()
})
it("不传入 jobId 时不发起请求", () => {
renderHook(() => useViralVideoPolling(null, vi.fn()))
expect(getViralVideoJobMock).not.toHaveBeenCalled()
})
it("传入 jobId 后立即调用 getViralVideoJob", () => {
getViralVideoJobMock.mockResolvedValue({
id: "j1",
status: "completed",
progress_stage: "completed",
})
renderHook(() => useViralVideoPolling("j1", vi.fn()))
expect(getViralVideoJobMock).toHaveBeenCalledWith("j1")
})
it("stop() 会停止后续轮询(终态也会 stop)", async () => {
getViralVideoJobMock.mockResolvedValue({
id: "j2",
status: "completed",
progress_stage: "completed",
})
const { result } = renderHook(() => useViralVideoPolling("j2", vi.fn(), { intervalMs: 50 }))
// 等第一次 promise 完成
await act(async () => {
await Promise.resolve()
await Promise.resolve()
})
// 终态后不会再调度新请求
const calls = getViralVideoJobMock.mock.calls.length
act(() => {
vi.advanceTimersByTime(2000)
})
expect(getViralVideoJobMock).toHaveBeenCalledTimes(calls)
expect(result.current.stop).toBeTypeOf("function")
})
})
@@ -1,35 +0,0 @@
import { describe, it, expect } from "vitest"
import {
genderLabel,
languageLabel,
genderClass,
formatTime,
formatFileSize,
} from "@/pages/voices/utils/format"
describe("voices utils/format", () => {
it("genderLabel returns label or falls back to value", () => {
expect(genderLabel("female")).toContain("女")
expect(genderLabel("male")).toContain("男")
expect(genderLabel("unknown" as never)).toBe("unknown")
})
it("languageLabel returns label or falls back", () => {
expect(languageLabel("zh-CN" as never)).toBeTruthy()
expect(languageLabel("xx-XX" as never)).toBe("xx-XX")
})
it("genderClass returns css class", () => {
expect(genderClass("female")).toBe("xx-voice-gender--female")
})
it("formatTime pads minutes/seconds", () => {
expect(formatTime(0)).toBe("00:00")
expect(formatTime(5)).toBe("00:05")
expect(formatTime(65)).toBe("01:05")
expect(formatTime(3600)).toBe("60:00")
})
it("formatFileSize human-readable", () => {
expect(formatFileSize(0)).toBe("0 B")
expect(formatFileSize(512)).toBe("512 B")
expect(formatFileSize(2048)).toBe("2.0 KB")
expect(formatFileSize(2 * 1024 * 1024)).toBe("2.0 MB")
})
})
+1 -2
View File
@@ -28,12 +28,11 @@ export default defineConfig({
"src/pages/editing-planner/EditingPlanner.tsx",
"src/pages/assets/AssetLibrary.tsx",
"src/pages/voice-materials/VoiceMaterialLibrary.tsx",
"src/pages/viral-video/ViralVideoPage.tsx",
],
// CI 覆盖率门禁(Phase 4 后提升,逐步逼近目标)
// 当前实际:行 ~62% / 分支 ~61% / 函数 ~25%
thresholds: {
lines: 49,
lines: 50,
branches: 50,
functions: 20,
},
+1 -14
View File
@@ -553,20 +553,7 @@ def concat_video_files(
if work_dir is None:
work_dir = output_path.parent
# Bug #2110: 探测每段是否真实包含音频流,避免 Seedance 生成的无声片段
# (gen_audio=False)让 concat filter `a=1` 找不到 [N:a] 而报 exit 234。
from video_processing.ffmpeg_utils import probe_has_audio as _probe_has_audio
segments: list[ConcatSegment] = []
for p in video_paths:
if not p:
continue
try:
has_audio = _probe_has_audio(p)
except Exception:
has_audio = True # 探测失败保守认为有音频
segments.append(ConcatSegment(video_path=p, has_audio=has_audio))
segments = [ConcatSegment(video_path=p) for p in video_paths if p]
config = ConcatConfig(segments=segments, force_reencode=force_reencode)
engine = ConcatEngine(work_dir)
-33
View File
@@ -1,33 +0,0 @@
"""爆款视频 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",
]
-961
View File
@@ -1,961 +0,0 @@
"""参考爆款视频风格分析模块(#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)
+312 -905
View File
File diff suppressed because it is too large Load Diff
-37
View File
@@ -104,43 +104,6 @@ Staging 当前可以保持 no-op;Production 开启前必须先验证 SMTP/Redi
---
## Staging 服务器 Docker 凭证配置
Staging 服务器(116.62.226.203)需要配置 ACR 和 Gitea Registry 凭证,否则 docker pull 和 Watchtower 自动更新会失败。
### 凭证文件位置
- Docker 配置文件:`/root/.docker/config.json`
- 包含两个 registry 的认证信息:
- `xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com`(阿里云 ACR)
- `git.xiaoxiajianji.com`(Gitea 容器镜像仓库)
### 服务器迁移后恢复步骤
```bash
# 1. 登录 ACR
docker login xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com -u <ACR_USERNAME>
# 2. 登录 Gitea Registry
docker login git.xiaoxiajianji.com -u xiaoxia -p <GITEA_REGISTRY_TOKEN>
# 3. 重启 Watchtower(确保挂载最新 config.json)
docker restart watchtower
```
### Watchtower 配置
- 容器名:`watchtower`
- 检查间隔:300 秒(5 分钟)
- 监控容器:`xiaoxia-api-staging`、`xiaoxia-worker-staging`、`xiaoxia-web-staging`
- 必须挂载 `-v /root/.docker/config.json:/config.json` 才能拉取私有镜像
- 必须挂载 `-v /var/run/docker.sock:/var/run/docker.sock` 才能管理容器
- 容器使用 `:dev` 稳定 tag,Watchtower 通过检测 `:dev` tag 的 digest 变化来发现更新
### 镜像 Tag 策略
- CI 每次构建推送三种 tag:`${GITHUB_SHA}`(精确版本)、`${GITHUB_REF_NAME}`(分支名)、`:dev`(滚动 tag,仅 develop 分支)
- Staging 容器统一使用 `:dev` tag 启动,确保 Watchtower 能自动发现新版本
- Migration(alembic)使用 commit SHA tag 执行,不依赖 Watchtower
---
## Gitea Actions 约定
- `develop` 分支触发 staging 部署。
+1 -5
View File
@@ -30,10 +30,6 @@ COPY deploy/configs/douyin_cookies.txt /app/configs/douyin_cookies.txt
# 强制升级 yt-dlp 到最新(抖音反爬经常变更,旧版 cookies 支持失效;#1968/#1963)
RUN pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com --upgrade "yt-dlp>=2026.8.19"
# API 启动入口(幂等迁移 + uvicorn)—— #2129: watchtower 自动部署兜底
COPY infra/docker/entrypoint-api.sh /usr/local/bin/entrypoint-api.sh
RUN chmod +x /usr/local/bin/entrypoint-api.sh
# 设置环境变量
ENV PATH="/opt/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
ENV PYTHONPATH=/app:/app/apps/api
@@ -45,4 +41,4 @@ HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"
# API 入口点
ENTRYPOINT ["/usr/local/bin/entrypoint-api.sh"]
CMD ["uvicorn", "apps.api.main:app", "--host", "0.0.0.0", "--port", "8000"]
-22
View File
@@ -1,22 +0,0 @@
#!/bin/bash
# API 启动入口:先幂等执行数据库迁移,再启动传入的 CMD(默认 uvicorn)
# 解决 watchtower 自动拉取新镜像后容器重启、未跑 alembic upgrade head 导致新列缺失 500 的问题(#2129)
set -e
cd /app
echo "[entrypoint-api] Running alembic upgrade head..."
if alembic upgrade head; then
echo "[entrypoint-api] Migrations ok."
else
echo "[entrypoint-api] WARNING: alembic upgrade failed, continuing (existing columns should be fine)..." >&2
fi
# 若有显式 CMD(CI 部署时 docker compose run --rm api sh -c '...' 传入),直接 exec 它
if [ "$#" -gt 0 ]; then
echo "[entrypoint-api] Exec custom command: $*"
exec "$@"
fi
echo "[entrypoint-api] Starting uvicorn..."
exec uvicorn apps.api.main:app --host 0.0.0.0 --port 8000
-11
View File
@@ -18,17 +18,6 @@
set -e
# #2129: 幂等执行数据库迁移(watchtower 自动部署兜底)
# worker 容器独立启动,不能依赖 API 容器先跑迁移
cd /app
echo "[entrypoint-worker] Running alembic upgrade head..."
if alembic upgrade head; then
echo "[entrypoint-worker] Migrations ok."
else
echo "[entrypoint-worker] WARNING: alembic upgrade failed, continuing to start workers..." >&2
fi
cd - >/dev/null
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
# ── 并发计算:显式 env 优先;否则从 WORKER_CONCURRENCY 按比例推导 ──
-1
View File
@@ -34,7 +34,6 @@ ENV APP_VERSION=$APP_VERSION
# 复制文件(按变化频率从低到高排序,最大化层缓存命中)
COPY alembic.ini /app/alembic.ini
COPY alembic/ /app/alembic/
COPY migrations/ /app/migrations/
COPY packages/ /app/packages/
# PR #1844 起,worker 还需要加载 apps.api.app.tasks.lipsync_tts,
@@ -502,19 +502,8 @@ class SQLAlchemyAssetRepository:
return [self._to_domain(m) for m in models]
def find_by_storage_key(self, storage_key: str) -> Asset | None:
"""按 storage_key 查找素材。
Bug #2110: 历史数据 file_url 列可能是旧路径(assets/...),新代码统一写入
storage_key 列。双列 OR 查询,避免占位 asset 因路径错配导致 ingest 兜底新建
第二条 READY 记录,原占位卡 PROCESSING → 前端缩略图出现后消失。
"""
if not storage_key:
return None
model = (
self.session.query(AssetModel)
.filter((AssetModel.storage_key == storage_key) | (AssetModel.file_url == storage_key))
.first()
)
"""按 storage_key(对应 DB 中的 file_url)查找素材。"""
model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first()
if model is None:
return None
return self._to_domain(model)
@@ -943,20 +943,9 @@ class ViralVideoJobModel(Base):
style_strength = Column(String(20), nullable=False, default="medium")
style_guide = Column(JSON, nullable=True)
style_template_id = Column(String(36), nullable=False, default="", index=True)
# v1.5 音频/视频参数
voice_id = Column(String(200), nullable=False, default="")
voice_source = Column(String(20), nullable=False, default="")
video_ratio = Column(String(10), nullable=False, default="9:16")
video_model = Column(String(100), nullable=False, default="")
# 结果与状态
status = Column(String(30), nullable=False, default="pending", index=True)
intent_result = Column(JSON, nullable=True)
image_analysis = Column(JSON, nullable=True)
storyboard = Column(JSON, nullable=True)
generated_copy_text = Column(Text, nullable=False, default="")
copy_result = Column(
JSON, nullable=True
) # v1.6: 编导脚本结构{overview,scene_and_lighting,shots,hard_constraints,negative_prompts,voiceover_script}
result_video_url = Column(String(1000), nullable=False, default="")
credits_cost = Column(Integer, nullable=False, default=0)
error_msg = Column(Text, nullable=False, default="")
@@ -81,41 +81,6 @@ def ensure_database_exists(database_url: str) -> None:
admin_engine.dispose()
_VIRAL_VIDEO_BACKFILL_COLS = [
("storyboard", "JSON"),
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
("copy_result", "JSON"),
]
def _ensure_viral_video_columns(connection) -> None:
"""Idempotently add new columns to viral_video_jobs; create_all will not ALTER existing tables."""
from sqlalchemy import inspect as _inspect
try:
insp = _inspect(connection)
if not insp.has_table("viral_video_jobs"):
return
existing = {c["name"] for c in insp.get_columns("viral_video_jobs")}
except Exception:
return
import logging as _logging
_log = _logging.getLogger(__name__)
for col, ddl in _VIRAL_VIDEO_BACKFILL_COLS:
if col in existing:
continue
try:
connection.execute(text(f"ALTER TABLE viral_video_jobs ADD COLUMN {col} {ddl}"))
_log.info("added column viral_video_jobs.%s", col)
except Exception as e:
_log.warning("add column %s failed: %s", col, e)
def initialize_database(engine) -> None:
"""初始化数据库 schema。
@@ -135,5 +100,4 @@ def initialize_database(engine) -> None:
text("SELECT pg_advisory_unlock(:lock_id)"),
{"lock_id": SCHEMA_INIT_LOCK_ID},
)
_ensure_viral_video_columns(connection)
connection.commit()
@@ -24,7 +24,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
viral_structure=model.viral_structure or "",
marketing_purpose=model.marketing_purpose or "",
bgm_preference=model.bgm_preference or "",
duration=model.duration or 15,
duration=model.duration or 30,
user_copy_text=model.user_copy_text or "",
fusion_level=model.fusion_level or "ai_polish",
reference_audio_path=model.reference_audio_path or "",
@@ -32,16 +32,8 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
style_strength=getattr(model, "style_strength", "medium") or "medium",
style_guide=dict(model.style_guide) if model.style_guide else None,
style_template_id=getattr(model, "style_template_id", "") or "",
voice_id=getattr(model, "voice_id", "") or "",
voice_source=getattr(model, "voice_source", "") or "",
video_ratio=getattr(model, "video_ratio", "9:16") or "9:16",
video_model=getattr(model, "video_model", "") or "",
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
intent_result=dict(model.intent_result) if model.intent_result else None,
image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None,
storyboard=list(model.storyboard) if getattr(model, "storyboard", None) else None,
generated_copy_text=getattr(model, "generated_copy_text", "") or "",
copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None,
result_video_url=model.result_video_url or "",
credits_cost=model.credits_cost or 0,
error_msg=model.error_msg or "",
@@ -78,16 +70,8 @@ class SQLAlchemyViralVideoJobRepository:
style_strength=job.style_strength,
style_guide=job.style_guide,
style_template_id=job.style_template_id,
voice_id=job.voice_id,
voice_source=job.voice_source,
video_ratio=job.video_ratio,
video_model=job.video_model,
status=job.status,
intent_result=job.intent_result,
image_analysis=job.image_analysis,
storyboard=job.storyboard,
generated_copy_text=job.generated_copy_text,
copy_result=job.copy_result,
result_video_url=job.result_video_url,
credits_cost=job.credits_cost,
error_msg=job.error_msg,
@@ -107,10 +91,6 @@ class SQLAlchemyViralVideoJobRepository:
raise ValueError(f"ViralVideoJob {job.id} not found")
model.status = job.status
model.intent_result = job.intent_result
model.image_analysis = job.image_analysis
model.storyboard = job.storyboard
model.generated_copy_text = job.generated_copy_text or ""
model.copy_result = job.copy_result
model.result_video_url = job.result_video_url
model.credits_cost = job.credits_cost
model.error_msg = job.error_msg
@@ -118,24 +98,6 @@ class SQLAlchemyViralVideoJobRepository:
model.started_at = job.started_at
model.completed_at = job.completed_at
model.style_guide = job.style_guide
# v1.5 three-stage: persist user-editable params so resume uses latest values
model.user_copy_text = job.user_copy_text
model.industry = job.industry
model.target_customer = job.target_customer
model.persona_id = job.persona_id
model.viral_structure = job.viral_structure
model.marketing_purpose = job.marketing_purpose
model.bgm_preference = job.bgm_preference
model.duration = job.duration
model.fusion_level = job.fusion_level
model.reference_audio_path = job.reference_audio_path
model.reference_video_url = job.reference_video_url
model.style_strength = job.style_strength
model.style_template_id = job.style_template_id
model.voice_id = job.voice_id or ""
model.voice_source = job.voice_source or ""
model.video_ratio = job.video_ratio or "9:16"
model.video_model = job.video_model or ""
model.updated_at = datetime.now(timezone.utc)
self.session.commit()
@@ -161,9 +123,7 @@ class SQLAlchemyViralVideoJobRepository:
self.session.query(ViralVideoJobModel)
.filter(
ViralVideoJobModel.user_id == user_id,
ViralVideoJobModel.status.in_(
["pending", "running", "wait_user_confirm", "image_analyzed", "copy_generated"]
),
ViralVideoJobModel.status.in_(["pending", "running", "wait_user_confirm"]),
)
.count()
)
-3
View File
@@ -96,9 +96,6 @@ 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 = ""
+34 -90
View File
@@ -1,11 +1,10 @@
"""ViralVideoJob 领域模型 — 爆款视频任务.
v1.6 重大简化:Seedance 2.5 单次最长30秒,单次调用直接出片,不再分段/拼接/ffmpeg concat。
状态机(三步分步):
pending -> running -> image_analyzed -> running -> copy_generated -> running -> completed
wait_user_confirm -> running -> completed (旧路径兼容)
任意阶段 fail; 任意非终态 cancel.
failed -> pending (retry 重置后重跑)。
状态机:
pending → running → completed
↘ failed → pending (retry)
↘ cancelled
running 中可暂停:running → wait_user_confirm → running (confirm-intent resume)
"""
from __future__ import annotations
@@ -27,10 +26,10 @@ from uuid import uuid4
class ViralVideoStatus(StrEnum):
"""爆款视频任务状态枚举。"""
PENDING = "pending"
RUNNING = "running"
IMAGE_ANALYZED = "image_analyzed"
COPY_GENERATED = "copy_generated"
WAIT_USER_CONFIRM = "wait_user_confirm"
COMPLETED = "completed"
FAILED = "failed"
@@ -38,32 +37,44 @@ class ViralVideoStatus(StrEnum):
class ViralVideoStage(StrEnum):
"""编排流水线阶段枚举(用于 WS 进度推送)。"""
IMAGE_ANALYSIS = "image_analysis"
VIDEO_ANALYSIS = "video_analysis"
INTENT_PARSING = "intent_parsing"
SCRIPT_GENERATION = "script_generation" # v1.6: 编导分镜脚本(融合原 copy_fusion+storyboard+review)
COPY_FUSION = "copy_fusion"
STORYBOARD = "storyboard"
REVIEW = "review"
TTS = "tts"
RENDERING = "rendering" # v1.6: 单次 Seedance 生成(BGM/音效/画面一次出片)
BGM_SELECT = "bgm_select"
RENDERING = "rendering"
MUSETALK = "musetalk"
UPLOADING = "uploading"
class FusionLevel(StrEnum):
"""文案融合级别。"""
AI_FULL = "ai_full"
AI_POLISH = "ai_polish"
USER_PRIMARY = "user_primary"
class StyleStrength(StrEnum):
"""风格强度。"""
LIGHT = "light"
MEDIUM = "medium"
STRICT = "strict"
class PromptType(StrEnum):
"""Prompt 模板类型(与 #2040 seed 对齐)。"""
IMAGE_ANALYSIS = "image_analysis"
INTENT_PARSING = "intent_parsing"
SCRIPT_GENERATION = "script_generation"
COPY_FUSION = "copy_fusion"
STORYBOARD = "storyboard"
REVIEW = "review"
VIDEO_STYLE_INTEGRATION = "video_style_integration"
STYLE_CONSTRAINT = "style_constraint"
@@ -75,17 +86,20 @@ STAGE_LABELS = {
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
ViralVideoStage.INTENT_PARSING: "意图解析",
ViralVideoStage.SCRIPT_GENERATION: "编导脚本生成",
ViralVideoStage.COPY_FUSION: "文案融合",
ViralVideoStage.STORYBOARD: "分镜脚本",
ViralVideoStage.REVIEW: "合规审核",
ViralVideoStage.TTS: "AI 配音",
ViralVideoStage.RENDERING: "视频生成",
ViralVideoStage.BGM_SELECT: "BGM 选择",
ViralVideoStage.RENDERING: "视频渲染",
ViralVideoStage.MUSETALK: "数字人口型",
ViralVideoStage.UPLOADING: "上传发布",
}
@dataclass
class ViralVideoJob:
"""爆款视频任务领域实体(v1.6 单次 Seedance 出片版)。"""
"""爆款视频任务领域实体。"""
user_id: str
images: list[str] = field(default_factory=list)
@@ -95,28 +109,19 @@ class ViralVideoJob:
viral_structure: str = ""
marketing_purpose: str = ""
bgm_preference: str = ""
duration: int = 15 # v1.6: 默认15秒,上限30秒(Seedance 2.5 单次最大30s)
duration: int = 30
user_copy_text: str = ""
fusion_level: str = FusionLevel.AI_POLISH
reference_audio_path: str = ""
# v1.3
reference_video_url: str = ""
style_strength: str = StyleStrength.MEDIUM
style_guide: dict | None = None
style_template_id: str = ""
# v1.5.1 音频/视频参数
voice_id: str = ""
voice_source: str = ""
video_ratio: str = "9:16"
video_model: str = ""
# v1.4+ 产物
image_analysis: dict | None = None
intent_result: dict | None = None
generated_copy_text: str = "" # v1.6: 存 voiceover_script(纯口播对白),字段名兼容
storyboard: list | None = None # v1.6: 存 copy_result.shots,字段名兼容
copy_result: dict | None = None # v1.6: 完整编导脚本结构
# 状态
id: str = field(default_factory=lambda: uuid4().hex)
status: ViralVideoStatus = ViralVideoStatus.PENDING
intent_result: dict | None = None
result_video_url: str = ""
credits_cost: int = 0
error_msg: str = ""
@@ -126,44 +131,13 @@ class ViralVideoJob:
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
# -- 状态转换 --
# ── 状态转换 ──
def mark_running(self) -> None:
if self.status not in (
ViralVideoStatus.PENDING,
ViralVideoStatus.IMAGE_ANALYZED,
ViralVideoStatus.COPY_GENERATED,
ViralVideoStatus.WAIT_USER_CONFIRM,
ViralVideoStatus.RUNNING,
):
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
raise ValueError(f"Cannot transition from {self.status} to running")
self.status = ViralVideoStatus.RUNNING
if self.started_at is None:
self.started_at = datetime.now(timezone.utc)
self.updated_at = datetime.now(timezone.utc)
def mark_image_analyzed(self) -> None:
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
raise ValueError(f"Cannot transition from {self.status} to image_analyzed")
self.status = ViralVideoStatus.IMAGE_ANALYZED
if self.started_at is None:
self.started_at = datetime.now(timezone.utc)
self.updated_at = datetime.now(timezone.utc)
def mark_copy_generated(self, copy_result: dict) -> None:
"""v1.6 阶段2完成:编导脚本(含 voiceover_script/shots/硬约束/负面词)已生成。"""
if self.status not in (
ViralVideoStatus.IMAGE_ANALYZED,
ViralVideoStatus.RUNNING,
ViralVideoStatus.PENDING,
):
raise ValueError(f"Cannot transition from {self.status} to copy_generated")
self.status = ViralVideoStatus.COPY_GENERATED
self.copy_result = copy_result or {}
if isinstance(copy_result, dict):
self.generated_copy_text = copy_result.get("voiceover_script", "") or ""
shots = copy_result.get("shots") or []
self.storyboard = list(shots) if isinstance(shots, list) else []
self.started_at = datetime.now(timezone.utc)
self.updated_at = datetime.now(timezone.utc)
def mark_wait_user_confirm(self, intent_result: dict) -> None:
@@ -173,25 +147,6 @@ class ViralVideoJob:
self.intent_result = intent_result
self.updated_at = datetime.now(timezone.utc)
def resume_from_image_analyzed(self, **kwargs) -> None:
if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING):
raise ValueError(f"Cannot resume from {self.status} to copy-gen")
for k, v in kwargs.items():
if hasattr(self, k) and v not in (None, "", []):
setattr(self, k, v)
self.status = ViralVideoStatus.RUNNING
self.updated_at = datetime.now(timezone.utc)
def resume_from_copy_generated(self, edited_copy: str | None = None) -> None:
"""阶段2->阶段3:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
if self.status != ViralVideoStatus.COPY_GENERATED:
raise ValueError(f"Cannot resume from {self.status} to render")
if edited_copy and isinstance(self.copy_result, dict):
self.copy_result = {**self.copy_result, "voiceover_script": edited_copy}
self.generated_copy_text = edited_copy
self.status = ViralVideoStatus.RUNNING
self.updated_at = datetime.now(timezone.utc)
def resume_from_confirm(self) -> None:
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
raise ValueError(f"Cannot resume from {self.status}")
@@ -224,14 +179,3 @@ class ViralVideoJob:
ViralVideoStatus.FAILED,
ViralVideoStatus.CANCELLED,
)
@property
def effective_copy_text(self) -> str:
"""TTS 用的最终口播文案:优先 copy_result.voiceover_script,兼容老字段。"""
if isinstance(self.copy_result, dict) and self.copy_result.get("voiceover_script"):
return self.copy_result["voiceover_script"]
return self.generated_copy_text or self.user_copy_text or "你好,给大家推荐一款好物"
@property
def voiceover_script(self) -> str:
return self.effective_copy_text
-217
View File
@@ -13,9 +13,7 @@ API 和 Worker 两边共用。基于火山引擎方舟平台的 OpenAI 兼容接
from __future__ import annotations
import logging
import os
import time
import uuid
from typing import Any, Optional
import httpx
@@ -240,221 +238,6 @@ 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 | None = "9:16",
resolution: str = "720p",
generate_audio: bool = True,
watermark: bool = False,
output_dir: str | None = None,
model: str | None = None,
reference_images: list[str] | None = None,
reference_audios: list[str] | None = None,
reference_videos: list[str] | None = None,
) -> str | None:
"""调用 Seedance 2.5 文生/图生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
v1.6: 支持多参考图(产品素材)+ 参考音频(TTS口型驱动)+ 参考视频,单次生成最长 30 秒。
Args:
prompt: 文本提示词(含完整编导脚本:总览+场景光线+逐镜头时间轴+硬约束+负面词)
image_url: 首帧参考图 URL(可选,提供则走图生视频首帧模式,ratio 跟随首帧)
duration: 视频时长 4~30 秒
ratio: 宽高比 16:9/9:16/1:1/4:3/3:4/21:9/adaptive;image_url 存在时自动忽略
resolution: 480p/720p/1080p
generate_audio: 是否让模型原生合成音效/BGM(v1.6 默认 True,配合 reference_audios 做口型驱动)
watermark: 是否加水印
output_dir: 下载目录,默认 /tmp
model: 指定模型 ID;空则用 settings.doubao_video_model
reference_images: 多参考图 URL 列表(产品素材,最多30张;注意 image_url 为首帧单独传)
reference_audios: 参考音频 URL 列表(TTS口播,驱动口型,最多10段)
reference_videos: 参考视频 URL 列表(风格参考)
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", 900) or 900
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
video_model = model or default_video_model
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
if image_url:
content.append({"type": "image_url", "image_url": {"url": image_url}})
# v1.6: 多参考图(产品素材)
if reference_images:
for url in reference_images[:30]:
if url and isinstance(url, str):
content.append({"type": "image_url", "image_url": {"url": url}})
create_payload: dict[str, Any] = {
"model": video_model,
"content": content,
"generate_audio": bool(generate_audio),
"duration": int(duration),
"resolution": resolution,
"watermark": bool(watermark),
}
# v1.6: 参考音频(TTS 驱动口型)
if reference_audios:
create_payload["reference_audios"] = [
{"url": u, "role": "audio_url"} for u in reference_audios[:10] if u and isinstance(u, str)
]
# v1.6: 参考视频(风格参考)
if reference_videos:
create_payload["reference_videos"] = [{"url": u} for u in reference_videos[:5] if u and isinstance(u, str)]
# Bug #2110 / v1.6: ratio=None 时不传(首帧图生视频跟随原图比例)
if ratio and not image_url:
create_payload["ratio"] = ratio
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
create_url = f"{self.base_url}/contents/generations/tasks"
logger.info(
"Seedance 创建任务请求: url=%s model=%s duration=%ds ratio=%s gen_audio=%s image_url=%s ref_imgs=%d ref_audios=%d ref_videos=%d",
create_url,
video_model,
duration,
ratio or "(follow-image)",
generate_audio,
bool(image_url),
len(reference_images or []),
len(reference_audios or []),
len(reference_videos or []),
)
# 1) 创建任务(带重试)
task_id: str | None = None
last_error: Exception | None = None
for attempt in range(self.max_retries + 1):
try:
resp = httpx.post(create_url, headers=headers, json=create_payload, timeout=self.timeout)
# 测试环境下 MagicMock().status_code 是 MagicMock,与 int 比较会抛 TypeError;
# 用显式 int() 转换+类型判断,避免误判。
try:
_status = int(resp.status_code)
except (TypeError, ValueError):
_status = 200
if _status >= 400:
# 把响应体完整打出来(通常含 error.code/message,能直接定位:模型未开通/Key 无权限/模型 ID 错误)
logger.error(
"Seedance 创建任务 HTTP %d: body=%s",
resp.status_code,
(resp.text or "")[:1000],
)
resp.raise_for_status()
data = resp.json()
task_id = data.get("id")
if task_id:
break
last_error = RuntimeError(f"create task returned no id: {str(data)[:200]}")
except Exception as e:
last_error = e
if attempt < self.max_retries:
wait = 0.5 * (2**attempt)
logger.warning(
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s", wait, attempt + 1, self.max_retries + 1, e
)
time.sleep(wait)
if not task_id:
logger.error(
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s 【排查建议】"
"1) 确认方舟控制台已开通 Doubao-Seedance-2.5 模型;"
"2) DOUBAO_API_KEY 对应的账号有该模型调用权限;"
"3) DOUBAO_BASE_URL 必须为 https://ark.cn-beijing.volces.com/api/v3;"
"4) 若控制台用「推理接入点」(endpoint),请把 DOUBAO_VIDEO_MODEL 改为 ep-xxx 接入点 ID。",
video_model,
self.base_url,
last_error,
)
return None
logger.info(
"Seedance 任务已创建: task_id=%s model=%s duration=%ds gen_audio=%s",
task_id,
video_model,
duration,
generate_audio,
)
# 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)
try:
if int(getattr(resp, "status_code", 200)) >= 400:
resp.raise_for_status()
except (TypeError, ValueError):
pass
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 httpx.HTTPStatusError as e:
last_error = e
logger.warning(
"Seedance 轮询 HTTP %d: body=%s",
e.response.status_code,
(e.response.text or "")[:500],
)
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
# ── 单例 ─────────────────────────────────────────────────────────────────────
+12 -95
View File
@@ -515,107 +515,24 @@ def call_llm(prompt: str, temperature: float = 0.7) -> object:
def call_vision(image_url: str, prompt: str) -> object:
"""调用豆包视觉大模型分析图片,返回解析后的 JSON 或原文字符串;失败返回 None。
Bug #2114 (VLM 牛头不对马嘴根因修复):
之前误走 client.chat_completion(用文本模型 doubao-seed-1.6),多模态 content list 被当成
纯文本发给文本模型 → 模型要么看不到图、要么抛 400,静默被 except 吞掉 → 返回 None →
_step_image_analysis fallback 到 {"name":"未识别"} → 后续文案/分镜完全没图的信息。
现改走 vision_completion,走视觉模型 doubao-1-5-vision-pro-250915。
"""
"""调用豆包视觉大模型分析图片,返回解析后的 JSON 或原文字符串;失败返回 None。"""
client = get_doubao_client()
if not client.is_available:
logger.warning("[call_vision] 豆包客户端未配置 (DOUBAO_API_KEY 缺失)")
return None
if not image_url:
logger.warning("[call_vision] 空 image_url,跳过视觉分析")
return None
system_prompt = (
"你是资深电商视觉分析师。请严格基于用户提供的图片观察回答,"
"图片里没有的信息不要凭空想象或编造;看不清或无法判断时明确说"
"「图片中无法判断」,不要猜测。输出必须是严格 JSON,不要附加 Markdown 或解释文字。"
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
{"role": "system", "content": "你是专业的视觉分析师。需要结构化输出时请严格使用 JSON。"},
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": image_url}},
],
},
]
logger.info(
"[call_vision] 调用豆包视觉模型 vision_model=%s image_url=%s prompt_len=%d",
getattr(client, "vision_model", "?"),
image_url[:120],
len(prompt),
)
raw = client.vision_completion(
messages=messages,
images=[image_url],
temperature=0.2,
max_tokens=2048,
timeout=60,
)
raw = client.chat_completion(messages, temperature=0.3, max_tokens=2048)
if raw is None:
logger.warning("[call_vision] 视觉模型返回 None (image_url=%s)", image_url[:80])
return None
logger.info("[call_vision] 视觉模型原始返回 (前400字): %s", raw[:400])
# 剥离 ```json ... ``` 包裹
stripped = raw.strip()
if stripped.startswith("```"):
stripped = stripped.strip("`")
if stripped.startswith("json"):
stripped = stripped[4:].lstrip()
try:
return json.loads(stripped)
except (json.JSONDecodeError, TypeError) as e:
logger.warning("[call_vision] JSON 解析失败(%s),返回原始文本: %s", e, raw[:200])
return json.loads(raw)
except (json.JSONDecodeError, TypeError):
return raw
def call_video_generation(
prompt: str,
*,
image_url: str | None = None,
duration: int = 15,
ratio: str | None = "9:16",
resolution: str = "720p",
output_dir: str | None = None,
model: str | None = None,
generate_audio: bool = True,
reference_images: list[str] | None = None,
reference_audios: list[str] | None = None,
reference_videos: list[str] | None = None,
) -> str | None:
"""调用 Seedance 2.5 生成视频(v1.6 单次出片版),返回本地 MP4 路径;失败返回 None。
v1.6:
- 默认 generate_audio=True,模型原生合成环境音效/BGM;
- reference_audios 传 TTS 音频 URL 数组做口型驱动;
- reference_images 传产品素材 URL 数组做视觉参考;
- 单次最长 30 秒,不分段不拼接;
- image_url 存在时为「首帧图生视频」模式,自动不传 ratio(Bug #2110)。
"""
client = get_doubao_client()
if not client.is_available:
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
return None
effective_ratio = None if image_url else ratio
try:
kwargs: dict = dict(
prompt=prompt,
image_url=image_url,
duration=int(duration),
resolution=resolution,
generate_audio=bool(generate_audio),
watermark=False,
output_dir=output_dir,
model=model,
reference_images=reference_images,
reference_audios=reference_audios,
reference_videos=reference_videos,
)
if effective_ratio:
kwargs["ratio"] = effective_ratio
return client.video_generation(**kwargs)
except Exception as e:
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
return None
-5
View File
@@ -15,8 +15,3 @@ 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
+4 -41
View File
@@ -137,28 +137,6 @@ fi
echo "✅ compose.yml ready: $COMPOSE_FILE_PATH ($(wc -l < "$COMPOSE_FILE_PATH") lines)"
ln -sf "$NGINX_CONF_FILE" "$INFRA_DOCKER_DIR/nginx-${COMPOSE_ENV_VALUE}.conf" 2>/dev/null || true
# 封装 docker compose 调用:统一 --env-file(compose 默认只读取 project 目录下的 .env,
# 我们的 .env 在 $INFRA_DOCKER_DIR/../../.env,必须显式传入才能读到 GENERATED_FILES_HOST_DIR 等变量)
# ── 防御:清理可能残留的 docker-compose.override.yml / compose.override.yml ──
# 历史上运维曾用 override 文件固定镜像 tag 排查问题,若忘记删除会导致新镜像 tag 不生效,
# Worker 一直跑旧镜像(本次 P0 404 排查中即踩过此坑)。这里每次部署都主动清理。
for override in "$INFRA_DOCKER_DIR/docker-compose.override.yml" "$INFRA_DOCKER_DIR/compose.override.yml" "$INFRA_DOCKER_DIR/override.yml"; do
if [ -f "$override" ]; then
echo "⚠️ Found stale override file, removing: $override"
rm -f "$override"
fi
done
# ── 防御:清理可能残留的 docker-compose.override.yml / compose.override.yml ──
# 历史上运维曾用 override 文件固定镜像 tag 排查问题,若忘记删除会导致新镜像 tag 不生效,
# Worker 一直跑旧镜像(本次 P0 404 排查中即踩过此坑)。这里每次部署都主动清理。
for override in "$INFRA_DOCKER_DIR/docker-compose.override.yml" "$INFRA_DOCKER_DIR/compose.override.yml" "$INFRA_DOCKER_DIR/override.yml"; do
if [ -f "$override" ]; then
echo "⚠️ Found stale override file, removing: $override"
rm -f "$override"
fi
done
# 封装 docker compose 调用:统一 --env-file(compose 默认只读取 project 目录下的 .env,
# 我们的 .env 在 $INFRA_DOCKER_DIR/../../.env,必须显式传入才能读到 GENERATED_FILES_HOST_DIR 等变量)
compose() {
@@ -368,21 +346,6 @@ fi
echo "All images pulled."
# ====== 打稳定 tag(:dev),供 Watchtower 监控 ======
# Watchtower 只能检测同一个 tag 的 digest 变化。
# commit SHA tag 每次构建都不同,Watchtower 无法感知更新。
# 因此每次部署都将最新镜像 tag 为 :dev,容器统一使用 :dev 启动。
DEV_API="${REGISTRY}/xiaoxia-saas-api:dev"
DEV_WORKER="${REGISTRY}/xiaoxia-saas-worker:dev"
DEV_WEB="${REGISTRY}/xiaoxia-saas-web:dev"
docker tag "$REGISTRY_API" "$DEV_API"
docker tag "$REGISTRY_WORKER" "$DEV_WORKER"
docker tag "$REGISTRY_WEB" "$DEV_WEB"
echo "✅ Tagged images as :dev for Watchtower monitoring"
echo " API: $DEV_API"
echo " Worker: $DEV_WORKER"
echo " Web: $DEV_WEB"
# ====== 镜像内容校验 ======
echo ""
echo "=========================================="
@@ -548,7 +511,7 @@ docker run -d \
--health-retries 3 \
--health-start-period 40s \
$LOG_OPTS \
"$DEV_API" &
"$REGISTRY_API" &
PID_API_START=$!
# ── Worker: 通过 compose 启动(单一事实来源)──
@@ -556,7 +519,7 @@ PID_API_START=$!
# healthcheck 匹配 'celery.*worker'(不把 beat 算活)、资源限制 4C/8G。
# WORKER_IMAGE 通过环境变量覆盖镜像 tag(compose.yml 默认 :dev)。
echo "Starting worker via docker compose (from $INFRA_DOCKER_DIR)..."
WORKER_IMAGE="$DEV_WORKER" APP_VERSION="$IMAGE_TAG" compose up -d --no-deps worker &
WORKER_IMAGE="$REGISTRY_WORKER" APP_VERSION="$IMAGE_TAG" compose up -d --no-deps worker &
PID_WORKER_START=$!
# ── Web: 暂保留 docker run(TODO: 后续收敛到 compose)──
@@ -572,7 +535,7 @@ docker run -d \
--health-timeout 5s \
--health-retries 3 \
$LOG_OPTS \
"$DEV_WEB" &
"$REGISTRY_WEB" &
PID_WEB_START=$!
wait $PID_API_START $PID_WORKER_START $PID_WEB_START
@@ -703,5 +666,5 @@ echo "=== Staging deployment complete ==="
echo "API: http://127.0.0.1:8000"
echo "Web: http://127.0.0.1:3001"
echo "Worker: managed by docker compose (project=$COMPOSE_PROJECT)"
echo "Version: $IMAGE_TAG (running as :dev for Watchtower)"
echo "Version: $IMAGE_TAG"
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging
-485
View File
@@ -1,485 +0,0 @@
"""#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_task_resp.status_code = 200
fake_task_resp.text = ""
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()
fake_poll_resp.status_code = 200
fake_poll_resp.text = ""
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
-3
View File
@@ -144,9 +144,6 @@ def _storage():
"expires_at": "2026-01-01T00:00:00Z",
"fields": {"key": "uploads/abc/test.mp4"},
}
# Bug #2110: duplicated 命中时 _get_existing_asset_url 调用 get_url 返回公网 URL 字符串,
# Mock 默认返回 MagicMock,会让 DirectUploadPrepareResponse.url: str 校验失败。
s.get_url.return_value = ""
return s
-360
View File
@@ -1,360 +0,0 @@
"""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"
+53 -104
View File
@@ -120,7 +120,7 @@ class TestViralVideoJobDefaults:
job = ViralVideoJob(user_id="u1")
assert job.images == []
assert job.industry == ""
assert job.duration == 15
assert job.duration == 30
assert job.fusion_level == FusionLevel.AI_POLISH
assert job.style_strength == StyleStrength.MEDIUM
assert job.status == ViralVideoStatus.PENDING
@@ -145,10 +145,13 @@ class TestViralVideoStage:
"image_analysis",
"video_analysis",
"intent_parsing",
"script_generation",
"copy_fusion",
"storyboard",
"review",
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
]
actual_order = [s.value for s in ViralVideoStage]
@@ -168,7 +171,7 @@ class TestViralVideoSchemas:
assert req.images == ["https://example.com/img.jpg"]
assert req.fusion_level == "ai_polish"
assert req.style_strength == "medium"
assert req.duration == 15
assert req.duration == 30
def test_create_request_empty_images_raises(self):
from app.schemas.viral_video import CreateViralVideoRequest
@@ -362,10 +365,9 @@ class TestViralVideoPipeline:
industry="美妆",
target_customer="年轻女性",
marketing_purpose="品牌推广",
duration=15,
duration=30,
user_copy_text="这款产品超好用",
fusion_level="ai_polish",
video_ratio="9:16",
)
@patch("packages.shared.ai_service.call_vision")
@@ -403,110 +405,63 @@ class TestViralVideoPipeline:
assert "intent" in result
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_returns_copy_result(self, mock_llm, mock_job):
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
def test_copy_fusion_ai_polish(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_copy_fusion
mock_llm.return_value = {
"overview": {"theme": "口红推荐", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台,柔和自然光",
"shots": [
{
"time_range": "0-5秒",
"shot_type_angle_movement": "近景平视,缓慢推镜",
"scene_and_dialogue": "女主微笑展示口红:大家好,今天分享一款口红",
"action_details": "手持口红特写",
"audio_bgm": "轻快流行BGM",
"transition": "硬切",
"reference_image_index": 0,
},
{
"time_range": "5-15秒",
"shot_type_angle_movement": "特写,固定镜头",
"scene_and_dialogue": "涂抹口红:颜色特别好看很显白",
"action_details": "嘴唇涂抹特写",
"audio_bgm": "轻快BGM继续",
"transition": "结束",
"reference_image_index": 1,
},
],
"hard_constraints": ["无字幕无水印"],
"negative_prompts": ["字幕", "水印"],
"voiceover_script": "大家好,今天分享一款口红,颜色特别好看很显白。",
}
result = _step_script_generation(
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
)
assert isinstance(result, dict)
assert "voiceover_script" in result
assert "shots" in result
assert isinstance(result["shots"], list)
assert len(result["shots"]) == 2
assert result["overview"]["total_duration"] == 15
# final_copy 必须 = voiceover_script(向后兼容)
assert result.get("final_copy") == result["voiceover_script"]
mock_llm.return_value = "融合后的文案内容"
result = _step_copy_fusion(mock_job, {"intent": "推广"}, {"products": []})
assert isinstance(result, str)
assert len(result) > 0
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_fallback(self, mock_llm, mock_job):
"""LLM 返回异常时使用兜底脚本(不会抛错)。"""
from apps.worker.worker_app.tasks.viral_video import _fallback_script
def test_storyboard_generation(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
result = _fallback_script(mock_job)
assert isinstance(result, dict)
assert result["voiceover_script"]
assert len(result["shots"]) >= 1
mock_llm.return_value = [
{"order": 0, "type": "product_shot", "duration": 10},
{"order": 1, "type": "closing", "duration": 5},
]
result = _step_storyboard(mock_job, "测试文案", {})
assert isinstance(result, list)
assert len(result) == 2
@patch("packages.shared.ai_service.call_llm")
def test_review_pass_v16(self, mock_llm, mock_job):
"""v1.6 _step_review 接收 copy_result dict。"""
def test_review_pass(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_review
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
cr = {"voiceover_script": "大家好", "shots": []}
result = _step_review(mock_job, cr)
result = _step_review(mock_job, "测试文案", [])
assert result["passed"] is True
def test_assemble_seedance_prompt(self, mock_job):
"""编导脚本必须能拼出完整的 Seedance prompt,含总览/场景/逐镜头/约束。"""
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
def test_bgm_select(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
cr = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [
{
"time_range": "0-15秒",
"shot_type_angle_movement": "中景平视",
"scene_and_dialogue": "你好分享",
"action_details": "展示",
"audio_bgm": "BGM",
"transition": "结束",
"reference_image_index": 0,
}
],
"hard_constraints": ["无字幕"],
"negative_prompts": ["水印"],
}
prompt = _assemble_seedance_prompt(cr, mock_job)
assert "【视频总览】" in prompt
assert "【逐镜头时间轴】" in prompt
assert "【硬性约束】" in prompt
assert "【负面提示词】" in prompt
assert "0-15秒" in prompt
mock_job.bgm_preference = "upbeat"
bgm = _step_bgm_select(mock_job)
assert "upbeat" in bgm
def test_bgm_select_default(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
mock_job.bgm_preference = ""
bgm = _step_bgm_select(mock_job)
assert bgm == "bgm_default.mp3"
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
class TestPipelineIntegration:
"""v1.6 流水线端到端集成测试(mock 外部依赖):TTS+单次 Seedance+上传。"""
"""流水线端到端集成测试(mock 外部依赖)。"""
@patch("apps.worker.worker_app.tasks.viral_video._step_upload")
@patch("apps.worker.worker_app.tasks.viral_video._step_musetalk")
@patch("apps.worker.worker_app.tasks.viral_video._step_render")
@patch("apps.worker.worker_app.tasks.viral_video._upload_tts_to_oss")
@patch("apps.worker.worker_app.tasks.viral_video._step_bgm_select")
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
@patch("apps.worker.worker_app.tasks.viral_video._step_review")
@patch("apps.worker.worker_app.tasks.viral_video._step_script_generation")
@patch("apps.worker.worker_app.tasks.viral_video._step_storyboard")
@patch("apps.worker.worker_app.tasks.viral_video._step_copy_fusion")
@patch("apps.worker.worker_app.tasks.viral_video._step_intent_parsing")
@patch("apps.worker.worker_app.tasks.viral_video._step_video_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
@@ -519,47 +474,41 @@ class TestPipelineIntegration:
mock_img_analysis,
mock_video_analysis,
mock_intent,
mock_script,
mock_copy_fusion,
mock_storyboard,
mock_review,
mock_tts,
mock_tts_upload,
mock_bgm,
mock_render,
mock_musetalk,
mock_upload,
):
"""v1.6: TTS整段合成 → 上传TTS到OSS → 单次 Seedance → 上传成片。"""
"""测试 resume 流水线能从确认状态走到完成。"""
from apps.worker.worker_app.tasks.viral_video import (
resume_viral_video_pipeline,
)
# 构造 mock job
job = ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
status=ViralVideoStatus.RUNNING,
intent_result={"intent": "推广"},
duration=15,
video_ratio="9:16",
)
mock_repo = MagicMock()
mock_session = MagicMock()
mock_get_repo.return_value = (mock_session, mock_repo, job)
# v1.6: 如果没有 copy_result 会现场补生成
mock_intent.return_value = {"intent": "推广", "key_messages": [], "tone": "亲切"}
mock_script.return_value = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [],
"hard_constraints": [],
"negative_prompts": [],
"voiceover_script": "大家好,分享一款口红。",
"final_copy": "大家好,分享一款口红。",
}
# 设置各步骤返回值
mock_copy_fusion.return_value = "融合文案"
mock_storyboard.return_value = [{"order": 0, "duration": 10}]
mock_review.return_value = {"passed": True, "score": 90}
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
mock_tts_upload.return_value = None
mock_tts.return_value = "https://audio.mp3"
mock_bgm.return_value = "bgm_default.mp3"
mock_render.return_value = "/tmp/video.mp4"
mock_musetalk.return_value = "/tmp/video_final.mp4"
mock_upload.return_value = "https://oss.example.com/final.mp4"
result = resume_viral_video_pipeline.run("job-001")
-309
View File
@@ -1,309 +0,0 @@
"""#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 TestScriptGenerationV16:
"""v1.6 编导分镜脚本生成相关纯函数测试。"""
def test_fallback_script_has_required_fields(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _fallback_script
out = _fallback_script(mock_job)
assert isinstance(out, dict)
assert "overview" in out
assert "shots" in out
assert "voiceover_script" in out
assert "hard_constraints" in out
assert "negative_prompts" in out
assert out["overview"]["total_duration"] == mock_job.duration
assert out["final_copy"] == out["voiceover_script"]
assert len(out["shots"]) >= 1
def test_safe_json_loads_parses_fenced_code(self):
from apps.worker.worker_app.tasks.viral_video import _safe_json_loads
fenced = '```json\n{"voiceover_script": "你好", "shots": []}\n```'
out = _safe_json_loads(fenced)
assert out is not None
assert out["voiceover_script"] == "你好"
def test_safe_json_loads_handles_none(self):
from apps.worker.worker_app.tasks.viral_video import _safe_json_loads
assert _safe_json_loads(None) is None
assert _safe_json_loads("not json") is None
def test_validate_normalize_fills_defaults(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _validate_and_normalize_script
raw = {"voiceover_script": "你好", "shots": [{"scene_and_dialogue": "测试"}]}
out = _validate_and_normalize_script(raw, mock_job)
assert out["voiceover_script"] == "你好"
assert len(out["shots"]) == 1
assert out["shots"][0]["shot_type_angle_movement"]
assert out["overview"]["total_duration"] == mock_job.duration
def test_assemble_seedance_prompt_contains_sections(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
cr = {
"overview": {"theme": "测试", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮",
"shots": [
{
"time_range": "0-15秒",
"shot_type_angle_movement": "中景",
"scene_and_dialogue": "你好",
"action_details": "展示",
"audio_bgm": "BGM",
"transition": "结束",
"reference_image_index": 0,
}
],
"hard_constraints": ["无字幕"],
"negative_prompts": ["水印"],
}
p = _assemble_seedance_prompt(cr, mock_job)
for key in ("【视频总览】", "【场景与光线】", "【逐镜头时间轴】", "【硬性约束】", "【负面提示词】"):
assert key in p
# ── 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("apps.worker.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("apps.worker.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("apps.worker.services.tts_service_factory.get_tts_service", return_value=fake_service):
result = vv._step_tts(mock_job, "文案")
# Bug #2110: 校验传入了 voice_id+format=mp3
call_kwargs = fake_service.synthesize.call_args.kwargs
assert call_kwargs.get("format") == "mp3"
assert isinstance(result, Path)
assert result.exists()
# ── P1: BGM 跳过 / MuseTalk 无 persona 跳过 ───────────────────────
class TestDurationClamp:
"""v1.6 mark_copy_generated 派生字段 + duration clamp。"""
def test_mark_copy_generated_derives_fields(self):
job = ViralVideoJob(user_id="u1", duration=15)
cr = {
"overview": {"theme": "x", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "亮",
"shots": [{"time_range": "0-15秒", "scene_and_dialogue": "对白"}],
"voiceover_script": "你好",
"hard_constraints": [],
"negative_prompts": [],
}
job.mark_copy_generated(cr)
assert job.copy_result is cr
assert job.generated_copy_text == "你好"
assert job.storyboard == cr["shots"]
assert job.effective_copy_text == "你好"
# ── 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
assert kwargs["generate_audio"] is True
# ── P0-1: _step_render 占位片段生成 ──────────────────────────────
class TestCallVideoGenerationV16:
"""v1.6 call_video_generation 透传 reference_audios/reference_images 等参数到 client。"""
def test_passes_reference_params_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=15,
ratio="9:16",
reference_images=["https://img/r1.jpg"],
reference_audios=["https://oss/tts.mp3"],
reference_videos=["https://oss/ref.mp4"],
generate_audio=True,
model="doubao-seedance-2-5-260628",
)
assert result == str(out)
kwargs = mock_client.video_generation.call_args.kwargs
# 首帧模式不传 ratio(Bug #2110)
assert "ratio" not in kwargs
assert kwargs["image_url"] == "https://img/x.jpg"
assert kwargs["reference_audios"] == ["https://oss/tts.mp3"]
assert kwargs["reference_images"] == ["https://img/r1.jpg"]
assert kwargs["reference_videos"] == ["https://oss/ref.mp4"]
assert kwargs["generate_audio"] is True
assert kwargs["model"] == "doubao-seedance-2-5-260628"
def test_ratio_passed_when_no_image(self, tmp_path):
from packages.shared.ai_service import call_video_generation
out = tmp_path / "v.mp4"
out.write_bytes(b"fake")
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
mock_client = MagicMock()
mock_client.is_available = True
mock_client.video_generation.return_value = str(out)
mock_get.return_value = mock_client
call_video_generation(prompt="测试", duration=10, ratio="16:9")
kwargs = mock_client.video_generation.call_args.kwargs
assert kwargs["ratio"] == "16:9"
assert kwargs["image_url"] is None
# ── 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/render pipeline 应从 job.image_analysis 读(v1.5 _run_render_pipeline 共享渲染逻辑)。"""
import inspect
from apps.worker.worker_app.tasks import viral_video as vv
# v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis
src = inspect.getsource(vv._run_render_pipeline)
assert "job.image_analysis" in src
assert "image_analysis" in src
# resume 本身应该调用 _run_render_pipeline
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
assert "_run_render_pipeline" in resume_src
-399
View File
@@ -1,399 +0,0 @@
"""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": 15,
"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,
"image_analysis": None,
"storyboard": None,
"copy_result": None,
"generated_copy_text": "",
"voice_id": "",
"voice_source": "",
"voice_mode": "global",
"video_ratio": "9:16",
"video_model": "",
"credits_cost": 0,
"updated_at": None,
"is_terminal": False,
"effective_copy_text": "",
"voiceover_script": "",
}.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"
# ── v1.5 three-stage endpoints ─────────────────────────────────────────
class TestAnalyzeImages:
def test_analyze_images_creates_job_and_dispatches(self):
"""POST /analyze-images: 创建任务 + 入队 run_viral_video_analyze。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import AnalyzeImagesRequest
user = _auth_user("u1")
session = MagicMock()
repo = MagicMock()
req = AnalyzeImagesRequest(images=["https://x.com/a.jpg"], reference_video_url="", style_template_id="")
saved = {}
def fake_save(job):
saved["job"] = job
return job
repo.save.side_effect = fake_save
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task") as mock_send,
):
resp = vv_mod.analyze_images(req, authenticated_user=user, session=session)
job = saved["job"]
assert job.user_id == "u1"
assert job.images == ["https://x.com/a.jpg"]
mock_send.assert_called_once_with("worker.run_viral_video_analyze", args=[job.id])
assert resp.status == "pending"
class TestGenerateCopy:
def test_generate_copy_updates_params_and_dispatches(self):
"""POST /{id}/generate-copy: 在 image_analyzed 状态下写营销参数 + 入队 run_viral_video_generate_copy。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
job = _make_job(job_id="job-gc", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED)
repo = MagicMock()
repo.get.return_value = job
req = GenerateCopyRequest(
industry="美妆",
target_customer="年轻女性",
duration=25,
fusion_level="ai_full",
user_copy_text="试试这个",
)
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task") as mock_send,
):
resp = vv_mod.generate_copy("job-gc", req, authenticated_user=user, session=session)
# 参数写入
assert job.industry == "美妆"
assert job.target_customer == "年轻女性"
assert job.duration == 25
assert job.fusion_level == "ai_full"
assert job.user_copy_text == "试试这个"
job.resume_from_image_analyzed.assert_called_once()
repo.update.assert_called()
mock_send.assert_called_once_with("worker.run_viral_video_generate_copy", args=["job-gc"])
assert resp.id == "job-gc"
def test_generate_copy_rejects_wrong_status(self):
"""任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。"""
import pytest
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
from fastapi import HTTPException
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
repo = MagicMock()
repo.get.return_value = job
with (patch.object(vv_mod, "_get_job_repo", return_value=repo),):
with pytest.raises(HTTPException) as exc:
vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session)
assert exc.value.status_code == 409
def test_generate_copy_persists_voice_and_ratio(self):
"""generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
job = _make_job(job_id="job-gc3", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED)
repo = MagicMock()
repo.get.return_value = job
req = GenerateCopyRequest(
voice_id="cosy_voice_001",
voice_source="library",
video_ratio="16:9",
)
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task"),
):
vv_mod.generate_copy("job-gc3", req, authenticated_user=user, session=session)
assert job.voice_id == "cosy_voice_001"
assert job.voice_source == "library"
assert job.video_ratio == "16:9"
class TestAnalyzeImagesPersist:
def test_analyze_images_persists_voice_and_ratio(self):
"""analyze-images 创建任务时应带上 voice/video_ratio 字段。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import AnalyzeImagesRequest
user = _auth_user("u1")
session = MagicMock()
saved = {}
class FakeRepo:
def save(self, job):
saved["job"] = job
def get(self, jid):
return None
req = AnalyzeImagesRequest(
images=["img-1"],
voice_id="preset_v1",
voice_source="preset",
video_ratio="1:1",
)
with (
patch.object(vv_mod, "_get_job_repo", return_value=FakeRepo()),
patch.object(vv_mod.celery_app, "send_task"),
):
resp = vv_mod.analyze_images(req, authenticated_user=user, session=session)
job = saved["job"]
assert job.voice_id == "preset_v1"
assert job.voice_source == "preset"
assert job.video_ratio == "1:1"
assert resp.images == ["img-1"]
class TestConfirmCopy:
def test_confirm_copy_dispatches_render(self):
"""POST /{id}/confirm-copy: copy_generated -> RUNNING + 入队 run_viral_video_render,编辑文案写入。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import ConfirmCopyRequest
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
job = _make_job(job_id="job-cc", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
repo = MagicMock()
repo.get.return_value = job
req = ConfirmCopyRequest(edited_copy="我改了文案")
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task") as mock_send,
):
resp = vv_mod.confirm_copy("job-cc", req, authenticated_user=user, session=session)
job.resume_from_copy_generated.assert_called_once_with(edited_copy="我改了文案")
mock_send.assert_called_once_with("worker.run_viral_video_render", args=["job-cc"])
assert resp.id == "job-cc"
def test_confirm_copy_rejects_wrong_status(self):
import pytest
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import ConfirmCopyRequest
from fastapi import HTTPException
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
job = _make_job(job_id="job-cc2", user_id="u1", status=ViralVideoStatus.IMAGE_ANALYZED)
repo = MagicMock()
repo.get.return_value = job
with patch.object(vv_mod, "_get_job_repo", return_value=repo):
with pytest.raises(HTTPException) as exc:
vv_mod.confirm_copy("job-cc2", ConfirmCopyRequest(), authenticated_user=user, session=session)
assert exc.value.status_code == 409
-432
View File
@@ -1,432 +0,0 @@
"""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)