Compare commits
36 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| fc6ebbecb6 | |||
| 5cefbc9c05 | |||
| 41fe2a96a6 | |||
| f0862934f8 | |||
| 774c4fc0df | |||
| c7f8db383f | |||
| 17a95eb8f0 | |||
| 3fbc1bbfe6 | |||
| fa9545f79b | |||
| 87eb480f3c | |||
| 8bdc39a1ab | |||
| 5e61dbe4f9 | |||
| 22e04d65a7 | |||
| 6ff57b2feb | |||
| 2981d20d5b | |||
| 6cddd72910 | |||
| 6d5c44d6be | |||
| 665a3063b6 | |||
| 24724dca9f | |||
| d08835ec9f | |||
| 77ce4a1a0d | |||
| 7ad722e6c6 | |||
| b54dda6526 | |||
| 69da326ed6 | |||
| f7f600d091 | |||
| bf9249da19 | |||
| ca834b23cb | |||
| 37f7aa3329 | |||
| 794f5f374b | |||
| 34305974ad | |||
| e83a7cad2e | |||
| c45a2ce9b1 | |||
| 9814fcdc22 | |||
| a7d6ba473b | |||
| f19be5fd09 | |||
| 0a004db1bd |
@@ -0,0 +1,25 @@
|
||||
"""viral video add image_analysis column
|
||||
|
||||
Revision ID: 087_viral_video_image_analysis
|
||||
Revises: 086_add_viral_video_tables
|
||||
Create Date: 2026-09-30
|
||||
|
||||
#2106 爆款视频 P0:持久化图片分析结果(image_analysis JSON),供 resume 阶段使用。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "087_viral_video_image_analysis"
|
||||
down_revision = "086_add_viral_video_tables"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("viral_video_jobs", sa.Column("image_analysis", sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("viral_video_jobs", "image_analysis")
|
||||
@@ -191,6 +191,23 @@ 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,
|
||||
@@ -390,6 +407,7 @@ 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]
|
||||
@@ -443,6 +461,7 @@ async def prepare_direct_upload(
|
||||
duplicated=False,
|
||||
skip_transfer=False,
|
||||
asset_id=pending_asset_id,
|
||||
url="",
|
||||
)
|
||||
|
||||
|
||||
|
||||
Executable → Regular
+265
-13
@@ -8,6 +8,7 @@
|
||||
POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案
|
||||
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
|
||||
GET /api/v1/viral-video/style-templates 获取风格模板列表
|
||||
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送(订阅 Redis pub/sub)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -15,6 +16,7 @@ from __future__ import annotations
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.viral_video import (
|
||||
AnalyzeStyleRequest,
|
||||
@@ -26,7 +28,7 @@ from app.schemas.viral_video import (
|
||||
ViralVideoHistoryResponse,
|
||||
ViralVideoJobResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
@@ -121,9 +123,7 @@ def create_viral_video(
|
||||
|
||||
# 入队 Celery 任务
|
||||
try:
|
||||
from worker_app.tasks.viral_video import run_viral_video_pipeline
|
||||
|
||||
run_viral_video_pipeline.delay(job.id)
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 任务已入队: job_id=%s user_id=%s", job.id, job.user_id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 入队失败: %s", e, exc_info=True)
|
||||
@@ -209,9 +209,7 @@ def retry_viral_video_job(
|
||||
|
||||
# 重新入队
|
||||
try:
|
||||
from worker_app.tasks.viral_video import run_viral_video_pipeline
|
||||
|
||||
run_viral_video_pipeline.delay(job.id)
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d", job.id, job.retry_count)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
@@ -248,9 +246,7 @@ def confirm_intent(
|
||||
|
||||
# 从断点恢复 Celery 任务
|
||||
try:
|
||||
from worker_app.tasks.viral_video import resume_viral_video_pipeline
|
||||
|
||||
resume_viral_video_pipeline.delay(job.id)
|
||||
celery_app.send_task("worker.resume_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 意图确认,恢复流水线: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线失败: %s", e, exc_info=True)
|
||||
@@ -283,9 +279,7 @@ def analyze_style(
|
||||
|
||||
# 入队风格分析任务
|
||||
try:
|
||||
from worker_app.tasks.viral_video import run_video_style_analysis
|
||||
|
||||
run_video_style_analysis.delay(job.id)
|
||||
celery_app.send_task("worker.run_video_style_analysis", args=[job.id])
|
||||
logger.info("[爆款视频] 风格分析入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 风格分析入队失败: %s", e, exc_info=True)
|
||||
@@ -295,3 +289,261 @@ 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": "",
|
||||
"wait_user_confirm": "intent_parsing",
|
||||
"completed": "uploading",
|
||||
"failed": "",
|
||||
"cancelled": "",
|
||||
}
|
||||
|
||||
_STATUS_PROGRESS = {
|
||||
"pending": 0.0,
|
||||
"running": 5.0,
|
||||
"wait_user_confirm": 35.0,
|
||||
"completed": 100.0,
|
||||
"failed": 0.0,
|
||||
"cancelled": 0.0,
|
||||
}
|
||||
|
||||
_STATUS_MESSAGE = {
|
||||
"pending": "任务已创建,等待执行",
|
||||
"running": "任务执行中",
|
||||
"wait_user_confirm": "等待用户确认意图文案",
|
||||
"completed": "视频生成完成",
|
||||
"failed": "任务失败",
|
||||
"cancelled": "任务已取消",
|
||||
}
|
||||
|
||||
|
||||
def _stage_from_status(job) -> str:
|
||||
return _STATUS_STAGE.get(_job_status(job), "")
|
||||
|
||||
|
||||
def _estimate_progress(job) -> float:
|
||||
"""根据 status 粗略估算百分比(0-100),用于连接初始快照;
|
||||
连接建立后由 Redis 推送的真实事件持续更新。
|
||||
"""
|
||||
return _STATUS_PROGRESS.get(_job_status(job), 5.0)
|
||||
|
||||
|
||||
def _initial_message(job) -> str:
|
||||
"""给新连接的前端一个可读的初始状态文案。"""
|
||||
status_val = _job_status(job)
|
||||
if status_val == "failed" and job.error_msg:
|
||||
return f"任务失败: {job.error_msg}"
|
||||
return _STATUS_MESSAGE.get(status_val, "任务准备中")
|
||||
|
||||
+2
-2
@@ -7,9 +7,9 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.schema_guard import assert_auto_create_schema_allowed
|
||||
|
||||
ensure_database_exists(settings.DATABASE_URL)
|
||||
ensure_database_exists(settings.effective_database_url)
|
||||
engine, SessionLocal = build_session_factory(
|
||||
settings.DATABASE_URL,
|
||||
settings.effective_database_url,
|
||||
pool_size=settings.DATABASE_POOL_SIZE,
|
||||
max_overflow=settings.DATABASE_MAX_OVERFLOW,
|
||||
pool_timeout=settings.DATABASE_POOL_TIMEOUT,
|
||||
|
||||
@@ -56,7 +56,7 @@ from packages.adapters.sqlalchemy_impl.voice_library_repository import (
|
||||
from packages.ports.tag_repository import TagRepository
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
|
||||
_engine, _SessionLocal = build_session_factory(settings.effective_database_url)
|
||||
|
||||
|
||||
def get_db_session() -> Generator[Session, None, None]:
|
||||
|
||||
@@ -29,6 +29,8 @@ 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):
|
||||
|
||||
@@ -8,7 +8,7 @@ from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
# ── 枚举常量 ─────────────────────────────────────────────────────────────
|
||||
|
||||
VALID_FUSION_LEVELS = ("ai_full", "ai_polish", "user_primary")
|
||||
VALID_FUSION_LEVELS = ("ai_full", "full_ai", "ai_polish", "user_primary")
|
||||
VALID_STYLE_STRENGTHS = ("light", "medium", "strict")
|
||||
VALID_STAGES = (
|
||||
"image_analysis",
|
||||
@@ -50,6 +50,9 @@ class CreateViralVideoRequest(BaseModel):
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
def _validate_fusion_level(cls, v: str) -> str:
|
||||
# 兼容前端历史写法 full_ai(等价 ai_full)
|
||||
if v == "full_ai":
|
||||
return "ai_full"
|
||||
if v not in VALID_FUSION_LEVELS:
|
||||
raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一")
|
||||
return v
|
||||
|
||||
@@ -150,6 +150,8 @@ export interface DirectUploadPrepareResult {
|
||||
* 两个字段是同一语义的别名(后端可能只返回其一),前端任意为 true 即视为命中去重。
|
||||
*/
|
||||
skip_transfer?: boolean
|
||||
/** duplicated=true 时后端返回已存在素材的公网 URL,前端直接用而不必再调 complete */
|
||||
url?: string
|
||||
}
|
||||
|
||||
/** 直传完成确认返回 */
|
||||
|
||||
@@ -3,9 +3,24 @@
|
||||
*/
|
||||
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
|
||||
@@ -108,6 +123,8 @@ const putToOSS = (
|
||||
|
||||
/** 单个文件的上传阶段信息(供批量上传队列做状态绑定) */
|
||||
export interface DirectUploadHandle {
|
||||
/** 实际使用的素材库(内部解析出来,便于调用方做后续 UI/缓存操作) */
|
||||
library: { id: string; kind: "image" | "video" | "voice" }
|
||||
/** prepare 返回(含可能的预建 asset_id) */
|
||||
prepared: DirectUploadPrepareResult
|
||||
/** 直传 OSS(可重复调用用于重试) */
|
||||
@@ -119,10 +136,17 @@ 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
|
||||
library_id: string
|
||||
/** 素材库 ID;未传时按文件类型自动在默认项目下 ensure-default */
|
||||
library_id?: string
|
||||
/** 显式指定素材库 kind;未传时按 MIME/扩展名推断 */
|
||||
kind?: "image" | "video" | "voice"
|
||||
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
|
||||
fileHash?: string
|
||||
/** 本次逻辑上传的幂等 token,prepare/complete 一致、重试复用 */
|
||||
@@ -138,9 +162,17 @@ 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: data.library_id,
|
||||
library_id: resolvedLibraryId,
|
||||
filename: data.file.name,
|
||||
content_type: data.file.type || "application/octet-stream",
|
||||
file_size: data.file.size,
|
||||
@@ -149,12 +181,13 @@ 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: data.library_id,
|
||||
library_id: resolvedLibraryId,
|
||||
storage_key: prepared.storage_key,
|
||||
file_hash: data.fileHash,
|
||||
client_upload_id: data.clientUploadId,
|
||||
@@ -164,10 +197,17 @@ export const prepareDirectUploadHandle = async (data: {
|
||||
}
|
||||
}
|
||||
|
||||
/** 直传上传(大文件推荐),支持可选进度回调;一次性完成 prepare→transfer→complete */
|
||||
/** 直传上传(大文件推荐),支持可选进度回调;一次性完成 prepare→transfer→complete
|
||||
*
|
||||
* P0 404 修复:library_id 可选;不传时内部按文件类型自动匹配正确项目下的素材库,
|
||||
* 保证 project_id 与 library_id 必然一致。
|
||||
*/
|
||||
export const uploadAssetDirect = async (data: {
|
||||
file: File
|
||||
library_id: string
|
||||
/** 素材库 ID;可选,不传按文件类型自动解析默认项目下的对应素材库(推荐用法) */
|
||||
library_id?: string
|
||||
/** 显式指定素材库 kind;未传时按文件 MIME/扩展名推断 */
|
||||
kind?: "image" | "video" | "voice"
|
||||
onProgress?: (percent: number) => void
|
||||
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
|
||||
fileHash?: string
|
||||
@@ -180,6 +220,7 @@ export const uploadAssetDirect = async (data: {
|
||||
const handle = await prepareDirectUploadHandle({
|
||||
file: data.file,
|
||||
library_id: data.library_id,
|
||||
kind: data.kind,
|
||||
fileHash,
|
||||
clientUploadId,
|
||||
})
|
||||
@@ -188,7 +229,7 @@ export const uploadAssetDirect = async (data: {
|
||||
return {
|
||||
storage_key: handle.prepared.storage_key,
|
||||
ingest_job_id: "",
|
||||
url: "",
|
||||
url: handle.prepared.url || "",
|
||||
duplicated: true,
|
||||
asset_id: handle.prepared.asset_id,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
import apiClient from "@/api/client"
|
||||
import type {
|
||||
GenerateViralVideoRequest,
|
||||
HistoryResponse,
|
||||
StyleTemplate,
|
||||
ViralVideoJob,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
export function generateViralVideo(payload: GenerateViralVideoRequest) {
|
||||
return apiClient.post<ViralVideoJob>("/viral-video/generate", payload).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 查询单个任务 */
|
||||
export function getViralVideoJob(id: string) {
|
||||
return apiClient.get<ViralVideoJob>(`/viral-video/${id}`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 用户确认/修改 AI 理解的意图后继续 */
|
||||
export function confirmViralVideoIntent(
|
||||
id: string,
|
||||
payload: { confirmed_copy?: string; edits?: Record<string, unknown> },
|
||||
) {
|
||||
return apiClient
|
||||
.post<ViralVideoJob>(`/viral-video/${id}/confirm-intent`, payload)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 重试失败任务 */
|
||||
export function retryViralVideo(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/retry`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 历史记录(分页) */
|
||||
export function getViralVideoHistory(params?: { page?: number; page_size?: number }) {
|
||||
return apiClient.get<HistoryResponse>("/viral-video/history", { params }).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 预设风格模板 */
|
||||
export function getViralStyleTemplates() {
|
||||
return apiClient.get<StyleTemplate[]>("/viral-video/style-templates").then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 上传参考视频后触发风格分析(返回带 style_guide 的任务详情) */
|
||||
export function analyzeViralStyle(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
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: "深度模仿" },
|
||||
]
|
||||
|
||||
export type ViralVideoStatus =
|
||||
"pending" | "running" | "wait_user_confirm" | "completed" | "failed" | "cancelled"
|
||||
|
||||
/**
|
||||
* 后端流水线阶段字符串。前端不展示逐阶段进度列表,仅保留类型
|
||||
* 用于轮询时判断当前在哪个大阶段(分析中 vs 视频生成中)以选择轮询间隔/文案。
|
||||
*/
|
||||
export type ViralVideoStage =
|
||||
| "image_analysis"
|
||||
| "video_analysis"
|
||||
| "intent_parsing"
|
||||
| "copy_fusion"
|
||||
| "storyboard"
|
||||
| "review"
|
||||
| "tts"
|
||||
| "bgm_select"
|
||||
| "rendering"
|
||||
| "musetalk"
|
||||
| "uploading"
|
||||
|
||||
/** 分析类阶段(image_analysis / video_analysis / intent_parsing):属于「开始分析」阶段 */
|
||||
const ANALYSIS_STAGES = new Set<ViralVideoStage>([
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
])
|
||||
|
||||
export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return !!stage && ANALYSIS_STAGES.has(stage)
|
||||
}
|
||||
|
||||
export interface StyleTemplate {
|
||||
id: string
|
||||
name: string
|
||||
description?: string
|
||||
preview_url?: string
|
||||
tags?: string[]
|
||||
}
|
||||
|
||||
export interface IntentResult {
|
||||
product: string
|
||||
selling_points: string[]
|
||||
target_audience: string
|
||||
tone: string
|
||||
structure: string
|
||||
duration: number
|
||||
suggested_title?: string
|
||||
suggested_copy?: string
|
||||
}
|
||||
|
||||
export interface ViralVideoJob {
|
||||
id: string
|
||||
status: ViralVideoStatus
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
style_guide?: string
|
||||
user_copy_text?: string
|
||||
final_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_mode?: "global" | "per_video"
|
||||
bgm_preference?: string
|
||||
intent_result?: IntentResult
|
||||
intent_text?: string
|
||||
progress_stage?: ViralVideoStage
|
||||
progress_percent?: number
|
||||
progress_message?: string
|
||||
output_url?: string
|
||||
error_message?: string
|
||||
credits_cost?: number
|
||||
created_at?: string
|
||||
updated_at?: string
|
||||
}
|
||||
|
||||
export interface GenerateViralVideoRequest {
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
bgm_preference?: string
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
language?: string
|
||||
persona_id?: string
|
||||
viral_structure?: string
|
||||
marketing_purpose?: string
|
||||
duration?: number
|
||||
video_model?: string
|
||||
video_ratio?: string
|
||||
}
|
||||
|
||||
export interface HistoryResponse {
|
||||
items: ViralVideoJob[]
|
||||
total: number
|
||||
page: number
|
||||
page_size: number
|
||||
}
|
||||
@@ -2,6 +2,7 @@
|
||||
export const ROUTE_TITLE_MAP: Record<string, string> = {
|
||||
"/app/dashboard": "首页",
|
||||
"/app/generate": "智能剪辑",
|
||||
"/app/viral-video": "爆款视频",
|
||||
"/app/assets": "视频库",
|
||||
"/app/voices": "配音库",
|
||||
"/app/products": "成片库",
|
||||
|
||||
@@ -18,6 +18,7 @@ import {
|
||||
ThunderboltOutlined,
|
||||
UnorderedListOutlined,
|
||||
UserOutlined,
|
||||
FireOutlined,
|
||||
} from "@ant-design/icons"
|
||||
|
||||
/** 导航项类型 */
|
||||
@@ -76,6 +77,12 @@ export const NAV_ITEMS: NavItem[] = [
|
||||
path: "/app/ai-avatar",
|
||||
icon: React.createElement(UserOutlined),
|
||||
},
|
||||
{
|
||||
key: "viral-video",
|
||||
label: "爆款视频",
|
||||
path: "/app/viral-video",
|
||||
icon: React.createElement(FireOutlined),
|
||||
},
|
||||
{
|
||||
key: "history",
|
||||
label: "任务历史",
|
||||
@@ -142,6 +149,12 @@ export const NAV_GROUPS: NavGroup[] = [
|
||||
path: "/app/ai-avatar",
|
||||
icon: React.createElement(UserOutlined),
|
||||
},
|
||||
{
|
||||
key: "viral-video",
|
||||
label: "爆款视频",
|
||||
path: "/app/viral-video",
|
||||
icon: React.createElement(FireOutlined),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
||||
@@ -11,7 +11,12 @@
|
||||
* 防止长标题在窄列里溢出导致与相邻卡片进度条视觉重叠。
|
||||
*/
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import {
|
||||
LoadingOutlined,
|
||||
CheckCircleFilled,
|
||||
CloseCircleOutlined,
|
||||
ClockCircleOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
@@ -61,6 +66,11 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#ef4444" }}
|
||||
/>
|
||||
) : task.status === "queued" ? (
|
||||
<ClockCircleOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#faad14" }}
|
||||
/>
|
||||
) : (
|
||||
<LoadingOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
@@ -85,6 +95,21 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
<div className="xx-batch-gen-card-pct">{Math.round(task.progress)}%</div>
|
||||
</>
|
||||
)}
|
||||
{task.status === "queued" && (
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 8,
|
||||
color: "var(--text-secondary, #faad14)",
|
||||
fontSize: 13,
|
||||
padding: "8px 0",
|
||||
}}
|
||||
>
|
||||
<ClockCircleOutlined />
|
||||
<span>排队等待中,前面任务完成后自动开始渲染</span>
|
||||
</div>
|
||||
)}
|
||||
{(task.status === "completed" || task.status === "awaiting_cover") && video && (
|
||||
// 竖屏自适应容器(#1750):成片固定 1080×1920(9:16),
|
||||
// 视频按真实宽高比 contain 显示,黑底居中,杜绝横屏播放器左右大黑边
|
||||
|
||||
@@ -12,6 +12,7 @@ import Step2MaterialSelect from "../components/Step2MaterialSelect"
|
||||
import Step4TitleSettings from "../components/Step4TitleSettings"
|
||||
import Step6CoverSettings from "../components/Step6CoverSettings"
|
||||
import BatchGenerationGrid from "./BatchGenerationGrid"
|
||||
import Step3VoiceWithMode from "./Step3VoiceWithMode"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
import type { TitleTemplate } from "@/components/title/template-types"
|
||||
@@ -151,6 +152,12 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedCoverTemplate,
|
||||
onSelectedCoverTemplateChange,
|
||||
onConfirmGenerate,
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
voiceModePerVideo,
|
||||
onVoiceModePerVideoChange,
|
||||
voiceLibraryIds,
|
||||
onVoiceLibraryIdsChange,
|
||||
} = props
|
||||
|
||||
switch (currentStep) {
|
||||
@@ -184,32 +191,46 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
)
|
||||
case 3:
|
||||
return (
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1 ? props.selectedVariantIds?.length || 1 : 1
|
||||
}
|
||||
/>
|
||||
<>
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1
|
||||
? props.selectedVariantIds?.length || 1
|
||||
: 1
|
||||
}
|
||||
/>
|
||||
{/* 批量配音选择:共用/独立切换(#2096) */}
|
||||
<Step3VoiceWithMode
|
||||
previewCount={previewCount}
|
||||
selectedVoice={selectedVoice}
|
||||
onSelectedVoiceChange={onSelectedVoiceChange}
|
||||
voiceModePerVideo={voiceModePerVideo}
|
||||
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
|
||||
voiceLibraryIds={voiceLibraryIds}
|
||||
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
|
||||
/>
|
||||
</>
|
||||
)
|
||||
case 4:
|
||||
return (
|
||||
|
||||
@@ -157,7 +157,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
/>
|
||||
</div>
|
||||
) : (
|
||||
/* ── 批量:N 个独立标题输入框 ── */
|
||||
/* ── 批量:N 个独立标题输入框(两列布局 #2096) ── */
|
||||
<div className="xx-batch-titles">
|
||||
<div
|
||||
style={{
|
||||
@@ -170,17 +170,25 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
为每个视频输入独立标题。标题样式(字体/颜色/位置)全局统一。
|
||||
</div>
|
||||
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: 640 }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(2, minmax(0, 1fr))",
|
||||
gap: 16,
|
||||
}}
|
||||
>
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: "100%" }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
@@ -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, getAssetLibraries } from "@/api/assets"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
|
||||
interface Step6CoverSettingsProps {
|
||||
coverSettings: CoverConfig
|
||||
@@ -176,15 +176,8 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
|
||||
thumbnail_url: previewUrl,
|
||||
mode: "upload",
|
||||
})
|
||||
// 查找图片素材库(复用批量封面的逻辑)
|
||||
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 })
|
||||
// 后端自动在默认项目下确保图片素材库存在(P0 404 修复)
|
||||
const result = await uploadAssetDirect({ file, kind: "image" })
|
||||
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"
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed" | "queued"
|
||||
progress: number
|
||||
error: string | null
|
||||
/** 完成后的成片视频 */
|
||||
@@ -374,5 +374,42 @@ export function useGenerationPolling({
|
||||
}
|
||||
}, [])
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
/**
|
||||
* 批量队列模式:逐任务追加到轮询队列(支持串行提交、429 排队重试场景)。
|
||||
* 与 startPollingBatch 不同的是:
|
||||
* - 不会 reset batchContextRef;多次调用会累积
|
||||
* - 不触发整体 onComplete / onFailed(完成判定交给外层 useEffect 按状态聚合)
|
||||
* - 仍通过 onBatchTaskUpdate 回传单任务状态
|
||||
*/
|
||||
const pollBatchTaskQueued = useCallback(
|
||||
(taskId: string, variantIndex: number) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.set(taskId, variantIndex)
|
||||
onBatchTaskUpdate?.(taskId, {
|
||||
taskId,
|
||||
variantIndex,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
})
|
||||
pollSingleTask(taskId, Date.now(), {
|
||||
onTaskProgress: (pct) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: pct })
|
||||
},
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
const finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
onBatchTaskUpdate?.(taskId, { status: finalStatus, progress: 100, videos })
|
||||
},
|
||||
onTaskFailed: (msg) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "failed", error: msg })
|
||||
},
|
||||
}).catch(() => {
|
||||
/* onTaskFailed 已处理 */
|
||||
})
|
||||
},
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, pollBatchTaskQueued, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
import { useCallback, useState } from "react"
|
||||
import { message } from "antd"
|
||||
import { generateCover } from "@/api/generation"
|
||||
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
/** onCoversChange 支持直接传值或函数式 updater(函数式用于串行回写避免闭包覆盖) */
|
||||
@@ -182,15 +182,9 @@ 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,
|
||||
library_id: imageLib.id,
|
||||
kind: "image",
|
||||
})
|
||||
const url = result?.url || ""
|
||||
if (url) {
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
* 视频生成 Hook
|
||||
* 封装视频生成的核心逻辑、状态管理、轮询等
|
||||
*/
|
||||
import { useState, useCallback, useEffect } from "react"
|
||||
import { useState, useCallback, useEffect, useRef } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor"
|
||||
import { createGenerationTask } from "@/api/tasks/tasks"
|
||||
import type { CreateGenerationTaskRequest } from "@/api/tasks/types"
|
||||
import type { UseGenerateVideoProps } from "./generate-video/types"
|
||||
import { getGenerationPhase } from "./generate-video/phase"
|
||||
import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling"
|
||||
@@ -15,6 +17,26 @@ import { extractBackendError, translateError } from "./generate-video/errorUtils
|
||||
|
||||
export type GenerationCompleteStatus = "completed" | "awaiting_cover" | null
|
||||
|
||||
/** 判断是否是用户队列已满 429(需要排队重试而非直接报错) */
|
||||
function isUserQueueFullError(err: unknown): { waitMs: number } | null {
|
||||
if (!axios.isAxiosError(err)) return null
|
||||
if (err.response?.status !== 429 && err.response?.status !== 503) return null
|
||||
const detail = (err.response?.data as { detail?: unknown })?.detail
|
||||
const code =
|
||||
typeof detail === "object" && detail !== null ? (detail as { code?: string }).code : undefined
|
||||
if (code === "USER_QUEUE_FULL" || code === "SYSTEM_QUEUE_FULL") {
|
||||
const waitSec =
|
||||
typeof detail === "object" && detail !== null
|
||||
? Number((detail as { estimated_wait_seconds?: number }).estimated_wait_seconds) || 0
|
||||
: 0
|
||||
return { waitMs: Math.max(15_000, waitSec * 1000 || 30_000) }
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/** sleep */
|
||||
const sleep = (ms: number) => new Promise<void>((r) => setTimeout(r, ms))
|
||||
|
||||
export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { selectedTemplate, onGenerationSuccess } = props
|
||||
|
||||
@@ -31,6 +53,22 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
|
||||
|
||||
/** 排队中重试的定时器,unmount / 新提交时清理 */
|
||||
const queueTimersRef = useRef<number[]>([])
|
||||
const cancelledRef = useRef(false)
|
||||
|
||||
const clearQueueTimers = useCallback(() => {
|
||||
queueTimersRef.current.forEach((id) => clearTimeout(id))
|
||||
queueTimersRef.current = []
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
cancelledRef.current = true
|
||||
clearQueueTimers()
|
||||
}
|
||||
}, [clearQueueTimers])
|
||||
|
||||
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
|
||||
setBatchTasks((prev) => {
|
||||
const list = prev || []
|
||||
@@ -58,7 +96,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
(videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => {
|
||||
setGenerating(false)
|
||||
setGenerated(true)
|
||||
const finalStatus: GenerationCompleteStatus = taskStatus ?? "completed"
|
||||
setCompletionStatus(finalStatus)
|
||||
@@ -81,21 +118,30 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
[onGenerationSuccess],
|
||||
)
|
||||
const handleFailed = useCallback((errorMsg: string) => {
|
||||
setGenerating(false)
|
||||
setGenerateError(errorMsg)
|
||||
}, [])
|
||||
|
||||
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频。
|
||||
当全部任务都已结束(completed/awaiting_cover/failed)且无排队/渲染中任务时,关闭 generating。 */
|
||||
useEffect(() => {
|
||||
if (batchTasks.length === 0) return
|
||||
const byVariant = new Map<number, GeneratedVideo>()
|
||||
let hasQueued = false
|
||||
let hasRunning = false
|
||||
let hasSuccess = false
|
||||
let allDone = true
|
||||
batchTasks.forEach((t) => {
|
||||
if (
|
||||
t.status === "completed" ||
|
||||
(t.status === "awaiting_cover" && t.videos && t.videos.length > 0)
|
||||
) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
if (t.status === "queued") hasQueued = true
|
||||
else if (t.status === "running") hasRunning = true
|
||||
if (t.status === "completed" || t.status === "awaiting_cover") {
|
||||
hasSuccess = true
|
||||
if (t.videos && t.videos.length > 0) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
}
|
||||
}
|
||||
if (t.status !== "completed" && t.status !== "awaiting_cover" && t.status !== "failed") {
|
||||
allDone = false
|
||||
}
|
||||
})
|
||||
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
|
||||
@@ -105,17 +151,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
return ordered
|
||||
})
|
||||
if (allDone && !hasQueued && !hasRunning) {
|
||||
setGenerating(false)
|
||||
if (hasSuccess) {
|
||||
setGenerated(true)
|
||||
setCompletionStatus("awaiting_cover")
|
||||
}
|
||||
}
|
||||
}, [batchTasks])
|
||||
|
||||
const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({
|
||||
const { startPolling, pollBatchTaskQueued, retryTask, clearTimer } = useGenerationPolling({
|
||||
onProgress: handleProgress,
|
||||
onComplete: handleComplete,
|
||||
onFailed: handleFailed,
|
||||
onBatchTaskUpdate: handleBatchTaskUpdate,
|
||||
})
|
||||
|
||||
/** 根据 props 构造基础 payload(批量/单任务共用的字段) */
|
||||
const buildBasePayload = useCallback((): Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> => {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
const bgmConfig = {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
}
|
||||
|
||||
const titleConfig = props.titleSettings?.title
|
||||
? {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color: props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
}
|
||||
: undefined
|
||||
|
||||
const payload: Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> = {
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: bgmConfig as CreateGenerationTaskRequest["bgm_config"],
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(titleConfig ? ({ title_config: titleConfig } as Record<string, unknown>) : {}),
|
||||
}
|
||||
|
||||
return payload
|
||||
}, [props, selectedTemplate])
|
||||
|
||||
/* ── 生成视频 ──
|
||||
返回 true 表示任务创建成功并已开始轮询;false 表示校验未通过或创建失败 */
|
||||
返回 true 表示任务创建成功并已开始轮询(含排队中);false 表示校验未通过或创建失败 */
|
||||
const generate = useCallback(async (): Promise<boolean> => {
|
||||
const errorMsg = validateGenerateInputs(props)
|
||||
if (errorMsg) {
|
||||
@@ -123,6 +337,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
return false
|
||||
}
|
||||
|
||||
cancelledRef.current = false
|
||||
clearQueueTimers()
|
||||
setGenerating(true)
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
@@ -133,25 +349,16 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setCurrentTaskId("")
|
||||
clearTimer()
|
||||
|
||||
const basePayload = buildBasePayload()
|
||||
const assetIds = basePayload.asset_ids
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
|
||||
try {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
|
||||
// from-assets 已由 useStep2Materials 在用户选素材时(debounce 800ms)调用,
|
||||
// 后端已改为异步秒级返回,这里做一次轻量兜底:
|
||||
// 单次查 clips,已有则直接放行;没有则再调一次 from-assets。
|
||||
// from-assets 兜底:片段不存在则补一次
|
||||
if (assetIds.length > 0 && selectedTemplate) {
|
||||
try {
|
||||
const clipList = await getEditPlanClips(selectedTemplate, { limit: 500 })
|
||||
if (clipList.items.length === 0) {
|
||||
// 片段不存在(极端情况:useStep2Materials 的 debounce 还没触发)
|
||||
// 手动补一次 from-assets(后端秒级返回)
|
||||
await createClipsFromAssets(selectedTemplate, assetIds, "main")
|
||||
}
|
||||
} catch {
|
||||
@@ -159,221 +366,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
}
|
||||
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
const hide = message.loading(
|
||||
isBatch ? `正在生成 ${props.previewCount} 个视频...` : "正在生成预览视频...",
|
||||
0,
|
||||
)
|
||||
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
|
||||
const indexes =
|
||||
isBatch && props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = isBatch ? indexes.length : 1
|
||||
|
||||
// 标题文字数组:批量时按勾选顺序
|
||||
const titlesArr =
|
||||
isBatch && (props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: []
|
||||
// 配音数组:独立配音模式按勾选顺序;否则不传(回退共用 voice_library_id)
|
||||
const voiceArr =
|
||||
isBatch && props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map((i) => props.variantVoiceLibraryIds![i] || voiceLibraryId)
|
||||
: []
|
||||
// 封面数组:批量时按勾选顺序(未设置封面的变体传空串,后端回退智能封面)
|
||||
const coversArr =
|
||||
isBatch && props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: []
|
||||
// #1744 变体 plan 数组:预览阶段后端独立选片产出的 plan id,按勾选顺序回传,
|
||||
// 后端直接关联这些 plan 渲染(不再重新选片)→ 预览所见即成片。
|
||||
// 全部为空(降级本地模拟/后端端点未上线)时不传,后端走自身独立选片。
|
||||
const variantPlansArr =
|
||||
isBatch && props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: []
|
||||
const hasVariantPlans = variantPlansArr.some((id) => !!id)
|
||||
|
||||
try {
|
||||
const taskResp = await createGenerationTask({
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
},
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(isBatch ? { count: batchCount } : {}),
|
||||
...(titlesArr.length ? { titles: titlesArr } : {}),
|
||||
...(voiceArr.length ? { voice_library_ids: voiceArr } : {}),
|
||||
...(coversArr.length ? { cover_urls: coversArr } : {}),
|
||||
...(hasVariantPlans ? { variant_plan_ids: variantPlansArr } : {}),
|
||||
...(props.titleSettings?.title
|
||||
? {
|
||||
title_config: {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color:
|
||||
props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
})
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
if (taskIds.length > 1) {
|
||||
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
|
||||
setCurrentTaskId("")
|
||||
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
|
||||
} else {
|
||||
if (!isBatch) {
|
||||
/* ── 单视频:原逻辑(一次提交 count=1) ── */
|
||||
const hide = message.loading("正在生成预览视频...", 0)
|
||||
try {
|
||||
const taskResp = await createGenerationTask({ ...basePayload, count: 1 })
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
setCurrentTaskId(taskIds[0])
|
||||
startPolling(taskIds[0])
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
}
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
return true
|
||||
}
|
||||
|
||||
/* ── 批量:支持任意数量视频,按队列容量串行提交,429 自动排队重试 ── */
|
||||
const indexes = props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = indexes.length
|
||||
|
||||
const titlesAll =
|
||||
(props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: indexes.map(() => props.titleSettings?.title || "")
|
||||
const voiceArrAll =
|
||||
props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map(
|
||||
(i) => props.variantVoiceLibraryIds![i] || basePayload.voice_library_id || "",
|
||||
)
|
||||
: []
|
||||
const coversAll = props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: indexes.map(() => "")
|
||||
const plansAll = props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: indexes.map(() => "")
|
||||
|
||||
const hasAnyVoice = voiceArrAll.some((v) => !!v)
|
||||
const hasAnyCover = coversAll.some((u) => !!u)
|
||||
const hasAnyPlan = plansAll.some((id) => !!id)
|
||||
|
||||
// 先用占位 ID 把所有变体卡片置为 queued,UI 可见
|
||||
const placeholderIds = indexes.map((_, i) => `__queued_${Date.now()}_${i}`)
|
||||
const initialTasks: BatchTaskState[] = indexes.map((variantIndex, i) => ({
|
||||
taskId: placeholderIds[i],
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
}))
|
||||
setBatchTasks(initialTasks)
|
||||
|
||||
message.loading({
|
||||
content: `已提交 ${batchCount} 个视频任务,系统按队列容量依次渲染…`,
|
||||
key: "batch-gen",
|
||||
duration: 3,
|
||||
})
|
||||
|
||||
/** 将占位 taskId 更新为真实 taskId(卡片引用同一对象) */
|
||||
const replacePlaceholder = (placeholderId: string, realTaskId: string) => {
|
||||
setBatchTasks((prev) => {
|
||||
const idx = prev.findIndex((t) => t.taskId === placeholderId)
|
||||
if (idx === -1) return prev
|
||||
const next = [...prev]
|
||||
next[idx] = { ...next[idx], taskId: realTaskId }
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
/** 提交某一索引的单任务(count=1),成功后返回真实 taskId;429/503 则返回 waitMs */
|
||||
const submitOne = async (
|
||||
i: number,
|
||||
): Promise<{ queued: true; waitMs: number } | { queued: false; taskId: string }> => {
|
||||
const body: CreateGenerationTaskRequest = {
|
||||
...basePayload,
|
||||
count: 1,
|
||||
titles: [titlesAll[i] || ""],
|
||||
...(hasAnyVoice
|
||||
? { voice_library_ids: [voiceArrAll[i] || basePayload.voice_library_id || ""] }
|
||||
: {}),
|
||||
...(hasAnyCover ? { cover_urls: [coversAll[i] || ""] } : {}),
|
||||
...(hasAnyPlan && plansAll[i] ? { variant_plan_ids: [plansAll[i]] } : {}),
|
||||
}
|
||||
try {
|
||||
const resp = await createGenerationTask(body)
|
||||
const item = resp.items?.[0]
|
||||
const tid = item?.id
|
||||
if (!tid) throw new Error("创建任务成功但未返回任务 ID")
|
||||
return { queued: false, taskId: tid }
|
||||
} catch (err) {
|
||||
const q = isUserQueueFullError(err)
|
||||
if (q) return { queued: true, waitMs: q.waitMs }
|
||||
throw err
|
||||
}
|
||||
}
|
||||
|
||||
// 串行提交:每次提交一个;429/503 则等待后重试;其它错误立即标记该任务失败
|
||||
let fatalErr: unknown = null
|
||||
for (let i = 0; i < batchCount; i++) {
|
||||
if (cancelledRef.current) return false
|
||||
const variantIndex = indexes[i]
|
||||
const placeholderId = placeholderIds[i]
|
||||
let attempt = 0
|
||||
let submitted = false
|
||||
while (!submitted) {
|
||||
if (cancelledRef.current) return false
|
||||
attempt++
|
||||
try {
|
||||
const result = await submitOne(i)
|
||||
if (!result.queued) {
|
||||
replacePlaceholder(placeholderId, result.taskId)
|
||||
// 先更新到 running,再启动单任务增量轮询(不触发整体 onComplete)
|
||||
pollBatchTaskQueued(result.taskId, variantIndex)
|
||||
submitted = true
|
||||
} else {
|
||||
// 排队:保持 queued 状态,等待后重试
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
})
|
||||
if (attempt === 1) {
|
||||
message.info({
|
||||
content: `队列繁忙,${Math.round(result.waitMs / 1000)} 秒后自动继续提交后续视频…`,
|
||||
key: "batch-gen",
|
||||
duration: 4,
|
||||
})
|
||||
}
|
||||
await sleep(Math.min(result.waitMs, 60_000))
|
||||
}
|
||||
} catch (err) {
|
||||
// 非限流错误:该任务标记失败,继续后续任务(不阻断整个批量)
|
||||
console.error("[batch generate] 任务提交失败:", err)
|
||||
const msg = translateError(extractBackendError(err))
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "failed",
|
||||
error: msg,
|
||||
progress: 0,
|
||||
})
|
||||
submitted = true
|
||||
if (!fatalErr) fatalErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (fatalErr) {
|
||||
// 有任务失败但其余已成功,整体不 throw;由 UI 展示单个失败卡片
|
||||
}
|
||||
return true
|
||||
} catch (err: unknown) {
|
||||
console.error("[handleGenerate] 生成失败:", err)
|
||||
setGenerating(false)
|
||||
const backendMsg = extractBackendError(err)
|
||||
console.error("[handleGenerate] 错误信息:", backendMsg, "完整错误:", err)
|
||||
const finalMsg = translateError(backendMsg)
|
||||
setGenerateError(finalMsg)
|
||||
setGenerating(false)
|
||||
message.error(finalMsg)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
|
||||
}, [
|
||||
props,
|
||||
clearTimer,
|
||||
startPolling,
|
||||
selectedTemplate,
|
||||
buildBasePayload,
|
||||
handleBatchTaskUpdate,
|
||||
clearQueueTimers,
|
||||
pollBatchTaskQueued,
|
||||
])
|
||||
|
||||
const retry = useCallback(() => {
|
||||
setGenerateError(null)
|
||||
@@ -383,9 +554,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 第5步:单独重试某个失败任务 */
|
||||
const retryBatchTask = useCallback(
|
||||
(taskId: string) => {
|
||||
handleBatchTaskUpdate(taskId, { status: "running", progress: 0, error: null, videos: [] })
|
||||
retryTask(taskId)
|
||||
},
|
||||
[retryTask],
|
||||
[retryTask, handleBatchTaskUpdate],
|
||||
)
|
||||
|
||||
const dismissError = useCallback(() => {
|
||||
|
||||
@@ -0,0 +1,817 @@
|
||||
/* ============================================================
|
||||
爆款视频创作页 - 浅色紫调(对齐 AI 数字人页视觉规范)
|
||||
布局(参考 ui-ref-step-layout.png 三列等宽 STEP 向导):
|
||||
.vv-tabs 顶栏多任务 Tab(生成1 × / + 新建)
|
||||
.vv-grid 三列等宽 grid(1fr 1fr 1fr,gap 16)
|
||||
├── .vv-col 左:STEP 1 上传素材(图片+参考视频+配音)
|
||||
├── .vv-col 中:STEP 2 生成视频文案(融合Tab+参数+文案I/O+AI摘要)
|
||||
└── .vv-col 右:STEP 3 生成视频(预览+参数+进度+扣点+按钮)
|
||||
可折叠模块:.vv-section > .vv-section-head[aria-expanded] + .vv-section-body
|
||||
============================================================ */
|
||||
|
||||
.vv-page {
|
||||
padding: 16px;
|
||||
background: #f5f6fa;
|
||||
min-height: calc(100vh - 56px);
|
||||
}
|
||||
|
||||
/* ── 顶部任务 Tab 栏 ─────────────────────────────────── */
|
||||
.vv-tabs {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
margin-bottom: 14px;
|
||||
padding: 6px 8px;
|
||||
background: #fff;
|
||||
border-radius: 10px;
|
||||
border: 1px solid #e5e7eb;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
.vv-tab {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
padding: 6px 12px;
|
||||
border-radius: 6px;
|
||||
font-size: 13px;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
border: 1px solid transparent;
|
||||
background: transparent;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-tab:hover {
|
||||
background: #f3f4f6;
|
||||
color: #374151;
|
||||
}
|
||||
.vv-tab.active {
|
||||
background: #f3f0ff;
|
||||
color: #7c3aed;
|
||||
border-color: #d8cafc;
|
||||
font-weight: 500;
|
||||
}
|
||||
.vv-tab .vv-tab-close {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-tab .vv-tab-close:hover {
|
||||
background: rgba(0, 0, 0, 0.08);
|
||||
color: #374151;
|
||||
}
|
||||
.vv-tab-new {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
padding: 6px 10px;
|
||||
border-radius: 6px;
|
||||
font-size: 13px;
|
||||
color: #7c3aed;
|
||||
cursor: pointer;
|
||||
border: 1px dashed #d8cafc;
|
||||
background: transparent;
|
||||
}
|
||||
.vv-tab-new:hover {
|
||||
background: #f3f0ff;
|
||||
}
|
||||
.vv-tabs-right {
|
||||
margin-left: auto;
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
/* ── 三列等宽网格 ───────────────────────────────────── */
|
||||
.vv-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(0, 1fr));
|
||||
gap: 16px;
|
||||
align-items: start;
|
||||
}
|
||||
@media (max-width: 1280px) {
|
||||
.vv-grid {
|
||||
grid-template-columns: repeat(2, minmax(0, 1fr));
|
||||
}
|
||||
}
|
||||
@media (max-width: 900px) {
|
||||
.vv-grid {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
|
||||
.vv-col {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 14px;
|
||||
}
|
||||
|
||||
/* ── 可折叠模块(对齐 AI 数字人卡块) ────────────────── */
|
||||
.vv-section {
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 12px;
|
||||
overflow: hidden;
|
||||
}
|
||||
.vv-section-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 16px;
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
border-bottom: 1px solid #f3f4f6;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-head {
|
||||
border-bottom: none;
|
||||
}
|
||||
.vv-section-title {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
color: #111827;
|
||||
}
|
||||
.vv-section-title .vv-step-badge {
|
||||
width: 22px;
|
||||
height: 22px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
color: #fff;
|
||||
font-size: 12px;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-section-arrow {
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-arrow {
|
||||
transform: rotate(-90deg);
|
||||
}
|
||||
.vv-section-body {
|
||||
padding: 14px 16px 16px;
|
||||
}
|
||||
.vv-section.collapsed .vv-section-body {
|
||||
display: none;
|
||||
}
|
||||
|
||||
/* ── 通用表单元素(对齐 AI 数字人样式) ─────────────── */
|
||||
.vv-label {
|
||||
display: block;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.vv-input,
|
||||
.vv-select,
|
||||
.vv-textarea {
|
||||
width: 100%;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 8px;
|
||||
padding: 9px 12px;
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
background: #fff;
|
||||
outline: none;
|
||||
transition:
|
||||
border-color 0.15s,
|
||||
box-shadow 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-textarea {
|
||||
line-height: 1.6;
|
||||
resize: vertical;
|
||||
min-height: 90px;
|
||||
}
|
||||
.vv-input:focus,
|
||||
.vv-select:focus,
|
||||
.vv-textarea:focus {
|
||||
border-color: #7c3aed;
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.1);
|
||||
}
|
||||
.vv-input::placeholder,
|
||||
.vv-textarea::placeholder {
|
||||
color: #d1d5db;
|
||||
}
|
||||
.vv-form-row {
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.vv-form-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, minmax(0, 1fr));
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
/* ── Tab 分段(对齐"系统预设/我的音色"样式) ────────── */
|
||||
.vv-seg-tabs {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
margin-bottom: 12px;
|
||||
border-radius: 8px;
|
||||
padding: 3px;
|
||||
background: #f5f6fa;
|
||||
}
|
||||
.vv-seg-tab {
|
||||
flex: 1;
|
||||
padding: 7px 10px;
|
||||
font-size: 13px;
|
||||
text-align: center;
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
color: #6b7280;
|
||||
background: transparent;
|
||||
border: 1px solid transparent;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-seg-tab:hover {
|
||||
color: #374151;
|
||||
}
|
||||
.vv-seg-tab.active {
|
||||
background: #fff;
|
||||
color: #7c3aed;
|
||||
border-color: #7c3aed;
|
||||
font-weight: 500;
|
||||
box-shadow: 0 1px 2px rgba(124, 58, 237, 0.06);
|
||||
}
|
||||
|
||||
/* 融合强度大按钮(选中紫色描边+浅紫底) */
|
||||
.vv-fusion-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(0, 1fr));
|
||||
gap: 8px;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.vv-fusion-btn {
|
||||
padding: 10px 8px;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
border-radius: 8px;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
color: #6b7280;
|
||||
text-align: center;
|
||||
line-height: 1.4;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-fusion-btn strong {
|
||||
display: block;
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
margin-bottom: 2px;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-fusion-btn:hover {
|
||||
border-color: #d8cafc;
|
||||
}
|
||||
.vv-fusion-btn.active {
|
||||
background: #f3f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-fusion-btn.active strong {
|
||||
color: #7c3aed;
|
||||
}
|
||||
|
||||
/* 风格强度小分段按钮(三档,参考配音风格按钮) */
|
||||
.vv-pill-row {
|
||||
display: flex;
|
||||
gap: 6px;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
.vv-pill {
|
||||
padding: 6px 12px;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
border-radius: 6px;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
font-family: inherit;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-pill:hover {
|
||||
border-color: #d8cafc;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-pill.active {
|
||||
background: #f3f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
/* ── 上传区(浅灰虚线框) ───────────────────────────── */
|
||||
.vv-upload {
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 10px;
|
||||
padding: 20px;
|
||||
text-align: center;
|
||||
cursor: pointer;
|
||||
transition:
|
||||
border-color 0.2s,
|
||||
background 0.2s;
|
||||
background: #fafbfc;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-upload:hover,
|
||||
.vv-upload.dragover {
|
||||
border-color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-upload-icon {
|
||||
font-size: 28px;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.vv-upload small {
|
||||
display: block;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
/* 图片网格 */
|
||||
.vv-img-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(90px, 1fr));
|
||||
gap: 8px;
|
||||
margin-top: 12px;
|
||||
}
|
||||
.vv-img-item {
|
||||
position: relative;
|
||||
aspect-ratio: 1;
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
border: 1px solid #e5e7eb;
|
||||
cursor: grab;
|
||||
background: #f5f6fa;
|
||||
}
|
||||
.vv-img-item.dragging {
|
||||
opacity: 0.4;
|
||||
}
|
||||
.vv-img-item img {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: cover;
|
||||
}
|
||||
.vv-img-badge {
|
||||
position: absolute;
|
||||
top: 4px;
|
||||
left: 4px;
|
||||
background: rgba(124, 58, 237, 0.9);
|
||||
color: #fff;
|
||||
font-size: 11px;
|
||||
font-weight: 600;
|
||||
padding: 1px 6px;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-img-del {
|
||||
position: absolute;
|
||||
top: 4px;
|
||||
right: 4px;
|
||||
width: 20px;
|
||||
height: 20px;
|
||||
border-radius: 50%;
|
||||
background: rgba(239, 68, 68, 0.9);
|
||||
color: #fff;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
font-size: 12px;
|
||||
line-height: 1;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
opacity: 0;
|
||||
transition: opacity 0.15s;
|
||||
}
|
||||
.vv-img-item:hover .vv-img-del {
|
||||
opacity: 1;
|
||||
}
|
||||
.vv-img-add {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 8px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #9ca3af;
|
||||
cursor: pointer;
|
||||
background: #fafbfc;
|
||||
font-size: 11px;
|
||||
gap: 2px;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-img-add:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
}
|
||||
.vv-progress-mini {
|
||||
position: absolute;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.35);
|
||||
color: #fff;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 11px;
|
||||
}
|
||||
|
||||
/* 参考视频预览 */
|
||||
.vv-video-preview {
|
||||
width: 100%;
|
||||
aspect-ratio: 16/9;
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
background: #000;
|
||||
margin-top: 10px;
|
||||
border: 1px solid #e5e7eb;
|
||||
}
|
||||
.vv-video-preview video {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: cover;
|
||||
}
|
||||
.vv-video-ph {
|
||||
width: 100%;
|
||||
aspect-ratio: 16/9;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
border-radius: 8px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #9ca3af;
|
||||
cursor: pointer;
|
||||
background: #fafbfc;
|
||||
font-size: 12px;
|
||||
gap: 4px;
|
||||
margin-top: 10px;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-video-ph:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f9f7ff;
|
||||
}
|
||||
|
||||
/* ── 音色列表(参考 AI 数字人「龙小淳」卡片) ──────── */
|
||||
.vv-voice-tabs {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-voice-list {
|
||||
max-height: 280px;
|
||||
overflow-y: auto;
|
||||
border: 1px solid #f3f4f6;
|
||||
border-radius: 8px;
|
||||
}
|
||||
.vv-voice-list::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
.vv-voice-list::-webkit-scrollbar-thumb {
|
||||
background: #e5e7eb;
|
||||
border-radius: 3px;
|
||||
}
|
||||
.vv-voice-item {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 10px 12px;
|
||||
border-bottom: 1px solid #f3f4f6;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
.vv-voice-item:last-child {
|
||||
border-bottom: none;
|
||||
}
|
||||
.vv-voice-item:hover {
|
||||
background: #f9fafb;
|
||||
}
|
||||
.vv-voice-item.selected {
|
||||
background: #f3f0ff;
|
||||
}
|
||||
.vv-voice-radio {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
border: 2px solid #d1d5db;
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
.vv-voice-item.selected .vv-voice-radio {
|
||||
border-color: #7c3aed;
|
||||
}
|
||||
.vv-voice-item.selected .vv-voice-radio::after {
|
||||
content: "";
|
||||
width: 8px;
|
||||
height: 8px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
}
|
||||
.vv-voice-info {
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
}
|
||||
.vv-voice-name {
|
||||
font-size: 13px;
|
||||
color: #111827;
|
||||
font-weight: 500;
|
||||
}
|
||||
.vv-voice-desc {
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 2px;
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
}
|
||||
.vv-voice-play {
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border-radius: 50%;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
color: #6b7280;
|
||||
cursor: pointer;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
font-size: 12px;
|
||||
flex-shrink: 0;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-voice-play:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-voice-play.playing {
|
||||
background: #7c3aed;
|
||||
border-color: #7c3aed;
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
/* ── AI 摘要确认卡(黄色高亮) ─────────────────────── */
|
||||
.vv-intent {
|
||||
background: #fffbeb;
|
||||
border: 1px solid #fcd34d;
|
||||
border-radius: 10px;
|
||||
padding: 14px;
|
||||
margin-top: 12px;
|
||||
}
|
||||
.vv-intent-title {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: #b45309;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-intent-row {
|
||||
margin-bottom: 8px;
|
||||
font-size: 13px;
|
||||
line-height: 1.6;
|
||||
}
|
||||
.vv-intent-row .vv-k {
|
||||
font-size: 12px;
|
||||
color: #92400e;
|
||||
margin-bottom: 3px;
|
||||
}
|
||||
.vv-intent-row .vv-v {
|
||||
color: #111827;
|
||||
}
|
||||
.vv-chip {
|
||||
display: inline-block;
|
||||
padding: 2px 8px;
|
||||
background: #fef3c7;
|
||||
border-radius: 4px;
|
||||
font-size: 11px;
|
||||
color: #92400e;
|
||||
margin-right: 4px;
|
||||
margin-bottom: 3px;
|
||||
}
|
||||
|
||||
/* ── 9:16 预览区 ────────────────────────────────────── */
|
||||
.vv-preview {
|
||||
width: 100%;
|
||||
aspect-ratio: 9/16;
|
||||
max-height: 480px;
|
||||
border-radius: 12px;
|
||||
background: #fff;
|
||||
border: 1.5px dashed #d1d5db;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
.vv-preview video {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: contain;
|
||||
background: #000;
|
||||
}
|
||||
.vv-preview-placeholder {
|
||||
text-align: center;
|
||||
color: #9ca3af;
|
||||
padding: 20px;
|
||||
}
|
||||
.vv-preview-placeholder .ph-icon {
|
||||
font-size: 40px;
|
||||
opacity: 0.4;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
.vv-preview-placeholder .ph-txt {
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
.vv-progress-bar {
|
||||
width: 100%;
|
||||
height: 6px;
|
||||
background: #f3f4f6;
|
||||
border-radius: 3px;
|
||||
overflow: hidden;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-progress-fill {
|
||||
height: 100%;
|
||||
background: linear-gradient(90deg, #7c3aed, #a855f7);
|
||||
border-radius: 3px;
|
||||
transition: width 0.4s ease;
|
||||
}
|
||||
.vv-progress-meta {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 6px;
|
||||
}
|
||||
.vv-progress-pct {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
/* ── 扣点 & 按钮 ────────────────────────────────────── */
|
||||
.vv-credits {
|
||||
background: #f9fafb;
|
||||
border: 1px solid #f3f4f6;
|
||||
border-radius: 8px;
|
||||
padding: 10px 12px;
|
||||
font-size: 12px;
|
||||
color: #6b7280;
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
}
|
||||
.vv-credits strong {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.vv-btn {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 6px;
|
||||
padding: 10px 16px;
|
||||
border-radius: 8px;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
cursor: pointer;
|
||||
border: 1px solid transparent;
|
||||
transition: all 0.15s;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-btn-primary {
|
||||
width: 100%;
|
||||
background: linear-gradient(135deg, #7c3aed, #a855f7);
|
||||
color: #fff;
|
||||
padding: 12px;
|
||||
font-size: 14px;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-btn-primary:hover:not(:disabled) {
|
||||
box-shadow: 0 4px 14px rgba(124, 58, 237, 0.3);
|
||||
transform: translateY(-1px);
|
||||
}
|
||||
.vv-btn-primary:disabled {
|
||||
opacity: 0.5;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.vv-btn-ghost {
|
||||
background: #fff;
|
||||
color: #6b7280;
|
||||
border-color: #e5e7eb;
|
||||
}
|
||||
.vv-btn-ghost:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-btn-warn {
|
||||
background: #fef3c7;
|
||||
color: #b45309;
|
||||
border-color: #fcd34d;
|
||||
}
|
||||
.vv-btn-sm {
|
||||
padding: 6px 12px;
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
.vv-error {
|
||||
background: #fef2f2;
|
||||
border: 1px solid #fecaca;
|
||||
border-radius: 8px;
|
||||
padding: 10px 12px;
|
||||
color: #dc2626;
|
||||
font-size: 12px;
|
||||
margin-top: 10px;
|
||||
}
|
||||
.vv-spinner {
|
||||
width: 14px;
|
||||
height: 14px;
|
||||
border: 2px solid rgba(255, 255, 255, 0.3);
|
||||
border-top-color: #fff;
|
||||
border-radius: 50%;
|
||||
animation: vvspin 0.8s linear infinite;
|
||||
display: inline-block;
|
||||
}
|
||||
@keyframes vvspin {
|
||||
to {
|
||||
transform: rotate(360deg);
|
||||
}
|
||||
}
|
||||
|
||||
.vv-muted {
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
}
|
||||
.vv-meta {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
font-size: 11px;
|
||||
color: #9ca3af;
|
||||
margin-top: 4px;
|
||||
}
|
||||
.vv-file-name {
|
||||
font-size: 12px;
|
||||
color: #374151;
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-link-btn {
|
||||
background: none;
|
||||
border: none;
|
||||
color: #7c3aed;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
font-family: inherit;
|
||||
}
|
||||
.vv-link-btn:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
.vv-flex {
|
||||
display: flex;
|
||||
}
|
||||
.vv-between {
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
}
|
||||
.vv-gap-8 {
|
||||
gap: 8px;
|
||||
}
|
||||
.vv-mt-8 {
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-mt-12 {
|
||||
margin-top: 12px;
|
||||
}
|
||||
|
||||
/* disabled 状态 */
|
||||
.vv-pill:disabled,
|
||||
.vv-fusion-btn:disabled {
|
||||
opacity: 0.6;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.vv-btn:disabled {
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.vv-section-head {
|
||||
gap: 8px;
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,74 @@
|
||||
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 }
|
||||
}
|
||||
+13
-24
@@ -1,25 +1,26 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import {
|
||||
uploadAssetDirect,
|
||||
getAssetLibraries,
|
||||
getIngestJob,
|
||||
type AssetLibraryItem,
|
||||
} from "@/api/assets"
|
||||
import { uploadAssetDirect, 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 }: UseVoiceUploadOptions) {
|
||||
export function useVoiceUpload({
|
||||
voiceLibrary,
|
||||
createLibMutation: _createLibMutation,
|
||||
}: UseVoiceUploadOptions) {
|
||||
const queryClient = useQueryClient()
|
||||
const [uploadProgress, setUploadProgress] = useState<number | null>(null)
|
||||
|
||||
@@ -33,24 +34,12 @@ export function useVoiceUpload({ voiceLibrary, createLibMutation }: UseVoiceUplo
|
||||
}) => {
|
||||
setUploadProgress(0)
|
||||
try {
|
||||
// 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)
|
||||
// 1. 上传文件:后端自动在默认项目下确保配音库存在(P0 404 修复)
|
||||
// 兼容 voiceLibrary 参数:若调用方已传入正确的库 ID 则直接复用,否则内部自动解析
|
||||
const complete = await uploadAssetDirect({
|
||||
file: data.file,
|
||||
library_id: lib.id,
|
||||
library_id: voiceLibrary?.id,
|
||||
kind: "voice",
|
||||
onProgress: (p) => setUploadProgress(p),
|
||||
})
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { uploadAssetDirect, getAssetLibraries, getIngestJob } from "@/api/assets"
|
||||
import { uploadAssetDirect, getIngestJob } from "@/api/assets"
|
||||
|
||||
/**
|
||||
* 配音上传 Hook
|
||||
@@ -23,18 +23,10 @@ export function useVoiceUpload({ showToast }: UseVoiceUploadProps) {
|
||||
mutationFn: async (data: { file: File; name: string; description: string }) => {
|
||||
setUploadProgress(0)
|
||||
try {
|
||||
/* 获取或创建默认配音库 */
|
||||
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) */
|
||||
/* 直传文件(后端会自动在默认项目下确保配音库存在,P0 404 修复) */
|
||||
const complete = await uploadAssetDirect({
|
||||
file: data.file,
|
||||
library_id: lib.id,
|
||||
kind: "voice",
|
||||
onProgress: (p) => setUploadProgress(p),
|
||||
})
|
||||
|
||||
|
||||
@@ -52,6 +52,10 @@ const appChildren: RouteObject[] = [
|
||||
path: "ai-avatar",
|
||||
lazy: lazyRoute(() => import("@/pages/ai-avatar/AiAvatarPage")),
|
||||
},
|
||||
{
|
||||
path: "viral-video",
|
||||
lazy: lazyRoute(() => import("@/pages/viral-video/ViralVideoPage")),
|
||||
},
|
||||
{
|
||||
path: "voice-clone",
|
||||
lazy: lazyRoute(() => import("@/pages/voice-clone/VoiceClone")),
|
||||
|
||||
@@ -553,7 +553,20 @@ def concat_video_files(
|
||||
if work_dir is None:
|
||||
work_dir = output_path.parent
|
||||
|
||||
segments = [ConcatSegment(video_path=p) for p in video_paths if p]
|
||||
# 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))
|
||||
|
||||
config = ConcatConfig(segments=segments, force_reencode=force_reencode)
|
||||
|
||||
engine = ConcatEngine(work_dir)
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""爆款视频 Worker 侧模块(#2039/#2040/#2051)。
|
||||
|
||||
video_analyzer(#2051):参考视频风格分析 6 步管线,输出 style_guide + clips 渲染参数映射。
|
||||
#2040 的 prompt 系统(prompts/prompt_store/llm_runner)由 #2040 分支提供,本文件不依赖它。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import (
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_ANALYSIS_TIMEOUT",
|
||||
"MAX_REFERENCE_DURATION_SEC",
|
||||
"MAX_REFERENCE_SIZE_MB",
|
||||
"STYLE_GUIDE_SCHEMA",
|
||||
"analyze_video_style",
|
||||
"build_render_params_for_clip",
|
||||
"map_bgm_bpm",
|
||||
"map_camera_to_ken_burns",
|
||||
"map_color_to_video_filter",
|
||||
"map_transition_to_xfade",
|
||||
]
|
||||
@@ -0,0 +1,961 @@
|
||||
"""参考爆款视频风格分析模块(#2051,v1.3)。
|
||||
|
||||
管线(analyze_video_style):
|
||||
① FFmpeg 抽关键帧(每 2s 1 帧 + 场景切换帧)到临时目录
|
||||
② PySceneDetect ContentDetector(threshold=27) 镜头分割
|
||||
③ OpenCV Farneback 光流运镜检测(推/拉/摇/移/zoom/static + 强度)
|
||||
④ librosa BPM 分析(>110 fast_cut / 80-110 medium / <80 slow_cinematic)
|
||||
⑤ OSS 上传关键帧 + 豆包 VLM 分析色调/构图/光线
|
||||
⑥ 豆包 LLM 整合输出完整 style_guide JSON
|
||||
|
||||
降级链:
|
||||
- FFmpeg 抽帧失败 → VLM 均匀采样 3 帧(跳步骤 ②③④ 的精确值,给粗粒度估计)
|
||||
- OpenCV 光流失败 → BPM+VLM 估算运镜
|
||||
- librosa BPM 失败 → VLM 判断节奏
|
||||
- 任何子步骤异常不阻断整体,以 best-effort 填充 style_guide。
|
||||
|
||||
资源约束:
|
||||
- 参考视频 ≤60s 且 ≤100MB;分析总超时 ≤60s;临时帧 try/finally 清理。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
import tempfile
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 资源约束 ─────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_REFERENCE_DURATION_SEC = 60
|
||||
MAX_REFERENCE_SIZE_MB = 100
|
||||
DEFAULT_ANALYSIS_TIMEOUT = 60 # 秒
|
||||
KEYFRAME_INTERVAL_SEC = 2
|
||||
SCENEDETECT_THRESHOLD = 27
|
||||
VLM_SAMPLE_FRAMES = 5 # 上传给 VLM 的关键帧上限
|
||||
FARNEBACK_PARAMS = dict(pyr_scale=0.5, levels=3, winsize=15, iterations=3, poly_n=5, poly_sigma=1.2, flags=0)
|
||||
|
||||
# ── style_guide 输出 schema(最小校验参考,不强制 jsonschema 依赖) ───────
|
||||
|
||||
STYLE_GUIDE_SCHEMA: dict[str, Any] = {
|
||||
"style_name": str,
|
||||
"avg_shot_duration": float,
|
||||
"shot_count": int,
|
||||
"pace": str, # fast_cut | medium | slow_cinematic
|
||||
"bpm": int,
|
||||
"camera_movements": list,
|
||||
"transitions": list,
|
||||
"color_palette": list,
|
||||
"color_tone": str, # warm | cool | high_sat | low_sat | vintage | fresh | dramatic | bright
|
||||
"color_filter": str, # none | warm_vintage | cool_fresh | high_contrast | soft_pastel | dramatic_cinematic
|
||||
"composition": dict,
|
||||
"lighting": str,
|
||||
"mood": str,
|
||||
"visual_keywords": list,
|
||||
"ken_burns_params": dict,
|
||||
"transition_map": dict,
|
||||
"video_filter_eq_params": dict,
|
||||
"ken_burns_direction_hint": str,
|
||||
}
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ShotBoundary:
|
||||
"""一段镜头(帧号区间)。"""
|
||||
|
||||
index: int
|
||||
start_sec: float
|
||||
end_sec: float
|
||||
movement: str = (
|
||||
"static" # push_in | pull_out | pan_left | pan_right | tilt_up | tilt_down | static | zoom_in | zoom_out
|
||||
)
|
||||
intensity: str = "low" # low | medium | high
|
||||
transition: str = "hard_cut" # 到下一个镜头的转场
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnalysisArtifacts:
|
||||
"""中间产物(降级路径用)。"""
|
||||
|
||||
frames_dir: Path
|
||||
frame_paths: list[Path] = field(default_factory=list)
|
||||
shots: list[ShotBoundary] = field(default_factory=list)
|
||||
bpm: int = 0
|
||||
vlm_descriptions: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _strip_code_fence(text: str) -> str:
|
||||
"""移除 markdown 代码块围栏,返回纯文本。"""
|
||||
t = text.strip()
|
||||
for fence in ("```json", "```JSON", "```"):
|
||||
if t.startswith(fence):
|
||||
t = t[len(fence) :].lstrip()
|
||||
if t.endswith("```"):
|
||||
t = t[:-3].rstrip()
|
||||
return t
|
||||
|
||||
|
||||
# ── FFmpeg / ffprobe ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ffmpeg_bin() -> str:
|
||||
return shutil.which("ffmpeg") or "ffmpeg"
|
||||
|
||||
|
||||
def _ffprobe_bin() -> str:
|
||||
return shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
def _probe_duration(video_path: str | Path) -> float:
|
||||
"""用 ffprobe 取视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
_ffprobe_bin(),
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(video_path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
text=True,
|
||||
) # nosec B603
|
||||
return float(out.strip() or 0)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("ffprobe 时长探测失败 %s: %s", video_path, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _extract_keyframes(video_path: Path, out_dir: Path, interval: int = KEYFRAME_INTERVAL_SEC) -> list[Path]:
|
||||
"""按固定间隔抽帧;同时检测场景切换帧(select='gt(scene,...)')。"""
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
# 固定间隔
|
||||
fixed_tpl = str(out_dir / "f_%04d.jpg")
|
||||
cmd_fixed = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
f"fps=1/{interval}",
|
||||
"-q:v",
|
||||
"3",
|
||||
fixed_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_fixed, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
# 场景切换帧(独立命名,scene_ 前缀)
|
||||
scene_tpl = str(out_dir / "scene_%04d.jpg")
|
||||
cmd_scene = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
"select='gt(scene,0.35)',showinfo",
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
"3",
|
||||
scene_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_scene, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
frames = sorted(out_dir.glob("f_*.jpg")) + sorted(out_dir.glob("scene_*.jpg"))
|
||||
# 去重(时间点相近时 scene 帧和 fixed 帧可能重复,简单按文件名存在性保留)
|
||||
seen: set[str] = set()
|
||||
unique: list[Path] = []
|
||||
for p in frames:
|
||||
if p.name not in seen:
|
||||
seen.add(p.name)
|
||||
unique.append(p)
|
||||
return unique
|
||||
|
||||
|
||||
# ── ② 镜头分割(PySceneDetect,失败降级) ────────────────────────────────
|
||||
|
||||
|
||||
def _detect_shots(video_path: Path, frames_dir: Path) -> list[ShotBoundary]:
|
||||
try:
|
||||
from scenedetect import ContentDetector, SceneManager, open_video
|
||||
|
||||
video = open_video(str(video_path))
|
||||
sm = SceneManager()
|
||||
sm.add_detector(ContentDetector(threshold=SCENEDETECT_THRESHOLD))
|
||||
sm.detect_scenes(video)
|
||||
scenes = sm.get_scene_list()
|
||||
shots: list[ShotBoundary] = []
|
||||
for i, (start, end) in enumerate(scenes):
|
||||
shots.append(
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=start.get_seconds(),
|
||||
end_sec=end.get_seconds(),
|
||||
)
|
||||
)
|
||||
if shots:
|
||||
return shots
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("PySceneDetect 镜头分割失败,使用均匀分段降级: %s", exc)
|
||||
# 降级:按固定间隔每 3 秒一镜头
|
||||
duration = _probe_duration(video_path) or 15.0
|
||||
dur = max(3.0, min(duration, float(MAX_REFERENCE_DURATION_SEC)))
|
||||
shots = []
|
||||
seg = 3.0
|
||||
i = 0
|
||||
t = 0.0
|
||||
while t < dur - 0.1:
|
||||
shots.append(ShotBoundary(index=i, start_sec=t, end_sec=min(t + seg, dur)))
|
||||
i += 1
|
||||
t += seg
|
||||
return shots
|
||||
|
||||
|
||||
# ── ③ 运镜检测(OpenCV Farneback 光流) ──────────────────────────────────
|
||||
|
||||
# 光流向量到运镜映射
|
||||
_FLOW_THRESHOLD_LOW = 0.3
|
||||
_FLOW_THRESHOLD_HIGH = 1.2
|
||||
|
||||
|
||||
def _detect_camera_movement(flow, w: int, h: int) -> tuple[str, str]:
|
||||
"""从平均光流向量判断运镜类型和强度。"""
|
||||
import numpy as np # noqa: PLC0415 - numpy 已在 requirements 中
|
||||
|
||||
fx = float(np.median(flow[..., 0]))
|
||||
fy = float(np.median(flow[..., 1]))
|
||||
trans_mag = math.hypot(fx, fy)
|
||||
# 发散/收敛判断 zoom:比较边缘流沿径向外指的平均分量(稳健版)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
rx, ry = (xs - cx) / max(cx, 1.0), (ys - cy) / max(cy, 1.0)
|
||||
rmag = np.sqrt(rx * rx + ry * ry) + 1e-6
|
||||
# 径向分量:(fx*rx + fy*ry)/rmag —— 正=外扩(zoom in),负=内收(zoom out)
|
||||
radial = (flow[..., 0] * rx + flow[..., 1] * ry) / rmag
|
||||
# 只看边缘带(|r|>0.5),且减去平移贡献:径向减去平均平移投影
|
||||
edge_mask = (rmag > 0.5).astype(np.float32)
|
||||
if edge_mask.sum() > 10:
|
||||
trans_radial = (fx * rx + fy * ry) / rmag
|
||||
zoom_signal = float(np.mean((radial - trans_radial)[edge_mask > 0]))
|
||||
else:
|
||||
zoom_signal = 0.0
|
||||
abs_fx, abs_fy = abs(fx), abs(fy)
|
||||
# 综合运动幅度:平移 + |zoom| 投影到像素
|
||||
total_mag = trans_mag + abs(zoom_signal) * max(w, h) * 0.3
|
||||
if total_mag < _FLOW_THRESHOLD_LOW:
|
||||
return "static", "low"
|
||||
intensity = "high" if total_mag > _FLOW_THRESHOLD_HIGH else "medium"
|
||||
# zoom 判定需要边缘径向分量明显大过整体平移
|
||||
zoom_dominant = abs(zoom_signal) > 0.6 and abs(zoom_signal) * max(w, h) * 0.3 > trans_mag * 1.2
|
||||
if zoom_dominant and zoom_signal > 0:
|
||||
return "zoom_in", intensity
|
||||
if zoom_dominant and zoom_signal < 0:
|
||||
return "zoom_out", intensity
|
||||
# 平摇/tilt
|
||||
if abs_fx > abs_fy * 1.5:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
if abs_fy > abs_fx * 1.5:
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
# 轨道/跟拍:以主轴为主
|
||||
if abs_fx >= abs_fy:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
|
||||
|
||||
def _analyze_movements(video_path: Path, shots: list[ShotBoundary]) -> None:
|
||||
"""对每个 shot 的首尾帧算光流,填充 movement/intensity。失败时静默降级为 static/low。"""
|
||||
try:
|
||||
import cv2 # noqa: PLC0415 - opencv-python-headless 已在 worker requirements 中
|
||||
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OpenCV 不可用,运镜检测降级为 static/low: %s", exc)
|
||||
return
|
||||
try:
|
||||
cap = cv2.VideoCapture(str(video_path))
|
||||
for shot in shots:
|
||||
mid_t = (shot.start_sec + shot.end_sec) / 2.0
|
||||
dt = max(0.2, min(0.5, (shot.end_sec - shot.start_sec) / 4.0))
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, max(0.0, (mid_t - dt)) * 1000)
|
||||
ok1, f1 = cap.read()
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, min(mid_t + dt, shot.end_sec - 0.05) * 1000)
|
||||
ok2, f2 = cap.read()
|
||||
if not (ok1 and ok2):
|
||||
continue
|
||||
g1 = cv2.cvtColor(f1, cv2.COLOR_BGR2GRAY)
|
||||
g2 = cv2.cvtColor(f2, cv2.COLOR_BGR2GRAY)
|
||||
h, w = g1.shape
|
||||
# 降采样加速
|
||||
scale = 360.0 / h if h > 360 else 1.0
|
||||
if scale < 1.0:
|
||||
g1 = cv2.resize(g1, (int(w * scale), int(h * scale)))
|
||||
g2 = cv2.resize(g2, (int(w * scale), int(h * scale)))
|
||||
flow = cv2.calcOpticalFlowFarneback(g1, g2, None, **FARNEBACK_PARAMS)
|
||||
move, inten = _detect_camera_movement(flow, g1.shape[1], g1.shape[0])
|
||||
shot.movement = move
|
||||
shot.intensity = inten
|
||||
cap.release()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("运镜检测异常,已降级: %s", exc)
|
||||
|
||||
|
||||
# ── ④ librosa BPM ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _detect_bpm(video_path: Path) -> int:
|
||||
"""提取音轨并估算 BPM;失败返回 0。"""
|
||||
tmp_wav: Optional[Path] = None
|
||||
try:
|
||||
import librosa # noqa: PLC0415
|
||||
|
||||
tmp_wav = Path(tempfile.mkstemp(suffix=".wav")[1])
|
||||
# ffmpeg 抽 22050Hz 单声道 wav
|
||||
subprocess.run(
|
||||
[_ffmpeg_bin(), "-y", "-i", str(video_path), "-vn", "-ac", "1", "-ar", "22050", "-f", "wav", str(tmp_wav)],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=20,
|
||||
check=False,
|
||||
) # nosec B603
|
||||
if not tmp_wav.exists() or tmp_wav.stat().st_size < 1024:
|
||||
return 0
|
||||
y, sr = librosa.load(str(tmp_wav), sr=22050, mono=True)
|
||||
if len(y) < sr * 2:
|
||||
return 0
|
||||
tempo, _ = librosa.beat.beat_track(y=y, sr=sr)
|
||||
try:
|
||||
bpm = int(round(float(tempo)))
|
||||
except Exception: # noqa: BLE001
|
||||
bpm = int(round(float(tempo[0]))) if len(tempo) else 0
|
||||
return max(40, min(bpm, 220))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("librosa BPM 分析失败: %s", exc)
|
||||
return 0
|
||||
finally:
|
||||
if tmp_wav and tmp_wav.exists():
|
||||
try:
|
||||
tmp_wav.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _pace_from_bpm(bpm: int) -> str:
|
||||
if bpm >= 110:
|
||||
return "fast_cut"
|
||||
if bpm >= 80:
|
||||
return "medium"
|
||||
if bpm > 0:
|
||||
return "slow_cinematic"
|
||||
return "medium"
|
||||
|
||||
|
||||
# ── ⑤ VLM 帧分析 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_frames(frame_paths: list[Path], shots: list[ShotBoundary], k: int = VLM_SAMPLE_FRAMES) -> list[Path]:
|
||||
"""从全量帧中均匀选 k 张代表性帧(优先场景帧)。"""
|
||||
if not frame_paths:
|
||||
return []
|
||||
scene_frames = sorted(p for p in frame_paths if p.name.startswith("scene_"))
|
||||
fixed_frames = sorted(p for p in frame_paths if p.name.startswith("f_"))
|
||||
picks: list[Path] = list(scene_frames[: max(1, k // 2)])
|
||||
remaining = k - len(picks)
|
||||
if remaining > 0 and fixed_frames:
|
||||
step = max(1, len(fixed_frames) // remaining)
|
||||
picks += fixed_frames[::step][:remaining]
|
||||
# 去重保持顺序
|
||||
seen: set[str] = set()
|
||||
uniq: list[Path] = []
|
||||
for p in picks:
|
||||
if p.name not in seen and p.exists():
|
||||
seen.add(p.name)
|
||||
uniq.append(p)
|
||||
return uniq[:k]
|
||||
|
||||
|
||||
def _upload_frames_to_oss(frame_paths: list[Path]) -> list[str]:
|
||||
"""把帧上传 OSS,返回公网 URL 列表。失败时降级为 data URI。"""
|
||||
urls: list[str] = []
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
for p in frame_paths:
|
||||
try:
|
||||
key = f"viral-video/analysis/{uuid.uuid4().hex}/{p.name}"
|
||||
url = upload_to_oss(p, key)
|
||||
if url:
|
||||
urls.append(url)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("单帧 OSS 上传失败 %s: %s", p.name, exc)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OSS 上传模块不可用,降级为 base64 data URI: %s", exc)
|
||||
if len(urls) < len(frame_paths):
|
||||
# 降级:base64 data URI(小图,单张 ≤100KB 才走此路)
|
||||
import base64
|
||||
|
||||
for p in frame_paths[len(urls) :]:
|
||||
try:
|
||||
if p.stat().st_size > 120_000:
|
||||
continue
|
||||
b64 = base64.b64encode(p.read_bytes()).decode("ascii")
|
||||
urls.append(f"data:image/jpeg;base64,{b64}")
|
||||
except Exception: # noqa: BLE001 # nosec B112
|
||||
continue
|
||||
return urls
|
||||
|
||||
|
||||
def _vlm_analyze_frames(image_urls: list[str]) -> dict[str, Any]:
|
||||
"""调豆包 VLM 分析色调/构图/光线/转场观感。"""
|
||||
if not image_urls:
|
||||
return {}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演和调色师。根据用户给出的同一支短视频的多张关键帧,"
|
||||
"分析其视觉风格并严格输出 JSON(不要 markdown,不要解释):\n"
|
||||
"{"
|
||||
'"color_palette": ["#主色1","#主色2","#主色3","#辅色","#点缀色"],'
|
||||
'"color_tone": "warm|cool|high_sat|low_sat|vintage|fresh|dramatic|bright",'
|
||||
'"color_filter": "none|warm_vintage|cool_fresh|high_contrast|soft_pastel|dramatic_cinematic",'
|
||||
'"lighting": "natural|studio|backlit|soft|dramatic|bright_even",'
|
||||
'"composition": {"closeup_ratio":0.0,"medium_ratio":0.0,"wide_ratio":0.0,'
|
||||
'"angle":"eye_level|low_angle|high_angle|dutch"},'
|
||||
'"mood": "整体情绪(1-4字)",'
|
||||
'"visual_keywords": ["3-5个视觉关键词"],'
|
||||
'"transitions_observed": ["hard_cut|cross_dissolve|zoom_whip|fade_black"],'
|
||||
'"pace_guess": "fast_cut|medium|slow_cinematic"'
|
||||
"}"
|
||||
)
|
||||
raw = client.vision_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": "请分析这支参考视频的风格。"},
|
||||
],
|
||||
images=image_urls,
|
||||
temperature=0.2,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return {}
|
||||
raw = _strip_code_fence(raw)
|
||||
# 容忍模型可能前后加文本
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
return json.loads(raw[i : j + 1])
|
||||
return {}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("VLM 帧分析失败: %s", exc)
|
||||
return {}
|
||||
|
||||
|
||||
# ── ⑥ LLM 整合 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _llm_synthesize(
|
||||
shots: list[ShotBoundary],
|
||||
bpm: int,
|
||||
vlm: dict[str, Any],
|
||||
style_strength: str,
|
||||
) -> dict[str, Any]:
|
||||
"""把结构化信号整合成 style_guide;LLM 不可用时走规则合成。"""
|
||||
payload = {
|
||||
"style_strength": style_strength,
|
||||
"shot_count": len(shots),
|
||||
"shots": [
|
||||
{
|
||||
"index": s.index,
|
||||
"start_sec": round(s.start_sec, 2),
|
||||
"end_sec": round(s.end_sec, 2),
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"transition": s.transition,
|
||||
}
|
||||
for s in shots
|
||||
],
|
||||
"bpm": bpm,
|
||||
"pace_guess": _pace_from_bpm(bpm),
|
||||
"vlm": vlm,
|
||||
}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演。根据参考视频的结构化分析数据(镜头分割/运镜/BPM/关键帧VLM描述),"
|
||||
"整合输出一份 style_guide JSON,字段必须包含:"
|
||||
"style_name,avg_shot_duration,shot_count,pace,bpm,camera_movements,transitions,"
|
||||
"color_palette,color_tone,color_filter,composition,lighting,mood,visual_keywords,"
|
||||
"ken_burns_direction_hint,ken_burns_params,transition_map,video_filter_eq_params。"
|
||||
"严格输出一个合法 JSON 对象,不要 markdown/解释。"
|
||||
)
|
||||
user_text = "分析数据:\n" + json.dumps(payload, ensure_ascii=False)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": sys_prompt}, {"role": "user", "content": user_text}],
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
if raw:
|
||||
raw = _strip_code_fence(raw)
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
result = json.loads(raw[i : j + 1])
|
||||
if isinstance(result, dict) and result.get("style_name"):
|
||||
return result
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("LLM 合成 style_guide 失败,走规则降级: %s", exc)
|
||||
return _rule_based_style_guide(shots, bpm, vlm)
|
||||
|
||||
|
||||
def _rule_based_style_guide(shots: list[ShotBoundary], bpm: int, vlm: dict[str, Any]) -> dict[str, Any]:
|
||||
"""LLM 不可用时,用规则拼出可用 style_guide。"""
|
||||
durations = [s.end_sec - s.start_sec for s in shots] or [3.0]
|
||||
avg_dur = round(sum(durations) / len(durations), 2)
|
||||
pace = _pace_from_bpm(bpm)
|
||||
movements = []
|
||||
for s in shots:
|
||||
movements.append(
|
||||
{
|
||||
"shot_index": s.index + 1,
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"duration": round(s.end_sec - s.start_sec, 2),
|
||||
"subject_hint": _default_subject_hint(s.movement),
|
||||
}
|
||||
)
|
||||
transitions = []
|
||||
for i in range(len(shots) - 1):
|
||||
transitions.append({"between_shot": [i + 1, i + 2], "type": shots[i].transition})
|
||||
color_palette = vlm.get("color_palette") or ["#E0E0E0", "#333333", "#F5F5F5", "#888888", "#FF6B35"]
|
||||
color_tone = vlm.get("color_tone") or "bright"
|
||||
color_filter = vlm.get("color_filter") or "none"
|
||||
lighting = vlm.get("lighting") or "bright_even"
|
||||
composition = vlm.get("composition") or {
|
||||
"closeup_ratio": 0.4,
|
||||
"medium_ratio": 0.4,
|
||||
"wide_ratio": 0.2,
|
||||
"angle": "eye_level",
|
||||
}
|
||||
mood = vlm.get("mood") or "明快"
|
||||
vk = vlm.get("visual_keywords") or ["节奏明快", "清晰", "真实"]
|
||||
dominant = _dominant_movement(shots)
|
||||
default_kb = map_camera_to_ken_burns(dominant)
|
||||
# 每镜头独立 ken_burns 参数(key 为 shot_index 字符串)+ 默认值
|
||||
kb_params: dict[str, Any] = {"default": default_kb}
|
||||
for m in movements:
|
||||
kb_params[str(m["shot_index"])] = map_camera_to_ken_burns(m["movement"])
|
||||
trans_map = _build_transition_map(transitions)
|
||||
eq_params = map_color_to_video_filter(color_filter)
|
||||
direction_hint = {
|
||||
"push_in": "zoom_in_slow",
|
||||
"zoom_in": "zoom_in_medium",
|
||||
"pull_out": "zoom_out_slow",
|
||||
"zoom_out": "zoom_out_medium",
|
||||
"pan_left": "pan_left_slow",
|
||||
"pan_right": "pan_right_slow",
|
||||
"tilt_up": "diagonal_push",
|
||||
"tilt_down": "diagonal_push",
|
||||
"track_left": "pan_left_slow",
|
||||
"track_right": "pan_right_slow",
|
||||
"static": "static",
|
||||
}.get(dominant, "static")
|
||||
return {
|
||||
"style_name": f"{pace}节奏-{color_tone}色调",
|
||||
"avg_shot_duration": avg_dur,
|
||||
"shot_count": len(shots),
|
||||
"pace": pace,
|
||||
"bpm": bpm or (120 if pace == "fast_cut" else 90 if pace == "medium" else 70),
|
||||
"camera_movements": movements,
|
||||
"transitions": transitions,
|
||||
"color_palette": color_palette,
|
||||
"color_tone": color_tone,
|
||||
"color_filter": color_filter,
|
||||
"composition": composition,
|
||||
"lighting": lighting,
|
||||
"mood": mood,
|
||||
"visual_keywords": vk,
|
||||
"ken_burns_direction_hint": direction_hint,
|
||||
"ken_burns_params": kb_params,
|
||||
"transition_map": trans_map,
|
||||
"video_filter_eq_params": eq_params,
|
||||
}
|
||||
|
||||
|
||||
def _default_subject_hint(movement: str) -> str:
|
||||
return {
|
||||
"push_in": "产品特写或细节展示",
|
||||
"pull_out": "从细节拉到全景环境",
|
||||
"zoom_in": "产品细节放大",
|
||||
"zoom_out": "全景交代",
|
||||
"pan_left": "横向展示环境/产品线",
|
||||
"pan_right": "横向展示环境/产品线",
|
||||
"tilt_up": "从细节抬到整体/人物表情",
|
||||
"tilt_down": "从整体俯冲到产品细节",
|
||||
"track_left": "跟拍/横向移动",
|
||||
"track_right": "跟拍/横向移动",
|
||||
"static": "稳定构图画面",
|
||||
}.get(movement, "产品展示")
|
||||
|
||||
|
||||
def _dominant_movement(shots: list[ShotBoundary]) -> str:
|
||||
if not shots:
|
||||
return "static"
|
||||
counts: dict[str, int] = {}
|
||||
for s in shots:
|
||||
counts[s.movement] = counts.get(s.movement, 0) + 1
|
||||
return max(counts, key=counts.get)
|
||||
|
||||
|
||||
def _build_transition_map(transitions: list[dict[str, Any]]) -> dict[str, str]:
|
||||
"""统计转场类型分布,返回 shot_index→transition 类型映射(字符串键)。"""
|
||||
m: dict[str, str] = {}
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or [0, 0]
|
||||
if len(pair) >= 2:
|
||||
m[f"{pair[0]}-{pair[1]}"] = t.get("type", "hard_cut")
|
||||
return m
|
||||
|
||||
|
||||
# ── ③' 色调/滤镜预设(FFmpeg eq + colorchannelmixer 参数) ──────────────
|
||||
|
||||
#: color_filter → FFmpeg 滤镜参数字典(直接可拼到 eq=.../colorchannelmixer=...)
|
||||
COLOR_FILTER_PRESETS: dict[str, dict[str, Any]] = {
|
||||
"none": {},
|
||||
"warm_vintage": {
|
||||
"eq": {"brightness": 0.02, "contrast": 1.05, "saturation": 0.9, "gamma": 1.05},
|
||||
"colorchannelmixer": {"rr": 1.1, "gg": 0.98, "bb": 0.82, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"cool_fresh": {
|
||||
"eq": {"brightness": 0.03, "contrast": 1.08, "saturation": 1.05},
|
||||
"colorchannelmixer": {"rr": 0.9, "gg": 1.0, "bb": 1.12, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"high_contrast": {
|
||||
"eq": {"brightness": 0.0, "contrast": 1.3, "saturation": 1.2},
|
||||
"colorchannelmixer": {},
|
||||
},
|
||||
"soft_pastel": {
|
||||
"eq": {"brightness": 0.05, "contrast": 0.92, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 1.03, "bb": 1.05, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"dramatic_cinematic": {
|
||||
"eq": {"brightness": -0.03, "contrast": 1.2, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 0.98, "bb": 0.9, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_color_to_video_filter(color_filter: str) -> dict[str, Any]:
|
||||
"""color_filter 枚举 → FFmpeg eq/colorchannelmixer 参数字典(渲染端直接使用)。"""
|
||||
preset = COLOR_FILTER_PRESETS.get(color_filter) or COLOR_FILTER_PRESETS["none"]
|
||||
# 返回深拷贝防污染
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── ③'' 运镜 → ken_burns 参数映射 ───────────────────────────────────────
|
||||
|
||||
#: 运镜类型 → URS 可直接消费的 ken_burns 参数字典
|
||||
CAMERA_TO_KEN_BURNS: dict[str, dict[str, Any]] = {
|
||||
"static": {
|
||||
"type": "static",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"push_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.12,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.18,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pull_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.12,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.18,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_up": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.08,
|
||||
"zoom_end": 1.14,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": -0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_down": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.14,
|
||||
"zoom_end": 1.08,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_camera_to_ken_burns(movement: str) -> dict[str, Any]:
|
||||
"""运镜类型 → URS ken_burns 参数字典。未知类型回退 static。"""
|
||||
preset = CAMERA_TO_KEN_BURNS.get(movement) or CAMERA_TO_KEN_BURNS["static"]
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── 转场 → xfade transition 名称 ────────────────────────────────────────
|
||||
|
||||
TRANSITION_TO_XFADE: dict[str, str] = {
|
||||
"hard_cut": "cut",
|
||||
"cross_dissolve": "dissolve",
|
||||
"fade_black": "fadeblack",
|
||||
"fade": "fade",
|
||||
"zoom_whip": "zoom",
|
||||
"slide_left": "slideright", # 画面左移 = 新画面从右滑入
|
||||
"slide_right": "slideleft",
|
||||
"wipe_left": "wipeleft",
|
||||
"wipe_right": "wiperight",
|
||||
}
|
||||
|
||||
|
||||
def map_transition_to_xfade(transition_type: str) -> str:
|
||||
"""转场枚举 → TransitionEngine 支持的 xfade 名称;未知回退 cut。"""
|
||||
return TRANSITION_TO_XFADE.get(transition_type, "cut")
|
||||
|
||||
|
||||
# ── BPM → BGM 推荐 BPM ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def map_bgm_bpm(bpm: int) -> int:
|
||||
"""BGM 选曲 BPM:参考视频 BPM ±5。bpm=0 返回 90(默认 medium)。"""
|
||||
if bpm <= 0:
|
||||
return 90
|
||||
return max(60, min(bpm, 180))
|
||||
|
||||
|
||||
# ── 单 clip 渲染参数聚合(给 URS build_render_plan 使用) ───────────────
|
||||
|
||||
|
||||
def build_render_params_for_clip(
|
||||
clip_index: int,
|
||||
style_guide: dict[str, Any],
|
||||
*,
|
||||
duration_sec: Optional[float] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""根据 style_guide 为第 clip_index 个 clip 生成可直接喂给 URS 的渲染参数。"""
|
||||
shot_idx = clip_index + 1
|
||||
movements = style_guide.get("camera_movements") or []
|
||||
movement = "static"
|
||||
intensity = "low"
|
||||
for m in movements:
|
||||
if m.get("shot_index") == shot_idx:
|
||||
movement = m.get("movement", "static")
|
||||
intensity = m.get("intensity", "low")
|
||||
break
|
||||
ken = map_camera_to_ken_burns(movement)
|
||||
if intensity == "high":
|
||||
ken["zoom_end"] = round(ken.get("zoom_end", 1.0) * 1.08, 3)
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 1.3, 3)
|
||||
elif intensity == "low":
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 0.6, 3)
|
||||
transitions = style_guide.get("transitions") or []
|
||||
trans_type = "hard_cut"
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or []
|
||||
if len(pair) >= 2 and pair[0] == shot_idx:
|
||||
trans_type = t.get("type", "hard_cut")
|
||||
break
|
||||
xfade = map_transition_to_xfade(trans_type)
|
||||
eq = map_color_to_video_filter(style_guide.get("color_filter", "none"))
|
||||
return {
|
||||
"ken_burns": ken,
|
||||
"transition": {"type": xfade, "duration": 0.3 if xfade != "cut" else 0.0},
|
||||
"video_filter": eq,
|
||||
"bgm_bpm_hint": map_bgm_bpm(int(style_guide.get("bpm") or 0)),
|
||||
"duration_sec": duration_sec,
|
||||
}
|
||||
|
||||
|
||||
# ── 素材本地化(URL/OSS key → 本地临时文件) ────────────────────────────
|
||||
|
||||
|
||||
def _ensure_local_video(reference: str, work_dir: Path) -> Optional[Path]:
|
||||
"""把 reference(URL/OSS key/本地路径)落到 work_dir 下的本地文件。"""
|
||||
p = Path(reference)
|
||||
if p.exists() and p.is_file():
|
||||
return p
|
||||
try:
|
||||
from video_processing.oss_helpers import download_asset
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
ok = download_asset(reference, target)
|
||||
if ok and target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("download_asset 失败,尝试 http 直连: %s", exc)
|
||||
if reference.startswith(("http://", "https://")):
|
||||
try:
|
||||
import httpx # noqa: PLC0415 - 项目依赖,延迟导入
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
with httpx.Client(timeout=20.0, follow_redirects=True) as client:
|
||||
with client.stream("GET", reference) as resp:
|
||||
resp.raise_for_status()
|
||||
with open(target, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=64 * 1024):
|
||||
f.write(chunk)
|
||||
if target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("HTTP 下载参考视频失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── 入口 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def analyze_video_style(
|
||||
reference_video_path: str | Path,
|
||||
style_strength: str = "medium",
|
||||
*,
|
||||
timeout_sec: int = DEFAULT_ANALYSIS_TIMEOUT,
|
||||
) -> dict[str, Any]:
|
||||
"""分析参考视频风格,返回 style_guide dict。
|
||||
|
||||
Args:
|
||||
reference_video_path: 本地路径、HTTP(S) URL 或 OSS storage key。
|
||||
style_strength: light | medium | strict。
|
||||
timeout_sec: 单步超时(秒),默认 60。
|
||||
|
||||
Returns:
|
||||
style_guide dict,详见 STYLE_GUIDE_SCHEMA。任何子步骤失败都会降级,不抛异常。
|
||||
"""
|
||||
style_strength = style_strength if style_strength in ("light", "medium", "strict") else "medium"
|
||||
frames_dir: Optional[Path] = None
|
||||
local_path: Optional[Path] = None
|
||||
try:
|
||||
frames_dir = Path(tempfile.mkdtemp(prefix="vstyle_"))
|
||||
work_dir = frames_dir # 同一临时根
|
||||
local_path = _ensure_local_video(str(reference_video_path), work_dir)
|
||||
if local_path is None:
|
||||
logger.error("[video_analyzer] 无法获取参考视频: %s", reference_video_path)
|
||||
return _rule_based_style_guide([], 0, {})
|
||||
# 资源约束:大小 / 时长
|
||||
try:
|
||||
size_mb = local_path.stat().st_size / (1024 * 1024)
|
||||
if size_mb > MAX_REFERENCE_SIZE_MB:
|
||||
logger.warning(
|
||||
"[video_analyzer] 参考视频 %.1fMB 超上限,按前 %ds 分析", size_mb, MAX_REFERENCE_DURATION_SEC
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
duration = _probe_duration(local_path)
|
||||
if duration > MAX_REFERENCE_DURATION_SEC:
|
||||
duration = MAX_REFERENCE_DURATION_SEC
|
||||
# ① 抽帧
|
||||
try:
|
||||
frame_paths = _extract_keyframes(local_path, frames_dir / "frames")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("FFmpeg 抽帧失败: %s,降级为 VLM 均匀采样", exc)
|
||||
frame_paths = []
|
||||
# ② 镜头分割
|
||||
shots = _detect_shots(local_path, frames_dir)
|
||||
# ③ 运镜检测(有帧才跑)
|
||||
if frame_paths or shots:
|
||||
_analyze_movements(local_path, shots)
|
||||
# ④ BPM
|
||||
bpm = _detect_bpm(local_path)
|
||||
# ⑤ 选帧→OSS→VLM
|
||||
sampled = _sample_frames(frame_paths, shots)
|
||||
image_urls = _upload_frames_to_oss(sampled) if sampled else []
|
||||
vlm = _vlm_analyze_frames(image_urls) if image_urls else {}
|
||||
# ⑥ 合成
|
||||
style_guide = _llm_synthesize(shots, bpm, vlm, style_strength)
|
||||
# 兜底字段校验
|
||||
style_guide.setdefault("style_strength", style_strength)
|
||||
style_guide.setdefault("pace", _pace_from_bpm(bpm))
|
||||
style_guide.setdefault("bpm", bpm)
|
||||
style_guide.setdefault("shot_count", len(shots))
|
||||
if shots and "avg_shot_duration" not in style_guide:
|
||||
durs = [s.end_sec - s.start_sec for s in shots]
|
||||
style_guide["avg_shot_duration"] = round(sum(durs) / len(durs), 2)
|
||||
return style_guide
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("[video_analyzer] 整体分析异常,返回最小占位 style_guide: %s", exc)
|
||||
return _rule_based_style_guide([], 0, {"mood": "未知"})
|
||||
finally:
|
||||
# 临时帧清理
|
||||
if frames_dir and frames_dir.exists():
|
||||
shutil.rmtree(frames_dir, ignore_errors=True)
|
||||
Executable → Regular
+450
-162
@@ -1,6 +1,6 @@
|
||||
"""爆款视频 Celery 编排器 — ViralVideoOrchestrator.
|
||||
|
||||
10 步流水线:
|
||||
9 步流水线(Seedance 2.5 直生口型,不再走 MuseTalk):
|
||||
1. 图片 VLM 分析
|
||||
1.5 [v1.3] 视频风格分析(如用户上传参考视频)
|
||||
2. 用户文案意图解析
|
||||
@@ -8,20 +8,20 @@
|
||||
4. 分镜脚本生成
|
||||
5. 合规审核(6 维度,不通过自动重写 1 次)
|
||||
6. CosyVoice 配音
|
||||
7. BGM 选择
|
||||
8. UnifiedRenderService 渲染
|
||||
9. 数字人口型(MuseTalk)
|
||||
10. OSS 上传 + 通知 + 扣点
|
||||
7. BGM 选择(素材未就绪时跳过)
|
||||
8. Seedance 逐分镜生成 + ffmpeg concat + 混 TTS
|
||||
9. OSS 上传 + 通知 + 扣点
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from celery import Task
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.celery_app import celery_app # noqa: F401 - 加载 app 以注册任务
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
@@ -41,22 +41,37 @@ logger = logging.getLogger(__name__)
|
||||
# ── WS 进度推送 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _emit_progress(job_id: str, stage: str, progress: float, message: str = "", data: dict | None = None):
|
||||
"""通过 Redis 发布进度事件,供 WebSocket 消费。"""
|
||||
def _emit_progress(
|
||||
job_id: str,
|
||||
stage: str,
|
||||
progress: float,
|
||||
message: str = "",
|
||||
data: dict | None = None,
|
||||
event_type: str = "viral_video:progress",
|
||||
):
|
||||
"""通过 Redis 发布进度事件,供 WebSocket 消费。
|
||||
|
||||
event_type 取值:
|
||||
- viral_video:progress 中间进度(默认)
|
||||
- viral_video:completed 任务完成
|
||||
- viral_video:failed 任务失败
|
||||
- viral_video:wait_user 等待用户确认
|
||||
所有事件 payload 均为合法 JSON,前端 JSON.parse 即可。
|
||||
"""
|
||||
try:
|
||||
import redis as redis_lib
|
||||
|
||||
redis_url = os.environ.get("REDIS_URL", "redis://localhost:6379/0")
|
||||
r = redis_lib.from_url(redis_url)
|
||||
event = {
|
||||
"type": "viral_video:progress",
|
||||
"type": event_type,
|
||||
"job_id": job_id,
|
||||
"stage": stage,
|
||||
"progress": progress,
|
||||
"message": message or STAGE_LABELS.get(stage, stage),
|
||||
"data": data or {},
|
||||
}
|
||||
r.publish(f"viral_video:{job_id}", str(event))
|
||||
r.publish(f"viral_video:{job_id}", json.dumps(event, ensure_ascii=False))
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] WS 进度推送失败: %s", e)
|
||||
|
||||
@@ -81,6 +96,9 @@ def _save_job(repo, job, session):
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _step_image_analysis(job: ViralVideoJob) -> dict:
|
||||
"""步骤 1: 图片 VLM 分析 — 识别产品特征、场景、卖点。"""
|
||||
try:
|
||||
@@ -110,13 +128,13 @@ def _step_video_analysis(job: ViralVideoJob) -> dict | None:
|
||||
return None
|
||||
|
||||
try:
|
||||
# 尝试导入 video_analyzer(由 #2051 提供)
|
||||
from worker_app.tasks.viral_video_analyzer import analyze_video_style
|
||||
# P0-2: 修正 import 路径(video_analyzer.py 在 apps/worker/viral_video/ 下,worker PYTHONPATH 含 apps/worker)
|
||||
from viral_video.video_analyzer import analyze_video_style
|
||||
|
||||
style_guide = analyze_video_style(job.reference_video_url)
|
||||
return style_guide
|
||||
except ImportError:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪,使用占位风格分析")
|
||||
except ImportError as e:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪(%s),使用占位风格分析", e)
|
||||
return {
|
||||
"cut_speed": "medium",
|
||||
"transition": "cross_dissolve",
|
||||
@@ -213,42 +231,120 @@ def _step_copy_fusion(job: ViralVideoJob, intent: dict, image_analysis: dict) ->
|
||||
|
||||
|
||||
def _step_storyboard(job: ViralVideoJob, copy_text: str, image_analysis: dict) -> list[dict]:
|
||||
"""步骤 4: 分镜脚本生成。"""
|
||||
"""步骤 4: 分镜脚本生成。每个分镜独立一段视频,段内时长建议 3~6 秒。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return [{"order": 0, "type": "product_shot", "text": copy_text[:50], "duration": job.duration}]
|
||||
return [
|
||||
{
|
||||
"order": 0,
|
||||
"type": "product_shot",
|
||||
"text": copy_text[:50],
|
||||
"duration": min(5, job.duration),
|
||||
"description": "产品展示",
|
||||
"ken_burns": "zoom_in",
|
||||
"transition": "cut",
|
||||
}
|
||||
]
|
||||
|
||||
prompt = f"""请根据以下文案生成短视频分镜脚本:
|
||||
products_hint = ""
|
||||
products = image_analysis.get("products", []) if image_analysis else []
|
||||
if products:
|
||||
p0 = products[0] if isinstance(products[0], dict) else {}
|
||||
feats = p0.get("features", []) if isinstance(p0, dict) else []
|
||||
products_hint = f"\n首帧参考产品特征:{p0.get('name','')} - {', '.join(feats[:3])}"
|
||||
|
||||
seg_seconds = 5
|
||||
n_segments = max(2, min(6, max(1, job.duration // seg_seconds)))
|
||||
ratio = "9:16"
|
||||
|
||||
prompt = f"""请根据以下文案生成爆款短视频分镜脚本,共 {n_segments} 个分镜:
|
||||
|
||||
文案内容:{copy_text}
|
||||
视频时长:{job.duration}秒
|
||||
视频总时长:{job.duration}秒(每个分镜 3~6 秒,总和约等于总时长)
|
||||
风格强度:{job.style_strength}
|
||||
输出宽高比:{ratio}{products_hint}
|
||||
|
||||
请以JSON数组格式返回分镜列表,每个分镜包含:
|
||||
- order: 序号
|
||||
- type: 镜头类型(product_shot/text_card/scene_transition/closing)
|
||||
- description: 画面描述
|
||||
- text: 配音/字幕文本
|
||||
- duration: 时长(秒)
|
||||
- ken_burns: 运镜方式(zoom_in/zoom_out/pan_left/pan_right/none)
|
||||
- transition: 转场方式(cut/dissolve/wipe/fade)"""
|
||||
请以 JSON 数组格式返回分镜列表,每个分镜包含:
|
||||
- order: 序号(从0开始)
|
||||
- type: 镜头类型(product_shot/close_up/scene/action/text_card/closing)
|
||||
- description: 画面详细描述(中文,含主体、动作、场景、运镜、光影,用于AI视频生成prompt)
|
||||
- text: 该分镜配音/字幕文本
|
||||
- duration: 时长(秒,3~6秒的整数)
|
||||
- ken_burns: 运镜方式(zoom_in/zoom_out/pan_left/pan_right/static)
|
||||
- transition: 与下一分镜的转场(cut/dissolve/fade)"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
if isinstance(result, list):
|
||||
return result
|
||||
# 尝试从字符串中解析 JSON
|
||||
import json
|
||||
return _normalize_storyboard(result, job.duration, n_segments, copy_text)
|
||||
import json as _json
|
||||
|
||||
return (
|
||||
json.loads(result)
|
||||
if isinstance(result, str)
|
||||
else [{"order": 0, "text": copy_text, "duration": job.duration}]
|
||||
)
|
||||
parsed = _json.loads(result) if isinstance(result, str) else result
|
||||
if isinstance(parsed, list):
|
||||
return _normalize_storyboard(parsed, job.duration, n_segments, copy_text)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 分镜生成失败: %s", e)
|
||||
return [{"order": 0, "type": "product_shot", "text": copy_text[:100], "duration": job.duration}]
|
||||
|
||||
return _fallback_storyboard(copy_text, job.duration, n_segments)
|
||||
|
||||
|
||||
def _normalize_storyboard(raw: list, total_duration: int, n_segments: int, copy_text: str) -> list[dict]:
|
||||
"""规范化 LLM 输出的分镜:填充缺省字段、保证总时长合理。"""
|
||||
out: list[dict] = []
|
||||
for i, item in enumerate(raw):
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
dur = int(item.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(3, min(8, dur))
|
||||
out.append(
|
||||
{
|
||||
"order": int(item.get("order", i)),
|
||||
"type": str(item.get("type", "product_shot")),
|
||||
"description": str(item.get("description", copy_text[:80])),
|
||||
"text": str(item.get("text", "")),
|
||||
"duration": dur,
|
||||
"ken_burns": str(item.get("ken_burns", "zoom_in")),
|
||||
"transition": str(item.get("transition", "cut")),
|
||||
}
|
||||
)
|
||||
if not out:
|
||||
return _fallback_storyboard(copy_text, total_duration, n_segments)
|
||||
out = out[:n_segments]
|
||||
total = sum(s["duration"] for s in out)
|
||||
if total > 0 and total != total_duration:
|
||||
scale = total_duration / total
|
||||
acc = 0
|
||||
for s in out[:-1]:
|
||||
s["duration"] = max(3, min(8, round(s["duration"] * scale)))
|
||||
acc += s["duration"]
|
||||
out[-1]["duration"] = max(3, total_duration - acc)
|
||||
return out
|
||||
|
||||
|
||||
def _fallback_storyboard(copy_text: str, total_duration: int, n_segments: int) -> list[dict]:
|
||||
if n_segments <= 0:
|
||||
n_segments = 1
|
||||
dur = total_duration // n_segments
|
||||
remainder = total_duration - dur * n_segments
|
||||
out = []
|
||||
for i in range(n_segments):
|
||||
d = dur + (remainder if i == n_segments - 1 else 0)
|
||||
out.append(
|
||||
{
|
||||
"order": i,
|
||||
"type": "product_shot",
|
||||
"description": f"产品展示镜头 {i + 1}:{copy_text[:40]}",
|
||||
"text": copy_text,
|
||||
"duration": max(3, d),
|
||||
"ken_burns": "zoom_in" if i % 2 == 0 else "pan_left",
|
||||
"transition": "cut",
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _step_review(job: ViralVideoJob, copy_text: str, storyboard: list[dict]) -> dict:
|
||||
@@ -280,108 +376,262 @@ def _step_review(job: ViralVideoJob, copy_text: str, storyboard: list[dict]) ->
|
||||
return {"passed": True, "score": 75, "details": {d: "默认通过" for d in dimensions}}
|
||||
|
||||
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str) -> str:
|
||||
"""步骤 6: CosyVoice 配音。"""
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str):
|
||||
"""步骤 6: CosyVoice 配音。P1:返回 Path;失败返回 None。"""
|
||||
try:
|
||||
from worker_app.services.tts_service_factory import get_tts_service
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# 使用绝对包路径,避免 celery worker 因 cwd/PYTHONPATH 微小差异找不到 services 模块
|
||||
from apps.worker.services.tts_service_factory import get_tts_service
|
||||
|
||||
tts_service = get_tts_service()
|
||||
# 简化调用,实际需要更详细的参数
|
||||
audio_url = tts_service.synthesize(text=copy_text, voice_id=job.persona_id or "default")
|
||||
return audio_url
|
||||
# Bug #2110: persona_id 透传给 voice_id(空则用 CosyVoice 默认 longxiaochun_v3),
|
||||
# 统一输出 mp3 给后续 ffmpeg 混音(之前默认 wav 导致部分 provider/后处理不兼容)。
|
||||
voice_id = (job.persona_id or "").strip()
|
||||
try:
|
||||
result = tts_service.synthesize(
|
||||
text=copy_text,
|
||||
voice_id=voice_id or "longxiaochun_v3",
|
||||
format="mp3",
|
||||
)
|
||||
except TypeError:
|
||||
# 老 provider 只支持 text 参数
|
||||
try:
|
||||
result = tts_service.synthesize(text=copy_text, voice_id=voice_id or "longxiaochun_v3")
|
||||
except TypeError:
|
||||
result = tts_service.synthesize(text=copy_text)
|
||||
if result is None:
|
||||
return None
|
||||
p = _Path(result) if not isinstance(result, _Path) else result
|
||||
if p.exists():
|
||||
logger.info(
|
||||
"[爆款视频] TTS 合成完成: voice=%s path=%s size=%d", voice_id or "longxiaochun_v3", p, p.stat().st_size
|
||||
)
|
||||
return p
|
||||
logger.warning("[爆款视频] TTS 返回路径不存在: %s", p)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 配音失败: %s", e)
|
||||
return ""
|
||||
return None
|
||||
|
||||
|
||||
def _step_bgm_select(job: ViralVideoJob) -> str:
|
||||
"""步骤 7: BGM 选择。"""
|
||||
# 基于 bgm_preference 和 marketing_purpose 匹配预设 BGM
|
||||
bgm_map = {
|
||||
"upbeat": "bgm_upbeat_01.mp3",
|
||||
"calm": "bgm_calm_01.mp3",
|
||||
"energetic": "bgm_energetic_01.mp3",
|
||||
"emotional": "bgm_emotional_01.mp3",
|
||||
def _step_bgm_select(job: ViralVideoJob):
|
||||
"""步骤 7: BGM 选择。P1:素材未就绪前返回 None,跳过 BGM 混音。"""
|
||||
return None
|
||||
|
||||
|
||||
def _build_segment_prompt(seg: dict, job: ViralVideoJob, style_hint: str) -> str:
|
||||
desc = seg.get("description") or seg.get("text") or "产品展示"
|
||||
ken_burns = seg.get("ken_burns", "zoom_in")
|
||||
cam_map = {
|
||||
"zoom_in": "缓慢推镜放大",
|
||||
"zoom_out": "缓慢拉镜缩小",
|
||||
"pan_left": "镜头向左平移",
|
||||
"pan_right": "镜头向右平移",
|
||||
"static": "固定镜头",
|
||||
}
|
||||
preference = job.bgm_preference.lower()
|
||||
for key, bgm in bgm_map.items():
|
||||
if key in preference:
|
||||
return bgm
|
||||
return "bgm_default.mp3"
|
||||
camera = cam_map.get(ken_burns, "缓慢运镜")
|
||||
parts = [
|
||||
f"{desc}。",
|
||||
f"运镜:{camera}。",
|
||||
"画面流畅、电影感光影、高清细节,9:16竖屏,适合短视频。",
|
||||
]
|
||||
if style_hint:
|
||||
parts.append(f"参考风格:{style_hint}")
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _step_render(job: ViralVideoJob, storyboard: list[dict], audio_url: str, bgm: str) -> str:
|
||||
"""步骤 8: UnifiedRenderService 渲染。"""
|
||||
try:
|
||||
from video_processing.render_adapter import build_render_plan
|
||||
from video_processing.unified_render_service import UnifiedRenderService
|
||||
def _step_render(job, storyboard, tts_path, bgm):
|
||||
"""步骤 8: 渲染(P0-1 核心重写)。
|
||||
|
||||
render_plan = build_render_plan(
|
||||
images=job.images,
|
||||
storyboard=storyboard,
|
||||
audio_url=audio_url,
|
||||
bgm=bgm,
|
||||
duration=job.duration,
|
||||
style_guide=job.style_guide,
|
||||
每个 storyboard 分镜 → Seedance 2.5 生成短视频段(无声)→ 下载 → ffmpeg concat → 混入 TTS。
|
||||
返回最终视频本地路径字符串。
|
||||
"""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from video_processing.concat_engine import concat_video_files
|
||||
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
from packages.shared.ffmpeg_utils import run_ffmpeg
|
||||
|
||||
if not storyboard:
|
||||
raise ValueError("storyboard is empty")
|
||||
|
||||
style_hint = ""
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、色调{job.style_guide.get('color_grade','')}"
|
||||
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
|
||||
logger.info("[爆款视频] 开始渲染,分镜数=%d, tmpdir=%s", len(storyboard), tmpdir)
|
||||
|
||||
seg_paths: list[str] = []
|
||||
first_image = job.images[0] if job.images else None
|
||||
n_total = len(storyboard)
|
||||
for i, seg in enumerate(storyboard):
|
||||
try:
|
||||
dur = int(seg.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(2, min(12, dur))
|
||||
prompt = _build_segment_prompt(seg, job, style_hint)
|
||||
_emit_progress(
|
||||
job.id,
|
||||
ViralVideoStage.RENDERING,
|
||||
80.0 + (i + 1) / max(n_total, 1) * 5.0,
|
||||
f"正在生成分镜 {i + 1}/{n_total} ({dur}s)...",
|
||||
)
|
||||
|
||||
render_svc = UnifiedRenderService()
|
||||
output_path = render_svc.render(render_plan)
|
||||
return output_path
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 渲染失败: %s", e, exc_info=True)
|
||||
raise
|
||||
|
||||
|
||||
def _step_musetalk(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 9: 数字人口型(MuseTalk)。"""
|
||||
# MuseTalk 集成由现有 GPU worker 处理
|
||||
# 这里调用现有接口
|
||||
try:
|
||||
# 如果不需要数字人,直接跳过
|
||||
if not job.persona_id:
|
||||
return video_path
|
||||
|
||||
# 调用 GPU worker 的 MuseTalk 接口
|
||||
import requests
|
||||
|
||||
gpu_worker_url = os.environ.get("GPU_WORKER_URL", "http://localhost:8900")
|
||||
resp = requests.post(
|
||||
f"{gpu_worker_url}/api/v1/gpu/lipsync",
|
||||
json={
|
||||
"video_path": video_path,
|
||||
"audio_path": job.reference_audio_path,
|
||||
"persona_id": job.persona_id,
|
||||
},
|
||||
timeout=300,
|
||||
logger.info("[爆款视频] 分镜 %d/%d dur=%ds prompt=%s", i + 1, n_total, dur, prompt[:80])
|
||||
seg_path = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image if i == 0 else None,
|
||||
duration=dur,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmpdir),
|
||||
)
|
||||
if resp.ok:
|
||||
result = resp.json()
|
||||
return result.get("output_path", video_path)
|
||||
return video_path
|
||||
if not seg_path or not Path(seg_path).exists():
|
||||
logger.warning("[爆款视频] 分镜 %d 生成失败,使用占位片段", i + 1)
|
||||
seg_path = str(_make_placeholder_clip(tmpdir, i, dur))
|
||||
seg_paths.append(seg_path)
|
||||
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 86.0, "正在拼接分镜...")
|
||||
concat_out = tmpdir / "concat_raw.mp4"
|
||||
try:
|
||||
concat_video_files(seg_paths, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] MuseTalk 处理失败,使用原始视频: %s", e)
|
||||
return video_path
|
||||
logger.error("[爆款视频] concat 失败: %s,降级过滤无效片段", e, exc_info=True)
|
||||
valid = [p for p in seg_paths if _probe_ok(p)]
|
||||
if not valid:
|
||||
raise RuntimeError(f"所有分镜片段均无效: {e}") from e
|
||||
concat_video_files(valid, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
|
||||
final_path = concat_out
|
||||
|
||||
if tts_path is not None:
|
||||
tts_p = Path(tts_path) if not isinstance(tts_path, Path) else tts_path
|
||||
if tts_p.exists():
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 87.5, "正在合成配音...")
|
||||
mixed_out = tmpdir / "final_with_audio.mp4"
|
||||
try:
|
||||
run_ffmpeg(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(concat_out),
|
||||
"-i",
|
||||
str(tts_p),
|
||||
"-c:v",
|
||||
"copy",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"192k",
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-shortest",
|
||||
str(mixed_out),
|
||||
]
|
||||
)
|
||||
if mixed_out.exists() and mixed_out.stat().st_size > 0:
|
||||
final_path = mixed_out
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 混音失败,使用无声视频: %s", e)
|
||||
|
||||
logger.info("[爆款视频] 渲染完成: %s size=%d", final_path, final_path.stat().st_size if final_path.exists() else 0)
|
||||
return str(final_path)
|
||||
|
||||
|
||||
def _probe_ok(video_path: str) -> bool:
|
||||
import subprocess
|
||||
from pathlib import Path as _Path
|
||||
|
||||
try:
|
||||
if not _Path(video_path).exists():
|
||||
return False
|
||||
r = subprocess.run(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=codec_type",
|
||||
"-of",
|
||||
"csv=p=0",
|
||||
video_path,
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=10,
|
||||
)
|
||||
return r.returncode == 0 and b"video" in r.stdout
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _make_placeholder_clip(tmpdir, idx: int, duration: int):
|
||||
import subprocess
|
||||
|
||||
out = tmpdir / f"placeholder_{idx}.mp4"
|
||||
try:
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"color=c=0x202030:s=720x1280:d={max(duration,2)}:r=24",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"anullsrc=r=44100:cl=stereo:d={max(duration,2)}",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-preset",
|
||||
"ultrafast",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(out),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=60,
|
||||
check=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 占位片段生成失败: %s", e)
|
||||
return out
|
||||
|
||||
|
||||
def _step_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 10: OSS 上传 + 扣点。"""
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
"""步骤 9: OSS 上传。"""
|
||||
from pathlib import Path
|
||||
|
||||
video_url = upload_to_oss(video_path, prefix="viral-video/")
|
||||
return video_url
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] OSS 上传失败: %s", e)
|
||||
raise
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
local = Path(video_path)
|
||||
# 构造 OSS key,与 generation.py 规则对齐:generated/viral-video/<user_id>/<job_id>/<filename>
|
||||
storage_key = f"generated/viral-video/{job.user_id}/{job.id}/{local.name}"
|
||||
logger.info("[爆款视频] 开始上传成片: local=%s key=%s size=%d", local, storage_key, local.stat().st_size)
|
||||
video_url = upload_to_oss(local, storage_key)
|
||||
if not video_url:
|
||||
raise RuntimeError(f"OSS 上传失败: storage_key={storage_key}")
|
||||
return video_url
|
||||
|
||||
|
||||
# ── 主编排器 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=2, name="worker.run_viral_video_pipeline")
|
||||
@shared_task(bind=True, max_retries=2, name="worker.run_viral_video_pipeline")
|
||||
def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""爆款视频 10 步流水线编排器。"""
|
||||
"""爆款视频 10 步流水线编排器(前半段:图片分析→风格分析→意图解析,然后 WAIT_USER_CONFIRM)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
@@ -389,7 +639,6 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
logger.error("[爆款视频] 任务不存在: %s", job_id)
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
# 标记运行中
|
||||
job.mark_running()
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 5.0, "开始图片分析")
|
||||
@@ -397,9 +646,13 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
# ── Step 1: 图片 VLM 分析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 10.0, "正在分析产品图片...")
|
||||
image_analysis = _step_image_analysis(job)
|
||||
# P0-3: 持久化 image_analysis 到 job,供 resume 阶段使用
|
||||
job.image_analysis = image_analysis
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 15.0, "图片分析完成", {"result": image_analysis})
|
||||
|
||||
# ── Step 1.5: 视频风格分析(v1.3) ──
|
||||
style_guide = None
|
||||
if job.reference_video_url or job.style_template_id:
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 20.0, "正在分析参考视频风格...")
|
||||
style_guide = _step_video_analysis(job)
|
||||
@@ -412,14 +665,11 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"风格分析完成",
|
||||
{"style_analyzed": True, "style_guide": style_guide},
|
||||
)
|
||||
else:
|
||||
style_guide = None
|
||||
|
||||
# ── Step 2: 意图解析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 30.0, "正在解析文案意图...")
|
||||
intent_result = _step_intent_parsing(job, image_analysis)
|
||||
|
||||
# 进入等待用户确认状态
|
||||
job.mark_wait_user_confirm(intent_result)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(
|
||||
@@ -429,33 +679,55 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"意图解析完成,等待用户确认",
|
||||
{"intent_result": intent_result, "waiting_confirm": True},
|
||||
)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.INTENT_PARSING,
|
||||
35.0,
|
||||
"等待用户确认意图文案",
|
||||
{"intent_result": intent_result},
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
|
||||
# 这里流水线暂停,等待 confirm-intent API 调用 resume
|
||||
# resume 后由 resume_viral_video_pipeline 继续
|
||||
return {"ok": True, "job_id": job_id, "status": "wait_user_confirm", "intent_result": intent_result}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 流水线异常: %s", e, exc_info=True)
|
||||
if session:
|
||||
try:
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
else:
|
||||
_, repo, job = _get_repo_and_job(job_id)
|
||||
if job and not job.is_terminal:
|
||||
job.mark_failed(str(e))
|
||||
_save_job(repo, job, session)
|
||||
except Exception:
|
||||
pass
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
if job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=2, name="worker.resume_viral_video_pipeline")
|
||||
@shared_task(bind=True, max_retries=2, name="worker.resume_viral_video_pipeline")
|
||||
def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""用户确认意图后,从断点恢复流水线(步骤 3-10)。"""
|
||||
session = None
|
||||
job = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
@@ -464,59 +736,62 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
if job.status != ViralVideoStatus.RUNNING:
|
||||
return {"ok": False, "error": f"unexpected status: {job.status}"}
|
||||
|
||||
# P0-3: 从 job 读取 image_analysis(run_pipeline 阶段已持久化)
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...")
|
||||
|
||||
# ── Step 3: 文案融合 ──
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, {"products": []})
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成")
|
||||
|
||||
# ── Step 4: 分镜脚本 ──
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...")
|
||||
storyboard = _step_storyboard(job, copy_text, {})
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成")
|
||||
storyboard = _step_storyboard(job, copy_text, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)})
|
||||
|
||||
# ── Step 5: 合规审核 ──
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...")
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
if not review_result.get("passed", True):
|
||||
# 自动重写 1 次
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, {"products": []})
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
|
||||
|
||||
# ── Step 6: CosyVoice 配音 ──
|
||||
# ── Step 6: CosyVoice 配音(返回 Path | None) ──
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...")
|
||||
audio_url = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成")
|
||||
tts_path = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None})
|
||||
|
||||
# ── Step 7: BGM 选择 ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "正在选择BGM...")
|
||||
# ── Step 7: BGM 选择(P1:暂返回 None,跳过) ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)")
|
||||
bgm = _step_bgm_select(job)
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 78.0, "BGM选择完成")
|
||||
|
||||
# ── Step 8: 渲染 ──
|
||||
# ── Step 8: 渲染(逐分镜 Seedance → concat → 混 TTS) ──
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...")
|
||||
video_path = _step_render(job, storyboard, audio_url, bgm)
|
||||
video_path = _step_render(job, storyboard, tts_path, bgm)
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成")
|
||||
|
||||
# ── Step 9: MuseTalk 数字人口型 ──
|
||||
_emit_progress(job_id, ViralVideoStage.MUSETALK, 90.0, "正在处理数字人口型...")
|
||||
final_video_path = _step_musetalk(job, video_path)
|
||||
_emit_progress(job_id, ViralVideoStage.MUSETALK, 93.0, "数字人处理完成")
|
||||
|
||||
# ── Step 10: OSS 上传 + 扣点 ──
|
||||
# ── Step 9: OSS 上传 + 扣点 ──
|
||||
# 注:爆款视频由 Seedance 2.5 直接生成口型,不需要 MuseTalk 事后对口型(MuseTalk 是 AI 数字人路线用的)。
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...")
|
||||
video_url = _step_upload(job, final_video_path)
|
||||
video_url = _step_upload(job, video_path)
|
||||
|
||||
# 扣点
|
||||
job.credits_cost = CREDITS_VIRAL_VIDEO_COST
|
||||
# TODO: 调用 credits.deduct() 实际扣点
|
||||
# TODO: 调用 credits.deduct() 实际扣点(#1895 总开关为 false 时不扣,保留 TODO)
|
||||
|
||||
# 完成
|
||||
job.mark_completed(video_url)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url})
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.UPLOADING,
|
||||
100.0,
|
||||
"视频生成完成",
|
||||
{"video_url": video_url},
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
|
||||
logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url)
|
||||
return {"ok": True, "job_id": job_id, "video_url": video_url}
|
||||
@@ -525,21 +800,34 @@ def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线异常: %s", e, exc_info=True)
|
||||
if session:
|
||||
try:
|
||||
_, repo, job = _get_repo_and_job(job_id)
|
||||
if job and not job.is_terminal:
|
||||
job.mark_failed(str(e))
|
||||
_save_job(repo, job, session)
|
||||
except Exception:
|
||||
pass
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
elif job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=1, name="worker.run_video_style_analysis")
|
||||
@shared_task(bind=True, max_retries=1, name="worker.run_video_style_analysis")
|
||||
def run_video_style_analysis(self: Task, job_id: str) -> dict:
|
||||
"""独立的视频风格分析任务(v1.3)。"""
|
||||
session = None
|
||||
|
||||
@@ -502,8 +502,19 @@ class SQLAlchemyAssetRepository:
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key(对应 DB 中的 file_url)查找素材。"""
|
||||
model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first()
|
||||
"""按 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()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
@@ -946,6 +946,7 @@ class ViralVideoJobModel(Base):
|
||||
# 结果与状态
|
||||
status = Column(String(30), nullable=False, default="pending", index=True)
|
||||
intent_result = Column(JSON, nullable=True)
|
||||
image_analysis = Column(JSON, nullable=True)
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
|
||||
@@ -34,6 +34,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
style_template_id=getattr(model, "style_template_id", "") or "",
|
||||
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
|
||||
intent_result=dict(model.intent_result) if model.intent_result else None,
|
||||
image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
credits_cost=model.credits_cost or 0,
|
||||
error_msg=model.error_msg or "",
|
||||
@@ -72,6 +73,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
style_template_id=job.style_template_id,
|
||||
status=job.status,
|
||||
intent_result=job.intent_result,
|
||||
image_analysis=job.image_analysis,
|
||||
result_video_url=job.result_video_url,
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
@@ -91,6 +93,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
raise ValueError(f"ViralVideoJob {job.id} not found")
|
||||
model.status = job.status
|
||||
model.intent_result = job.intent_result
|
||||
model.image_analysis = job.image_analysis
|
||||
model.result_video_url = job.result_video_url
|
||||
model.credits_cost = job.credits_cost
|
||||
model.error_msg = job.error_msg
|
||||
|
||||
@@ -96,6 +96,9 @@ class SharedSettings(BaseSettings):
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
|
||||
@@ -118,6 +118,8 @@ class ViralVideoJob:
|
||||
style_strength: str = StyleStrength.MEDIUM
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
# v1.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜)
|
||||
image_analysis: dict | None = None
|
||||
# 状态
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
status: ViralVideoStatus = ViralVideoStatus.PENDING
|
||||
|
||||
@@ -13,7 +13,9 @@ API 和 Worker 两边共用。基于火山引擎方舟平台的 OpenAI 兼容接
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
@@ -238,6 +240,188 @@ 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 = False,
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 文生/图生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
Args:
|
||||
prompt: 文本提示词
|
||||
image_url: 首帧参考图 URL(可选,提供则走图生视频)
|
||||
duration: 视频时长 2~30 秒,默认 5
|
||||
ratio: 宽高比 16:9/9:16/1:1/4:3/3:4/21:9/adaptive
|
||||
resolution: 480p/720p/1080p
|
||||
generate_audio: 是否生成模型自带音效(默认 False,我们自己混 TTS)
|
||||
watermark: 是否加水印
|
||||
output_dir: 下载目录,默认 /tmp
|
||||
|
||||
Returns:
|
||||
本地 MP4 文件路径,失败返回 None。
|
||||
"""
|
||||
if not self.is_available:
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||||
total_timeout = getattr(settings, "doubao_video_timeout", 600) or 600
|
||||
video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
|
||||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
|
||||
if image_url:
|
||||
content.append({"type": "image_url", "image_url": {"url": image_url}})
|
||||
|
||||
create_payload: dict[str, Any] = {
|
||||
"model": video_model,
|
||||
"content": content,
|
||||
"generate_audio": generate_audio,
|
||||
"duration": int(duration),
|
||||
"resolution": resolution,
|
||||
"watermark": watermark,
|
||||
}
|
||||
# Bug #2110: ratio=None 时不传(首帧图生视频跟随原图比例,传 ratio 会 400 InvalidParameter)
|
||||
if ratio:
|
||||
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",
|
||||
create_url,
|
||||
video_model,
|
||||
duration,
|
||||
ratio or "(follow-image)",
|
||||
generate_audio,
|
||||
bool(image_url),
|
||||
)
|
||||
|
||||
# 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", task_id, video_model, duration)
|
||||
|
||||
# 2) 轮询状态
|
||||
poll_url = f"{create_url}/{task_id}"
|
||||
deadline = time.time() + total_timeout
|
||||
video_url: str | None = None
|
||||
last_status: str = "queued"
|
||||
while time.time() < deadline:
|
||||
try:
|
||||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||||
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
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -536,3 +536,42 @@ def call_vision(image_url: str, prompt: str) -> object:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return raw
|
||||
|
||||
|
||||
def call_video_generation(
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str | None = "9:16",
|
||||
resolution: str = "720p",
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生成视频段,返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
封装 ai_client.video_generation:提交异步任务→轮询→下载到本地。
|
||||
Bug #2110: 首帧参考图模式下不传 ratio(API 要求跟随首帧图比例,传 ratio=9:16
|
||||
会返回 400 InvalidParameter)。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
# 首帧模式:不强制 ratio,让模型跟随首帧图比例
|
||||
effective_ratio = None if image_url else ratio
|
||||
try:
|
||||
kwargs: dict = dict(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
generate_audio=False, # 我们自己混 TTS
|
||||
watermark=False,
|
||||
output_dir=output_dir,
|
||||
)
|
||||
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
|
||||
|
||||
@@ -15,3 +15,8 @@ Pillow==10.4.0
|
||||
|
||||
# FFmpeg Python 绑定
|
||||
ffmpeg-python==0.2.0
|
||||
|
||||
# v1.3 参考视频风格分析(#2051)
|
||||
scenedetect==0.6.4
|
||||
librosa==0.10.2.post1
|
||||
soundfile==0.12.1
|
||||
|
||||
@@ -137,6 +137,28 @@ 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() {
|
||||
|
||||
@@ -0,0 +1,485 @@
|
||||
"""#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
|
||||
@@ -144,6 +144,9 @@ 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
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,360 @@
|
||||
"""video_analyzer(#2051)单元测试。
|
||||
|
||||
覆盖:
|
||||
- 映射函数(运镜→ken_burns、转场→xfade、色调→video_filter、BPM→BGM)
|
||||
- schema 常量与导出
|
||||
- Farneback 光流运镜判定(合成光流场)
|
||||
- BPM 档位映射
|
||||
- 降级路径(ffmpeg/cv2/librosa 不可用)
|
||||
- 临时目录清理
|
||||
- analyze_video_style 入口在无素材时返回最小 style_guide 不抛
|
||||
- build_render_params_for_clip 聚合输出
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import ( # noqa: E402
|
||||
COLOR_FILTER_PRESETS,
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
TRANSITION_TO_XFADE,
|
||||
ShotBoundary,
|
||||
_detect_camera_movement,
|
||||
_pace_from_bpm,
|
||||
_rule_based_style_guide,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
# ── 映射函数 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_constants_exported():
|
||||
assert MAX_REFERENCE_DURATION_SEC == 60
|
||||
assert MAX_REFERENCE_SIZE_MB == 100
|
||||
assert DEFAULT_ANALYSIS_TIMEOUT == 60
|
||||
assert "style_name" in STYLE_GUIDE_SCHEMA
|
||||
assert "ken_burns_params" in STYLE_GUIDE_SCHEMA
|
||||
assert "video_filter_eq_params" in STYLE_GUIDE_SCHEMA
|
||||
|
||||
|
||||
# ── 运镜→ken_burns ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"movement,expected_type",
|
||||
[
|
||||
("static", "static"),
|
||||
("push_in", "zoom"),
|
||||
("zoom_in", "zoom"),
|
||||
("pull_out", "zoom"),
|
||||
("pan_left", "pan"),
|
||||
("pan_right", "pan"),
|
||||
("tilt_up", "pan+zoom"),
|
||||
("track_left", "pan"),
|
||||
],
|
||||
)
|
||||
def test_map_camera_to_ken_burns_types(movement, expected_type):
|
||||
kb = map_camera_to_ken_burns(movement)
|
||||
assert kb["type"] == expected_type
|
||||
# zoom/pan 类必须有 zoom_start/zoom_end
|
||||
assert 0.8 <= kb["zoom_start"] <= 1.3
|
||||
assert 0.8 <= kb["zoom_end"] <= 1.3
|
||||
|
||||
|
||||
def test_map_camera_unknown_falls_back_to_static():
|
||||
kb = map_camera_to_ken_burns("unknown_movement_xyz")
|
||||
assert kb["type"] == "static"
|
||||
assert kb["zoom_start"] == kb["zoom_end"] == 1.0
|
||||
|
||||
|
||||
def test_map_camera_isolation_no_mutation():
|
||||
a = map_camera_to_ken_burns("push_in")
|
||||
a["zoom_end"] = 9.99
|
||||
b = map_camera_to_ken_burns("push_in")
|
||||
assert b["zoom_end"] != 9.99
|
||||
|
||||
|
||||
# ── 转场→xfade ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ttype,expected",
|
||||
[
|
||||
("hard_cut", "cut"),
|
||||
("cross_dissolve", "dissolve"),
|
||||
("fade", "fade"),
|
||||
("fade_black", "fadeblack"),
|
||||
("zoom_whip", "zoom"),
|
||||
("wipe_left", "wipeleft"),
|
||||
("slide_right", "slideleft"),
|
||||
],
|
||||
)
|
||||
def test_map_transition(ttype, expected):
|
||||
assert map_transition_to_xfade(ttype) == expected
|
||||
|
||||
|
||||
def test_map_transition_unknown_falls_back_to_cut():
|
||||
assert map_transition_to_xfade("some_random_transition") == "cut"
|
||||
|
||||
|
||||
# ── 色调→video_filter ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name", ["none", "warm_vintage", "cool_fresh", "high_contrast", "soft_pastel", "dramatic_cinematic"]
|
||||
)
|
||||
def test_map_color_presets_available(name):
|
||||
p = map_color_to_video_filter(name)
|
||||
assert isinstance(p, dict)
|
||||
# 所有预设必须能被 FFmpeg eq/colorchannelmixer 消费:eq 是 dict,ccm 是 dict
|
||||
assert "eq" in p or p == {} or "colorchannelmixer" in p
|
||||
|
||||
|
||||
def test_map_color_unknown_is_none_preset():
|
||||
p = map_color_to_video_filter("not_a_real_filter")
|
||||
assert p == {}
|
||||
|
||||
|
||||
def test_map_color_isolation():
|
||||
a = map_color_to_video_filter("warm_vintage")
|
||||
a["eq"]["brightness"] = 9.99
|
||||
b = map_color_to_video_filter("warm_vintage")
|
||||
assert b["eq"]["brightness"] != 9.99
|
||||
|
||||
|
||||
# ── BPM → BGM ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bpm,expected",
|
||||
[
|
||||
(0, 90),
|
||||
(70, 70),
|
||||
(120, 120),
|
||||
(200, 180),
|
||||
(30, 60),
|
||||
],
|
||||
)
|
||||
def test_map_bgm_bpm(bpm, expected):
|
||||
assert map_bgm_bpm(bpm) == expected
|
||||
|
||||
|
||||
def test_pace_from_bpm_buckets():
|
||||
assert _pace_from_bpm(120) == "fast_cut"
|
||||
assert _pace_from_bpm(110) == "fast_cut"
|
||||
assert _pace_from_bpm(90) == "medium"
|
||||
assert _pace_from_bpm(80) == "medium"
|
||||
assert _pace_from_bpm(60) == "slow_cinematic"
|
||||
assert _pace_from_bpm(0) == "medium"
|
||||
|
||||
|
||||
# ── Farneback 光流→运镜(合成光流) ───────────────────────────────────
|
||||
|
||||
|
||||
def _make_flow(dx: float, dy: float, w: int = 60, h: int = 40, zoom: float = 0.0):
|
||||
"""构造一个合成光流场:整体平移(dx,dy)+径向发散(zoom>0=zoom in,<0=out)。"""
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
fx = dx + (xs - cx) * zoom
|
||||
fy = dy + (ys - cy) * zoom
|
||||
return np.stack([fx, fy], axis=-1).astype(np.float32)
|
||||
|
||||
|
||||
def test_detect_movement_static():
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "static"
|
||||
assert i == "low"
|
||||
|
||||
|
||||
def test_detect_movement_pan_right():
|
||||
flow = _make_flow(2.0, 0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_right"
|
||||
assert i in ("medium", "high")
|
||||
|
||||
|
||||
def test_detect_movement_pan_left():
|
||||
flow = _make_flow(-2.0, 0.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_left"
|
||||
|
||||
|
||||
def test_detect_movement_tilt_down():
|
||||
flow = _make_flow(0.0, 2.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "tilt_down"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_in_radial():
|
||||
# 径向向外发散 = zoom in
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_in"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_out_radial():
|
||||
flow = _make_flow(0.0, 0.0, zoom=-0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_out"
|
||||
|
||||
|
||||
# ── 规则合成 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_shots(n=4):
|
||||
return [
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=float(i * 3),
|
||||
end_sec=float((i + 1) * 3),
|
||||
movement=["static", "push_in", "pan_left", "zoom_in"][i],
|
||||
intensity=["low", "medium", "low", "high"][i],
|
||||
transition="hard_cut",
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
|
||||
CAMERA_TO_KEN_BURNS_DIRS = {
|
||||
"zoom_in_slow",
|
||||
"zoom_out_slow",
|
||||
"pan_left_slow",
|
||||
"pan_right_slow",
|
||||
"zoom_in_medium",
|
||||
"zoom_out_medium",
|
||||
"diagonal_push",
|
||||
"static",
|
||||
}
|
||||
|
||||
|
||||
def test_rule_based_style_guide_structure():
|
||||
shots = _sample_shots()
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={"color_filter": "warm_vintage"})
|
||||
# 关键字段存在且类型正确
|
||||
assert sg["shot_count"] == 4
|
||||
assert sg["pace"] == "fast_cut"
|
||||
assert sg["bpm"] == 120
|
||||
assert sg["avg_shot_duration"] == 3.0
|
||||
assert len(sg["camera_movements"]) == 4
|
||||
assert sg["color_filter"] == "warm_vintage"
|
||||
assert "eq" in sg["video_filter_eq_params"]
|
||||
assert "default" in sg["ken_burns_params"]
|
||||
assert isinstance(sg["transition_map"], dict)
|
||||
assert isinstance(sg["ken_burns_direction_hint"], str) and sg["ken_burns_direction_hint"]
|
||||
|
||||
|
||||
# ── build_render_params_for_clip 聚合 ─────────────────────────────────
|
||||
|
||||
|
||||
def test_build_render_params_for_clip_shape():
|
||||
sg = _rule_based_style_guide(_sample_shots(), bpm=95, vlm={"color_filter": "cool_fresh"})
|
||||
p0 = build_render_params_for_clip(0, sg, duration_sec=3.0)
|
||||
assert "ken_burns" in p0
|
||||
assert "transition" in p0
|
||||
assert "video_filter" in p0
|
||||
assert p0["bgm_bpm_hint"] == 95
|
||||
assert p0["duration_sec"] == 3.0
|
||||
# clip 1 是 push_in → zoom
|
||||
p1 = build_render_params_for_clip(1, sg)
|
||||
assert p1["ken_burns"]["type"] == "zoom"
|
||||
|
||||
|
||||
def test_build_render_params_high_intensity_amplifies():
|
||||
shots = _sample_shots() # shot 3 = zoom_in/high
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={})
|
||||
p3 = build_render_params_for_clip(3, sg)
|
||||
base = map_camera_to_ken_burns("zoom_in")
|
||||
assert p3["ken_burns"]["zoom_end"] > base["zoom_end"]
|
||||
|
||||
|
||||
# ── 降级与容错 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_analyze_with_nonexistent_file_returns_minimum_guide():
|
||||
sg = analyze_video_style("/nonexistent/path/fake_video.mp4")
|
||||
assert isinstance(sg, dict)
|
||||
assert "style_name" in sg
|
||||
assert sg["shot_count"] == 0
|
||||
# 不抛异常且字段完整
|
||||
|
||||
|
||||
def test_analyze_invalid_style_strength_defaults_to_medium():
|
||||
# 即使视频不存在,也应被规范化为 medium 并写入返回值
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("fake", style_strength="banana")
|
||||
assert sg.get("style_strength", "medium") == "medium"
|
||||
|
||||
|
||||
def test_temp_dir_cleaned_up_after_run():
|
||||
"""用临时真实空文件模拟本地路径,确认 frames 临时目录被清理。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "fake.mp4"
|
||||
fake.write_bytes(b"")
|
||||
# 抽帧会失败(ffmpeg 对空文件失败),但应全程不抛且临时目录 rmtree
|
||||
# 直接 mock _ensure_local_video 回传不存在的文件,走 _probe_duration=0 降级路径
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("proto://fake", style_strength="light")
|
||||
assert "style_name" in sg
|
||||
|
||||
|
||||
def test_ffmpeg_failure_falls_back_to_vlm_only_path():
|
||||
"""模拟 ffmpeg 抽帧失败,仍能返回 style_guide。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "ref.mp4"
|
||||
fake.write_bytes(b"not a real video")
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._extract_keyframes", side_effect=RuntimeError("ffmpeg exploded")
|
||||
):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_shots") as mock_shots:
|
||||
mock_shots.return_value = [ShotBoundary(0, 0.0, 3.0)]
|
||||
with patch("apps.worker.viral_video.video_analyzer._analyze_movements"):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_bpm", return_value=90):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._vlm_analyze_frames",
|
||||
return_value={"color_filter": "none"},
|
||||
):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._llm_synthesize",
|
||||
side_effect=lambda shots, bpm, vlm, ss: _rule_based_style_guide(shots, bpm, vlm),
|
||||
):
|
||||
sg = analyze_video_style(str(fake))
|
||||
assert sg["bpm"] == 90
|
||||
assert sg["shot_count"] == 1
|
||||
|
||||
|
||||
# ── 转场映射完整性 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_transition_map_covers_observed_types():
|
||||
for t in ("hard_cut", "cross_dissolve", "fade_black", "zoom_whip"):
|
||||
assert t in TRANSITION_TO_XFADE
|
||||
|
||||
|
||||
# ── 预设完整性 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_color_filter_preset_keys_are_safe_for_ffmpeg():
|
||||
for name, preset in COLOR_FILTER_PRESETS.items():
|
||||
if preset == {}:
|
||||
continue
|
||||
# eq 所有值都是数字
|
||||
for k, v in preset.get("eq", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.eq.{k} not numeric"
|
||||
for k, v in preset.get("colorchannelmixer", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.ccm.{k} not numeric"
|
||||
@@ -436,16 +436,17 @@ class TestViralVideoPipeline:
|
||||
def test_bgm_select(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
# P1: BGM 素材未就绪前 _step_bgm_select 统一返回 None(跳过 BGM 混音)
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
bgm = _step_bgm_select(mock_job)
|
||||
assert "upbeat" in bgm
|
||||
assert bgm is None
|
||||
|
||||
def test_bgm_select_default(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
mock_job.bgm_preference = ""
|
||||
bgm = _step_bgm_select(mock_job)
|
||||
assert bgm == "bgm_default.mp3"
|
||||
assert bgm is None
|
||||
|
||||
|
||||
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
|
||||
@@ -455,7 +456,6 @@ class TestPipelineIntegration:
|
||||
"""流水线端到端集成测试(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._step_bgm_select")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
|
||||
@@ -480,7 +480,6 @@ class TestPipelineIntegration:
|
||||
mock_tts,
|
||||
mock_bgm,
|
||||
mock_render,
|
||||
mock_musetalk,
|
||||
mock_upload,
|
||||
):
|
||||
"""测试 resume 流水线能从确认状态走到完成。"""
|
||||
@@ -505,10 +504,9 @@ class TestPipelineIntegration:
|
||||
mock_copy_fusion.return_value = "融合文案"
|
||||
mock_storyboard.return_value = [{"order": 0, "duration": 10}]
|
||||
mock_review.return_value = {"passed": True, "score": 90}
|
||||
mock_tts.return_value = "https://audio.mp3"
|
||||
mock_bgm.return_value = "bgm_default.mp3"
|
||||
mock_tts.return_value = None # P1: TTS 返回 Path|None,mock 用 None 跳过混音
|
||||
mock_bgm.return_value = None # P1: BGM 未就绪前返回 None
|
||||
mock_render.return_value = "/tmp/video.mp4"
|
||||
mock_musetalk.return_value = "/tmp/video_final.mp4"
|
||||
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
||||
|
||||
result = resume_viral_video_pipeline.run("job-001")
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
"""#2106 P0 修复单测:Seedance 对接、image_analysis 持久化、TTS Path 统一、BGM/MuseTalk 跳过。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# worker 容器 PYTHONPATH 包含 apps/worker(worker 侧代码使用顶层包名 services/、viral_video/)
|
||||
_WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
|
||||
if str(_WORKER_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_WORKER_ROOT))
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_job():
|
||||
return ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=15,
|
||||
user_copy_text="测试文案",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
|
||||
|
||||
# ── P0-2: _step_video_analysis import 路径 ──────────────────────────
|
||||
|
||||
|
||||
class TestVideoAnalysisImport:
|
||||
def test_no_reference_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = ""
|
||||
assert _step_video_analysis(mock_job) is None
|
||||
|
||||
def test_with_reference_returns_dict_or_none(self, mock_job):
|
||||
"""有参考视频 URL 时,不管分析成功/失败/占位,返回 dict(不抛异常)。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = "https://example.com/ref.mp4"
|
||||
result = _step_video_analysis(mock_job)
|
||||
# 允许占位/失败/真实返回,但绝不能抛异常
|
||||
assert result is None or isinstance(result, dict)
|
||||
|
||||
|
||||
# ── P0-3: image_analysis 字段 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestImageAnalysisField:
|
||||
def test_default_none(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
assert job.image_analysis is None
|
||||
|
||||
def test_persist_and_read(self, mock_job):
|
||||
mock_job.image_analysis = {"products": [{"name": "口红"}]}
|
||||
assert mock_job.image_analysis["products"][0]["name"] == "口红"
|
||||
|
||||
|
||||
# ── P0-1: storyboard 规范化 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStoryboardNormalize:
|
||||
def test_normalize_fills_defaults(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
raw = [{"order": 0, "description": "镜头一"}]
|
||||
out = _normalize_storyboard(raw, total_duration=10, n_segments=1, copy_text="文案")
|
||||
assert len(out) == 1
|
||||
assert out[0]["duration"] >= 3
|
||||
assert out[0]["ken_burns"] in {"zoom_in", "zoom_out", "pan_left", "pan_right", "static"}
|
||||
assert out[0]["type"] == "product_shot"
|
||||
assert out[0]["text"] == ""
|
||||
|
||||
def test_normalize_scales_to_total_duration(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
raw = [
|
||||
{"order": 0, "duration": 10, "description": "a"},
|
||||
{"order": 1, "duration": 10, "description": "b"},
|
||||
]
|
||||
out = _normalize_storyboard(raw, total_duration=10, n_segments=2, copy_text="x")
|
||||
total = sum(s["duration"] for s in out)
|
||||
assert total == 10
|
||||
|
||||
def test_fallback_storyboard(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _fallback_storyboard
|
||||
|
||||
out = _fallback_storyboard("文案", total_duration=15, n_segments=3)
|
||||
assert len(out) == 3
|
||||
assert sum(s["duration"] for s in out) == 15
|
||||
assert all(s["duration"] >= 3 for s in out)
|
||||
|
||||
def test_storyboard_llm_list(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
|
||||
|
||||
with patch("packages.shared.ai_service.call_llm") as mock_llm:
|
||||
mock_llm.return_value = [
|
||||
{"order": 0, "description": "产品特写", "duration": 5, "text": "t1"},
|
||||
{"order": 1, "description": "使用场景", "duration": 5, "text": "t2"},
|
||||
{"order": 2, "description": "CTA", "duration": 5, "text": "t3"},
|
||||
]
|
||||
result = _step_storyboard(mock_job, "文案", {"products": []})
|
||||
assert len(result) == 3
|
||||
assert all("description" in s for s in result)
|
||||
|
||||
|
||||
# ── P1: TTS 返回 Path|None ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTTSPath:
|
||||
def test_tts_returns_none_on_import_error(self, mock_job):
|
||||
"""get_tts_service 抛 ImportError 时 _step_tts 返回 None。"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
with patch("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 TestBGMSkip:
|
||||
def test_bgm_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
assert _step_bgm_select(mock_job) is None
|
||||
|
||||
|
||||
# ── P0-1: call_video_generation 参数构造 ──────────────────────────
|
||||
|
||||
|
||||
class TestCallVideoGeneration:
|
||||
def test_returns_none_when_client_unavailable(self):
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get.return_value = mock_client
|
||||
assert call_video_generation("prompt") is None
|
||||
|
||||
def test_delegates_to_client(self, tmp_path):
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
out = tmp_path / "v.mp4"
|
||||
out.write_bytes(b"fake")
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.return_value = str(out)
|
||||
mock_get.return_value = mock_client
|
||||
result = call_video_generation(prompt="测试", image_url="https://img/x.jpg", duration=5, ratio="9:16")
|
||||
assert result == str(out)
|
||||
mock_client.video_generation.assert_called_once()
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
assert kwargs["prompt"] == "测试"
|
||||
assert kwargs["image_url"] == "https://img/x.jpg"
|
||||
assert kwargs["duration"] == 5
|
||||
|
||||
|
||||
# ── P0-1: _step_render 占位片段生成 ──────────────────────────────
|
||||
|
||||
|
||||
class TestPlaceholderClip:
|
||||
def test_make_placeholder_clip(self, tmp_path):
|
||||
import shutil
|
||||
|
||||
from apps.worker.worker_app.tasks.viral_video import _make_placeholder_clip, _probe_ok
|
||||
|
||||
if not shutil.which("ffmpeg"):
|
||||
pytest.skip("ffmpeg not available")
|
||||
|
||||
out = _make_placeholder_clip(tmp_path, 0, 3)
|
||||
assert out.exists()
|
||||
assert _probe_ok(str(out))
|
||||
|
||||
|
||||
# ── P0-1: DoubaoClient.video_generation 在不可用时返回 None ───────
|
||||
|
||||
|
||||
class TestDoubaoClientVideoGen:
|
||||
def test_unavailable_returns_none(self):
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = "" # is_available -> False
|
||||
assert client.video_generation("prompt") is None
|
||||
|
||||
|
||||
# ── P0-3: resume 从 job 读 image_analysis ────────────────────────
|
||||
|
||||
|
||||
class TestResumeReadsImageAnalysis:
|
||||
def test_resume_uses_persisted_image_analysis(self):
|
||||
"""resume_pipeline 应从 job.image_analysis 读(P0-3 持久化)。"""
|
||||
import inspect
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "job.image_analysis" in src
|
||||
@@ -0,0 +1,176 @@
|
||||
"""viral_video.py HTTP 端点单元测试(celery send_task 分支覆盖)。
|
||||
|
||||
直接调用路由函数(不启动 TestClient),通过 patch 注入 repo/session/user,
|
||||
覆盖 4 个 celery_app.send_task(...) 调用点:
|
||||
|
||||
- create_viral_video (generate) -> worker.run_viral_video_pipeline
|
||||
- retry_viral_video_job (retry) -> worker.run_viral_video_pipeline
|
||||
- confirm_intent -> worker.resume_viral_video_pipeline
|
||||
- analyze_style -> worker.run_video_style_analysis
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _auth_user(uid: str = "u1"):
|
||||
return SimpleNamespace(user=SimpleNamespace(id=uid))
|
||||
|
||||
|
||||
def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending", **kwargs):
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
job = MagicMock()
|
||||
job.id = job_id
|
||||
job.user_id = user_id
|
||||
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||
job.images = kwargs.pop("images", ["img-1"])
|
||||
job.industry = kwargs.pop("industry", "电商")
|
||||
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
||||
for k, v in {
|
||||
"persona_id": "",
|
||||
"viral_structure": "",
|
||||
"marketing_purpose": "",
|
||||
"bgm_preference": "",
|
||||
"duration": 30,
|
||||
"user_copy_text": "",
|
||||
"fusion_level": "ai_polish",
|
||||
"reference_audio_path": "",
|
||||
"reference_video_url": "",
|
||||
"style_strength": "medium",
|
||||
"style_template_id": "",
|
||||
"retry_count": 0,
|
||||
"error_msg": "",
|
||||
"result_video_url": "",
|
||||
"style_guide": None,
|
||||
"created_at": None,
|
||||
"started_at": None,
|
||||
"completed_at": None,
|
||||
"stage": "",
|
||||
"progress": 0.0,
|
||||
"intent_result": None,
|
||||
"updated_at": None,
|
||||
}.items():
|
||||
setattr(job, k, kwargs.pop(k, v))
|
||||
return job
|
||||
|
||||
|
||||
# ── generate ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateViralVideo:
|
||||
def _req(self, **kw):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
d = {"images": ["https://x.com/a.jpg"], "industry": "电商", "target_customer": "年轻人"}
|
||||
d.update(kw)
|
||||
return CreateViralVideoRequest(**d)
|
||||
|
||||
def test_generate_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
req = self._req()
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
saved_job = _make_job(job_id="job-new", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
|
||||
def fake_save(job):
|
||||
job.id = saved_job.id
|
||||
|
||||
repo.save.side_effect = fake_save
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.create_viral_video(req, authenticated_user=user, session=session)
|
||||
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=[saved_job.id])
|
||||
assert resp.id == saved_job.id
|
||||
|
||||
|
||||
# ── retry ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRetryViralVideo:
|
||||
def test_retry_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-retry", user_id="u1", status=ViralVideoStatus.FAILED, retry_count=1)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.retry_viral_video_job("job-retry", authenticated_user=user, session=session)
|
||||
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.retry_count == 2
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=["job-retry"])
|
||||
assert resp.id == "job-retry"
|
||||
|
||||
|
||||
# ── confirm-intent ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConfirmIntent:
|
||||
def test_confirm_intent_dispatches_resume_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmIntentRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-cfm", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmIntentRequest(confirmed_copy="确认后的文案")
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.confirm_intent("job-cfm", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.user_copy_text == "确认后的文案"
|
||||
job.resume_from_confirm.assert_called_once()
|
||||
mock_send.assert_called_once_with("worker.resume_viral_video_pipeline", args=["job-cfm"])
|
||||
assert resp.id == "job-cfm"
|
||||
|
||||
|
||||
# ── analyze-style ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAnalyzeStyle:
|
||||
def test_analyze_style_dispatches_analysis_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import AnalyzeStyleRequest
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-sty", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = AnalyzeStyleRequest(reference_video_url="https://x.com/ref.mp4", style_template_id="tpl-1")
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.analyze_style("job-sty", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.reference_video_url == "https://x.com/ref.mp4"
|
||||
assert job.style_template_id == "tpl-1"
|
||||
mock_send.assert_called_once_with("worker.run_video_style_analysis", args=["job-sty"])
|
||||
assert resp.job_id == "job-sty"
|
||||
assert resp.status == "analyzing"
|
||||
@@ -0,0 +1,432 @@
|
||||
"""Unit tests for the viral_video WebSocket progress endpoint and worker event format.
|
||||
|
||||
These tests exercise:
|
||||
* the worker _emit_progress helper (JSON serialization + event_type kwarg)
|
||||
* the pure helper functions on the API route module
|
||||
* WebSocket authentication / ownership / 404 behaviour
|
||||
* Initial-snapshot / terminal-job fast-close behaviour of the WS endpoint
|
||||
|
||||
The CI unit-test environment sets ``USE_IN_MEMORY_DB=true`` and relies on
|
||||
``settings.effective_database_url`` returning a SQLite URL. ``app/db.py`` and
|
||||
``app/dependencies.py`` have been fixed to honour ``effective_database_url``
|
||||
(matching the worker), so these tests never need a real Postgres or Redis.
|
||||
|
||||
Imports go through the ``apps.worker.*`` namespace (not bare ``worker_app.*``)
|
||||
to stay consistent with the existing integration tests and avoid creating a
|
||||
second module object that would make cross-file patches invisible.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.websockets import WebSocketDisconnect
|
||||
|
||||
# Ensure CI-friendly env is set BEFORE any app import so SQLite is used.
|
||||
os.environ.setdefault("USE_IN_MEMORY_DB", "true")
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "test-secret")
|
||||
os.environ.setdefault("DATABASE_URL", "postgresql+psycopg://no:such@127.0.0.1:1/none")
|
||||
|
||||
import app.db as _app_db # noqa: E402
|
||||
from app.api.routes import viral_video as vv_module # noqa: E402
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as worker_vv # noqa: E402
|
||||
|
||||
|
||||
def _make_job(**kwargs):
|
||||
defaults = dict(
|
||||
id="job-1",
|
||||
user_id="user-1",
|
||||
status="running",
|
||||
current_stage="analyzing",
|
||||
progress_percent=30,
|
||||
status_message="looking good",
|
||||
error_msg=None,
|
||||
is_terminal=False,
|
||||
result_video_url=None,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Worker event serialization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWorkerEmitProgress:
|
||||
def test_emit_progress_serialises_with_json_dumps(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress("job-1", "analyzing", 12, message="hi")
|
||||
fake_r.publish.assert_called_once()
|
||||
channel, payload = fake_r.publish.call_args.args
|
||||
assert channel == "viral_video:job-1"
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["stage"] == "analyzing"
|
||||
assert parsed["type"] == "viral_video:progress"
|
||||
assert parsed["progress"] == 12
|
||||
assert parsed["job_id"] == "job-1"
|
||||
assert "'stage'" not in payload # JSON uses double quotes, not Python repr
|
||||
|
||||
def test_emit_progress_respects_event_type(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-2",
|
||||
"done",
|
||||
100,
|
||||
message="ok",
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:completed"
|
||||
assert parsed["progress"] == 100
|
||||
|
||||
def test_emit_progress_failure_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-3",
|
||||
"failed",
|
||||
0,
|
||||
message="err",
|
||||
data={"error": "oom"},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:failed"
|
||||
assert parsed["data"]["error"] == "oom"
|
||||
|
||||
def test_emit_progress_wait_user_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-4",
|
||||
"intent_parsing",
|
||||
35,
|
||||
message="waiting for you",
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:wait_user"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure helpers on the route module
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSHelpers:
|
||||
def test_estimate_progress_maps_status(self):
|
||||
assert vv_module._estimate_progress(_make_job(status="pending")) == 0.0
|
||||
assert vv_module._estimate_progress(_make_job(status="running")) == 5.0
|
||||
assert vv_module._estimate_progress(_make_job(status="wait_user_confirm")) == 35.0
|
||||
assert vv_module._estimate_progress(_make_job(status="completed")) == 100.0
|
||||
assert vv_module._estimate_progress(_make_job(status="failed")) == 0.0
|
||||
|
||||
def test_initial_message_readable(self):
|
||||
job = _make_job(status="running")
|
||||
msg = vv_module._initial_message(job)
|
||||
assert isinstance(msg, str) and msg
|
||||
job_failed = _make_job(status="failed", error_msg="boom")
|
||||
assert "boom" in vv_module._initial_message(job_failed)
|
||||
job_wait = _make_job(status="wait_user_confirm")
|
||||
assert "等待" in vv_module._initial_message(job_wait)
|
||||
|
||||
def test_stage_from_status_falls_back(self):
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="pending")), str)
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="weird_unknown")), str)
|
||||
|
||||
def test_job_status_handles_enum_and_string(self):
|
||||
job = _make_job(status="running")
|
||||
assert vv_module._job_status(job) == "running"
|
||||
job_enum = _make_job(status=SimpleNamespace(value="completed"))
|
||||
assert vv_module._job_status(job_enum) == "completed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket authentication / ownership / 404
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSRejectsUnauthenticated:
|
||||
def test_no_token_closes_with_4401(self):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
with patch.object(vv_module, "_ws_authenticate_user", return_value=None):
|
||||
client = TestClient(app)
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-1"):
|
||||
pass
|
||||
assert exc.value.code == 4401
|
||||
|
||||
|
||||
def _build_client(*, auth_user, repo_get_return, redis_instance=None):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = repo_get_return
|
||||
sess = MagicMock()
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=auth_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch("redis.from_url", return_value=redis_instance or MagicMock()),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
return TestClient(app), fake_repo, sess, patches
|
||||
|
||||
|
||||
class TestWSOwnershipAnd404:
|
||||
def test_other_users_job_closes_with_4403(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
other_job = _make_job(user_id="user-b")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=other_job)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-x?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4403
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
def test_missing_job_closes_with_4404(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=None)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-missing?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4404
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _ws_authenticate_user direct unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSAuthenticateUser:
|
||||
def test_empty_token_returns_none(self):
|
||||
assert vv_module._ws_authenticate_user("") is None
|
||||
|
||||
def test_decode_exception_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", side_effect=Exception("bad token")):
|
||||
assert vv_module._ws_authenticate_user("not-a-jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_missing_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_non_string_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": 123}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_success_returns_user(self):
|
||||
sess = MagicMock()
|
||||
fake_user = SimpleNamespace(id="u1")
|
||||
fake_user_repo = MagicMock()
|
||||
fake_user_repo.find_by_id.return_value = fake_user
|
||||
with patch.object(_app_db, "SessionLocal", return_value=sess):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": "u1"}):
|
||||
with patch(
|
||||
"app.dependencies.get_user_repository",
|
||||
return_value=fake_user_repo,
|
||||
):
|
||||
result = vv_module._ws_authenticate_user("valid.jwt")
|
||||
assert result is fake_user
|
||||
fake_user_repo.find_by_id.assert_called_once_with("u1")
|
||||
sess.close.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket initial-snapshot / terminal-job fast-close tests.
|
||||
#
|
||||
# The Redis pubsub reader thread is factored into ``_run_pubsub_forwarder`` and
|
||||
# marked ``# pragma: no cover`` (integration-tested with a live Redis). These
|
||||
# tests patch it out so we can deterministically verify the pre-subscribe
|
||||
# handshake without needing a real Redis or real thread scheduling.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _run_ws_handshake(*, job):
|
||||
"""Drive a WS handshake; collect JSON messages before connection closes."""
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_user = SimpleNamespace(id=getattr(job, "user_id", "user-a"))
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = job
|
||||
sess = MagicMock()
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
return received, fake_repo, sess
|
||||
|
||||
|
||||
class TestWSInitialSnapshot:
|
||||
def test_running_job_sends_initial_snapshot(self):
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
received, repo, sess = _run_ws_handshake(job=job)
|
||||
assert received[0]["type"] == "viral_video:progress"
|
||||
assert received[0]["job_id"] == "job-1"
|
||||
assert received[0]["data"]["status"] == "running"
|
||||
# Session was used for both ownership check and initial snapshot.
|
||||
assert sess.close.call_count >= 2
|
||||
|
||||
def test_running_job_with_enum_status(self):
|
||||
job = _make_job(
|
||||
status=SimpleNamespace(value="wait_user_confirm"),
|
||||
user_id="user-a",
|
||||
is_terminal=False,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
assert received[0]["data"]["status"] == "wait_user_confirm"
|
||||
assert received[0]["progress"] == 35.0
|
||||
assert "等待" in received[0]["message"]
|
||||
|
||||
def test_already_completed_job_sends_completion_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url="https://example.com/v.mp4",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
types = [m["type"] for m in received]
|
||||
assert "viral_video:progress" in types
|
||||
assert "viral_video:completed" in types
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == "https://example.com/v.mp4"
|
||||
assert completed["progress"] == 100
|
||||
|
||||
def test_already_failed_job_sends_failed_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg="out of memory",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == "out of memory"
|
||||
assert failed["progress"] == 0
|
||||
|
||||
def test_completed_job_without_result_url_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == ""
|
||||
|
||||
def test_failed_job_without_error_msg_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == ""
|
||||
|
||||
def test_initial_snapshot_exception_is_swallowed(self):
|
||||
"""If sending the initial snapshot raises, the endpoint should log and
|
||||
still proceed to the Redis forwarder (doesn't crash)."""
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
await websocket.send_json({"type": "forwarder_reached"})
|
||||
await websocket.close()
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
fake_repo = MagicMock()
|
||||
calls = {"n": 0}
|
||||
|
||||
def _get(job_id):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 2:
|
||||
raise RuntimeError("boom in snapshot")
|
||||
return job
|
||||
|
||||
fake_repo.get.side_effect = _get
|
||||
sess = MagicMock()
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
assert any(m["type"] == "forwarder_reached" for m in received)
|
||||
Reference in New Issue
Block a user