Compare commits
43 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 2332b5ef98 | |||
| 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 | |||
| 6636dc45f7 | |||
| e11e4f0e99 | |||
| a7d6ba473b | |||
| 966da04c9c | |||
| fdeb792bab | |||
| ff1d878c62 | |||
| f19be5fd09 | |||
| eeb8a05b69 | |||
| efb7fa5729 | |||
| 5e1520230f | |||
| 96bcec5fdd | |||
| dfc5e5a5b6 |
+100
@@ -0,0 +1,100 @@
|
||||
"""add viral video tables
|
||||
|
||||
Revision ID: 086_add_viral_video_tables
|
||||
Revises: 085_atom_clip_caption_embedding
|
||||
Create Date: 2026-09-28
|
||||
|
||||
新增爆款视频相关表:
|
||||
- viral_video_jobs: 爆款视频任务
|
||||
- viral_video_style_templates: 风格模板配置
|
||||
- viral_video_prompt_templates: Prompt 模板(由 #2040 seed)
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "086_add_viral_video_tables"
|
||||
down_revision = "085_atom_clip_caption_embedding"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# viral_video_jobs
|
||||
op.create_table(
|
||||
"viral_video_jobs",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, index=True),
|
||||
sa.Column("images", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("industry", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("target_customer", sa.String(500), nullable=False, server_default=""),
|
||||
sa.Column("persona_id", sa.String(36), nullable=False, server_default=""),
|
||||
sa.Column("viral_structure", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("marketing_purpose", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("bgm_preference", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("duration", sa.Integer(), nullable=False, server_default="30"),
|
||||
sa.Column("user_copy_text", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("fusion_level", sa.String(20), nullable=False, server_default="ai_polish"),
|
||||
sa.Column("reference_audio_path", sa.String(1000), nullable=False, server_default=""),
|
||||
# v1.3 新增
|
||||
sa.Column("reference_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_strength", sa.String(20), nullable=False, server_default="medium"),
|
||||
sa.Column("style_guide", sa.JSON(), nullable=True),
|
||||
sa.Column("style_template_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
# 状态与结果
|
||||
sa.Column("status", sa.String(30), nullable=False, server_default="pending", index=True),
|
||||
sa.Column("intent_result", sa.JSON(), nullable=True),
|
||||
sa.Column("result_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("credits_cost", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error_msg", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("retry_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_style_templates
|
||||
op.create_table(
|
||||
"viral_video_style_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("thumbnail_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_config", sa.JSON(), nullable=False, server_default="{}"),
|
||||
sa.Column("is_system", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("sort_order", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_prompt_templates
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("prompt_type", sa.String(50), nullable=False, index=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("variables", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("version", sa.Integer(), nullable=False, server_default="1"),
|
||||
sa.Column("is_active", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# Seed 默认风格模板
|
||||
op.execute("""
|
||||
INSERT INTO viral_video_style_templates (id, name, description, style_config, is_system, sort_order)
|
||||
VALUES
|
||||
('style-tpl-001', '快节奏冲击', '高频切镜+动感BGM,适合食品饮料等快消品', '{"cut_speed": "fast", "transition": "jump_cut", "energy": "high"}', true, 1),
|
||||
('style-tpl-002', '质感慢镜', '慢节奏+电影感调色,适合美妆护肤珠宝', '{"cut_speed": "slow", "transition": "dissolve", "energy": "low", "color_grade": "cinematic"}', true, 2),
|
||||
('style-tpl-003', '口播种草', '数字人口播+产品特写穿插', '{"cut_speed": "medium", "transition": "cross_dissolve", "has_talking_head": true}', true, 3),
|
||||
('style-tpl-004', '场景叙事', '多场景切换+故事线叙述', '{"cut_speed": "medium", "transition": "wipe", "narrative": true}', true, 4)
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
op.drop_table("viral_video_style_templates")
|
||||
op.drop_table("viral_video_jobs")
|
||||
@@ -0,0 +1,25 @@
|
||||
"""viral video add image_analysis column
|
||||
|
||||
Revision ID: 087_viral_video_image_analysis
|
||||
Revises: 086_add_viral_video_tables
|
||||
Create Date: 2026-09-30
|
||||
|
||||
#2106 爆款视频 P0:持久化图片分析结果(image_analysis JSON),供 resume 阶段使用。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "087_viral_video_image_analysis"
|
||||
down_revision = "086_add_viral_video_tables"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("viral_video_jobs", sa.Column("image_analysis", sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("viral_video_jobs", "image_analysis")
|
||||
@@ -36,6 +36,7 @@ from app.api.routes.titles import router as titles_router
|
||||
from app.api.routes.tts import router as tts_router
|
||||
from app.api.routes.upload import router as upload_router
|
||||
from app.api.routes.videos import router as videos_router
|
||||
from app.api.routes.viral_video import router as viral_video_router
|
||||
from app.api.routes.voice_clones import router as voice_clones_router
|
||||
from app.api.routes.voices import router as voices_router
|
||||
from fastapi import APIRouter
|
||||
@@ -240,3 +241,4 @@ api_router.include_router(
|
||||
prefix="/gpu",
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
api_router.include_router(viral_video_router, prefix="/viral-video", tags=["爆款视频"])
|
||||
|
||||
@@ -11,9 +11,7 @@ from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.core.task_enqueue import (
|
||||
GLOBAL_PENDING_LIMIT,
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
@@ -302,26 +300,17 @@ def create_preview_generation_task(
|
||||
count,
|
||||
)
|
||||
|
||||
# 预检查队列限流(按变体总数计)
|
||||
try:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending + count > USER_PENDING_LIMIT:
|
||||
raise UserPendingLimitExceeded(
|
||||
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
|
||||
)
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
# 预检查队列限流(按变体总数计)——仅保留全局硬上限,用户上限改为软 warning 在 safe_enqueue 内处理(#2098)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
|
||||
) from e
|
||||
detail=build_rate_limit_detail(
|
||||
GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT),
|
||||
generation_task_repository,
|
||||
scope="global",
|
||||
),
|
||||
)
|
||||
|
||||
# 确定视频比例:优先前端传入,否则从模板 mode 推断
|
||||
video_ratio = request.video_ratio or ""
|
||||
@@ -589,9 +578,6 @@ def create_preview_generation_task(
|
||||
if not enqueued:
|
||||
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
|
||||
_mark_task_failed(generation_task_repository, task, "任务入队失败")
|
||||
except UserPendingLimitExceeded as e:
|
||||
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
|
||||
rate_limit_exc = rate_limit_exc or e
|
||||
except GlobalQueueFull as e:
|
||||
_mark_task_failed(generation_task_repository, task, "系统队列已满")
|
||||
rate_limit_exc = rate_limit_exc or e
|
||||
@@ -603,11 +589,6 @@ def create_preview_generation_task(
|
||||
|
||||
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败")
|
||||
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
|
||||
if isinstance(rate_limit_exc, UserPendingLimitExceeded):
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"),
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"),
|
||||
|
||||
@@ -7,7 +7,6 @@ from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.core.task_enqueue import (
|
||||
GLOBAL_PENDING_LIMIT,
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
@@ -50,12 +49,98 @@ from packages.domain.smart_match import smart_select_assets
|
||||
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
|
||||
# AssetClassification 枚举: scenic / product / person / animal / food / tech / sport / music / other
|
||||
_CATEGORY_KEYWORDS: dict[str, set[str]] = {
|
||||
"scenic": {"风景", "自然", "山水", "大海", "天空", "日落", "日出", "森林", "城市", "建筑", "夜景", "街道", "公园", "景区", "旅行", "旅游", "户外"},
|
||||
"product": {"产品", "商品", "展示", "演示", "开箱", "评测", "好物", "推荐", "种草", "购物", "电商", "带货", "品牌", "广告", "包装"},
|
||||
"person": {"人物", "人物采访", "对话", "说话", "讲解", "演讲", "采访", "聊天", "开会", "工作", "办公室", "团队", "员工", "老板", "女性", "男性", "美女", "帅哥"},
|
||||
"scenic": {
|
||||
"风景",
|
||||
"自然",
|
||||
"山水",
|
||||
"大海",
|
||||
"天空",
|
||||
"日落",
|
||||
"日出",
|
||||
"森林",
|
||||
"城市",
|
||||
"建筑",
|
||||
"夜景",
|
||||
"街道",
|
||||
"公园",
|
||||
"景区",
|
||||
"旅行",
|
||||
"旅游",
|
||||
"户外",
|
||||
},
|
||||
"product": {
|
||||
"产品",
|
||||
"商品",
|
||||
"展示",
|
||||
"演示",
|
||||
"开箱",
|
||||
"评测",
|
||||
"好物",
|
||||
"推荐",
|
||||
"种草",
|
||||
"购物",
|
||||
"电商",
|
||||
"带货",
|
||||
"品牌",
|
||||
"广告",
|
||||
"包装",
|
||||
},
|
||||
"person": {
|
||||
"人物",
|
||||
"人物采访",
|
||||
"对话",
|
||||
"说话",
|
||||
"讲解",
|
||||
"演讲",
|
||||
"采访",
|
||||
"聊天",
|
||||
"开会",
|
||||
"工作",
|
||||
"办公室",
|
||||
"团队",
|
||||
"员工",
|
||||
"老板",
|
||||
"女性",
|
||||
"男性",
|
||||
"美女",
|
||||
"帅哥",
|
||||
},
|
||||
"animal": {"动物", "宠物", "狗", "猫", "鸟", "鱼", "马", "牛", "羊", "野生动物", "动物园"},
|
||||
"food": {"美食", "食物", "餐饮", "餐厅", "做饭", "烹饪", "厨房", "菜品", "饮料", "水果", "甜点", "蛋糕", "咖啡", "茶", "零食", "吃"},
|
||||
"tech": {"科技", "数码", "电脑", "手机", "屏幕", "软件", "APP", "互联网", "AI", "人工智能", "机器人", "办公", "程序员", "代码", "屏幕录制"},
|
||||
"food": {
|
||||
"美食",
|
||||
"食物",
|
||||
"餐饮",
|
||||
"餐厅",
|
||||
"做饭",
|
||||
"烹饪",
|
||||
"厨房",
|
||||
"菜品",
|
||||
"饮料",
|
||||
"水果",
|
||||
"甜点",
|
||||
"蛋糕",
|
||||
"咖啡",
|
||||
"茶",
|
||||
"零食",
|
||||
"吃",
|
||||
},
|
||||
"tech": {
|
||||
"科技",
|
||||
"数码",
|
||||
"电脑",
|
||||
"手机",
|
||||
"屏幕",
|
||||
"软件",
|
||||
"APP",
|
||||
"互联网",
|
||||
"AI",
|
||||
"人工智能",
|
||||
"机器人",
|
||||
"办公",
|
||||
"程序员",
|
||||
"代码",
|
||||
"屏幕录制",
|
||||
},
|
||||
"sport": {"运动", "健身", "跑步", "篮球", "足球", "游泳", "瑜伽", "户外", "锻炼", "体育", "比赛", "球场"},
|
||||
"music": {"音乐", "歌曲", "演唱会", "乐器", "唱歌", "跳舞", "舞蹈", "MV", "演出", "乐队", "钢琴", "吉他", "节奏"},
|
||||
}
|
||||
@@ -77,6 +162,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
break
|
||||
return matched or None
|
||||
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -193,10 +279,13 @@ def _select_assets_from_library(
|
||||
# #2035:加载片段级 AI 标签,供叙事模式 AI 加权和 smart 模式语义匹配使用。
|
||||
# 失败降级为空(不影响选片主流程)。
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] = {}
|
||||
ai_tags_by_asset: dict[str, dict] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重)
|
||||
ai_tags_by_asset: dict[
|
||||
str, dict
|
||||
] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重)
|
||||
try:
|
||||
if db is not None:
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel
|
||||
|
||||
ready_ids = [a.id for a in ready_video_assets]
|
||||
clip_rows = (
|
||||
db.query(AssetAtomClipModel.asset_id, AssetAtomClipModel.ai_tags)
|
||||
@@ -594,21 +683,13 @@ def create_generation_task(
|
||||
# 同批次任务共享 batch_id,用于视频查重时批次内比对
|
||||
batch_id = uuid.uuid4().hex if count > 1 else ""
|
||||
|
||||
# 预检查:批量提交前先看会不会超限,避免建一半才拒
|
||||
# 预检查(Bug B #2098):只保留全局 503 保护,用户级不再硬拒 429;
|
||||
# 超额任务直接入队等待 worker 自然消费,前端展示排队位置而非阻止提交。
|
||||
# USER_PENDING_LIMIT 作为软上限(safe_enqueue 兜底),提高到 20 支持批量提交。
|
||||
try:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending + count > USER_PENDING_LIMIT:
|
||||
raise UserPendingLimitExceeded(
|
||||
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
|
||||
)
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -923,14 +1004,10 @@ def create_generation_task(
|
||||
else:
|
||||
failed_tasks.append(task)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
# 兜底:如果预检查后又并发提交了,在这里也拦住
|
||||
failed_tasks.append(task)
|
||||
if not created_tasks:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from _e
|
||||
break
|
||||
# Bug B #2098: 用户级限流已改为软限制,此分支理论上不再触发;
|
||||
# 极端并发兜底仍入队(safe_enqueue 内部会打 warning 日志),不 429 拒绝
|
||||
logger.warning("[生成任务] 用户 pending 超软限制,仍允许入队: task_id=%s", task.id)
|
||||
created_tasks.append(task)
|
||||
except GlobalQueueFull as _e:
|
||||
failed_tasks.append(task)
|
||||
if not created_tasks:
|
||||
@@ -1081,10 +1158,8 @@ def confirm_generation(
|
||||
):
|
||||
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
# Bug B #2098: 用户级限流已软处理,理论上不再触发;作为防御仍放行
|
||||
logger.warning("[任务] 用户 pending 超软限制,任务已入队")
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -1263,22 +1338,8 @@ def retry_generation_task(
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
|
||||
user_id = authenticated_user.user.id
|
||||
# 预检查:创建前判断,>= 上限就拒绝
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
# 预检查(Bug B #2098):只保留全局 503,用户级不再硬拒
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending >= USER_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(
|
||||
UserPendingLimitExceeded(
|
||||
user_id=user_id,
|
||||
pending_count=user_pending,
|
||||
limit=USER_PENDING_LIMIT,
|
||||
),
|
||||
generation_task_repository,
|
||||
scope="user",
|
||||
),
|
||||
)
|
||||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -1321,10 +1382,8 @@ def retry_generation_task(
|
||||
):
|
||||
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
# Bug B #2098: 用户级限流已软处理,理论上不再触发;作为防御仍放行
|
||||
logger.warning("[任务] 用户 pending 超软限制,任务已入队")
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
|
||||
@@ -390,6 +390,7 @@ async def prepare_direct_upload(
|
||||
duplicated=True,
|
||||
skip_transfer=True,
|
||||
asset_id=existing.id,
|
||||
url=existing.file_url or storage_service.get_url(existing.storage_key) or "",
|
||||
)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
@@ -443,6 +444,7 @@ async def prepare_direct_upload(
|
||||
duplicated=False,
|
||||
skip_transfer=False,
|
||||
asset_id=pending_asset_id,
|
||||
url="",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,549 @@
|
||||
"""爆款视频 API 路由。
|
||||
|
||||
端点:
|
||||
POST /api/v1/viral-video/generate 创建爆款视频任务
|
||||
GET /api/v1/viral-video/{job_id} 查询任务状态
|
||||
GET /api/v1/viral-video/history 历史记录
|
||||
POST /api/v1/viral-video/{job_id}/retry 重试失败任务
|
||||
POST /api/v1/viral-video/{job_id}/confirm-intent 确认意图文案
|
||||
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
|
||||
GET /api/v1/viral-video/style-templates 获取风格模板列表
|
||||
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送(订阅 Redis pub/sub)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.viral_video import (
|
||||
AnalyzeStyleRequest,
|
||||
AnalyzeStyleResponse,
|
||||
ConfirmIntentRequest,
|
||||
CreateViralVideoRequest,
|
||||
StyleTemplateListResponse,
|
||||
StyleTemplateResponse,
|
||||
ViralVideoHistoryResponse,
|
||||
ViralVideoJobResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _to_response(job) -> ViralVideoJobResponse:
|
||||
return ViralVideoJobResponse(
|
||||
id=job.id,
|
||||
user_id=job.user_id,
|
||||
images=job.images,
|
||||
industry=job.industry,
|
||||
target_customer=job.target_customer,
|
||||
persona_id=job.persona_id,
|
||||
viral_structure=job.viral_structure,
|
||||
marketing_purpose=job.marketing_purpose,
|
||||
bgm_preference=job.bgm_preference,
|
||||
duration=job.duration,
|
||||
user_copy_text=job.user_copy_text,
|
||||
fusion_level=job.fusion_level,
|
||||
reference_audio_path=job.reference_audio_path,
|
||||
reference_video_url=job.reference_video_url,
|
||||
style_strength=job.style_strength,
|
||||
style_guide=job.style_guide,
|
||||
style_template_id=job.style_template_id,
|
||||
status=job.status,
|
||||
intent_result=job.intent_result,
|
||||
result_video_url=job.result_video_url,
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
completed_at=job.completed_at,
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _get_job_repo(session: Session) -> SQLAlchemyViralVideoJobRepository:
|
||||
return SQLAlchemyViralVideoJobRepository(session)
|
||||
|
||||
|
||||
def _get_style_repo(session: Session) -> SQLAlchemyViralVideoStyleTemplateRepository:
|
||||
return SQLAlchemyViralVideoStyleTemplateRepository(session)
|
||||
|
||||
|
||||
# ── Endpoints ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/generate", response_model=ViralVideoJobResponse)
|
||||
def create_viral_video(
|
||||
request: CreateViralVideoRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""创建爆款视频任务,入队 Celery 编排器。"""
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
repo = _get_job_repo(session)
|
||||
|
||||
# 创建领域实体
|
||||
job = ViralVideoJob(
|
||||
user_id=authenticated_user.user.id,
|
||||
images=list(request.images),
|
||||
industry=request.industry,
|
||||
target_customer=request.target_customer,
|
||||
persona_id=request.persona_id,
|
||||
viral_structure=request.viral_structure,
|
||||
marketing_purpose=request.marketing_purpose,
|
||||
bgm_preference=request.bgm_preference,
|
||||
duration=request.duration,
|
||||
user_copy_text=request.user_copy_text,
|
||||
fusion_level=request.fusion_level,
|
||||
reference_audio_path=request.reference_audio_path,
|
||||
reference_video_url=request.reference_video_url,
|
||||
style_strength=request.style_strength,
|
||||
style_template_id=request.style_template_id,
|
||||
)
|
||||
|
||||
# 持久化
|
||||
repo.save(job)
|
||||
|
||||
# 入队 Celery 任务
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 任务已入队: job_id=%s user_id=%s", job.id, job.user_id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"任务入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.get("/history", response_model=ViralVideoHistoryResponse)
|
||||
def list_viral_video_history(
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoHistoryResponse:
|
||||
"""获取用户的爆款视频历史列表。"""
|
||||
repo = _get_job_repo(session)
|
||||
jobs = repo.list_by_user(authenticated_user.user.id, limit=limit, offset=offset)
|
||||
items = [_to_response(j) for j in jobs]
|
||||
return ViralVideoHistoryResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.get("/style-templates", response_model=StyleTemplateListResponse)
|
||||
def list_style_templates(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> StyleTemplateListResponse:
|
||||
"""获取风格模板列表。"""
|
||||
repo = _get_style_repo(session)
|
||||
templates = repo.list_all()
|
||||
items = [
|
||||
StyleTemplateResponse(
|
||||
id=t["id"],
|
||||
name=t["name"],
|
||||
description=t["description"],
|
||||
thumbnail_url=t["thumbnail_url"],
|
||||
style_config=t["style_config"],
|
||||
)
|
||||
for t in templates
|
||||
]
|
||||
return StyleTemplateListResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
|
||||
def get_viral_video_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""查询爆款视频任务状态。"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权查看此任务")
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/retry", response_model=ViralVideoJobResponse)
|
||||
def retry_viral_video_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""重试失败的爆款视频任务。"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status != ViralVideoStatus.FAILED:
|
||||
raise HTTPException(status_code=409, detail="只有失败的任务可以重试")
|
||||
|
||||
# 重置状态
|
||||
job.retry_count += 1
|
||||
job.status = ViralVideoStatus.PENDING
|
||||
job.error_msg = ""
|
||||
job.started_at = None
|
||||
job.completed_at = None
|
||||
repo.update(job)
|
||||
|
||||
# 重新入队
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d", job.id, job.retry_count)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"重试入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/confirm-intent", response_model=ViralVideoJobResponse)
|
||||
def confirm_intent(
|
||||
job_id: str,
|
||||
request: ConfirmIntentRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""用户确认/修改 AI 生成的意图文案,恢复流水线。"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise HTTPException(status_code=409, detail="任务当前不在等待确认状态")
|
||||
|
||||
# 更新文案
|
||||
if request.confirmed_copy:
|
||||
job.user_copy_text = request.confirmed_copy
|
||||
|
||||
# 恢复流水线
|
||||
job.resume_from_confirm()
|
||||
repo.update(job)
|
||||
|
||||
# 从断点恢复 Celery 任务
|
||||
try:
|
||||
celery_app.send_task("worker.resume_viral_video_pipeline", args=[job.id])
|
||||
logger.info("[爆款视频] 意图确认,恢复流水线: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"恢复流水线失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/analyze-style", response_model=AnalyzeStyleResponse)
|
||||
def analyze_style(
|
||||
job_id: str,
|
||||
request: AnalyzeStyleRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> AnalyzeStyleResponse:
|
||||
"""触发参考视频风格分析(独立步骤,可在生成前单独调用)。"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
|
||||
# 更新参考视频 URL
|
||||
job.reference_video_url = request.reference_video_url
|
||||
if request.style_template_id:
|
||||
job.style_template_id = request.style_template_id
|
||||
repo.update(job)
|
||||
|
||||
# 入队风格分析任务
|
||||
try:
|
||||
celery_app.send_task("worker.run_video_style_analysis", args=[job.id])
|
||||
logger.info("[爆款视频] 风格分析入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 风格分析入队失败: %s", e, exc_info=True)
|
||||
|
||||
return AnalyzeStyleResponse(
|
||||
job_id=job.id,
|
||||
status="analyzing",
|
||||
style_guide=None,
|
||||
)
|
||||
|
||||
|
||||
# ── WebSocket 进度推送 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ws_authenticate_user(token: str):
|
||||
"""从 token 字符串解析用户(复用 HTTP Bearer 的解码 + 黑名单逻辑)。
|
||||
|
||||
WebSocket 握手阶段不能发自定义 Authorization header,
|
||||
因此统一通过 query 参数 ``?token=...`` 传 JWT。
|
||||
"""
|
||||
from app.auth import _decode_user_token
|
||||
from app.dependencies import get_user_repository
|
||||
|
||||
if not token:
|
||||
return None
|
||||
try:
|
||||
payload = _decode_user_token(token)
|
||||
except Exception:
|
||||
return None
|
||||
user_id = payload.get("sub")
|
||||
if not isinstance(user_id, str) or not user_id:
|
||||
return None
|
||||
# 同步场景下手动拉 repository 实例
|
||||
from app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
user_repo = get_user_repository(session)
|
||||
user = user_repo.find_by_id(user_id)
|
||||
return user
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
async def _run_pubsub_forwarder(
|
||||
websocket, redis_lib, settings, job_id: str
|
||||
) -> None: # pragma: no cover - integration tested (real Redis + thread)
|
||||
"""订阅 Redis 频道并把消息桥接到 WebSocket,终态消息后自动关闭。
|
||||
|
||||
该函数封装了线程 + asyncio.Queue 桥接逻辑,在单测中可被整体替换为桩,
|
||||
避免引入真实 Redis 与线程调度的不确定性。
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
|
||||
r = redis_lib.from_url(settings.REDIS_URL, decode_responses=True)
|
||||
pubsub = r.pubsub(ignore_subscribe_messages=True)
|
||||
channel = f"viral_video:{job_id}"
|
||||
pubsub.subscribe(channel)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
queue: asyncio.Queue = asyncio.Queue(maxsize=64)
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
def _reader() -> None:
|
||||
try:
|
||||
while not stop_event.is_set():
|
||||
msg = pubsub.get_message(timeout=0.5)
|
||||
if msg is None or msg.get("type") != "message":
|
||||
continue
|
||||
raw = msg.get("data")
|
||||
if not isinstance(raw, str):
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(raw)
|
||||
except Exception:
|
||||
payload = {"type": "viral_video:progress", "data": {"raw": raw}}
|
||||
loop.call_soon_threadsafe(queue.put_nowait, payload)
|
||||
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
|
||||
loop.call_soon_threadsafe(stop_event.set)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频WS] pubsub reader 异常退出: %s", e)
|
||||
loop.call_soon_threadsafe(stop_event.set)
|
||||
|
||||
try:
|
||||
reader_thread = threading.Thread(target=_reader, name=f"viral-video-ws-{job_id}", daemon=True)
|
||||
reader_thread.start()
|
||||
|
||||
while not stop_event.is_set():
|
||||
try:
|
||||
payload = await asyncio.wait_for(queue.get(), timeout=1.0)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
try:
|
||||
await websocket.send_json(payload)
|
||||
except Exception:
|
||||
break
|
||||
if payload.get("type") in ("viral_video:completed", "viral_video:failed"):
|
||||
break
|
||||
except WebSocketDisconnect:
|
||||
logger.info("[爆款视频WS] 客户端断开: job_id=%s", job_id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频WS] 转发异常: %s", e, exc_info=True)
|
||||
try:
|
||||
await websocket.send_json({"type": "viral_video:error", "message": f"服务异常: {e}"})
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
stop_event.set()
|
||||
try:
|
||||
pubsub.unsubscribe(channel)
|
||||
pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
r.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@router.websocket("/ws/{job_id}")
|
||||
async def viral_video_websocket(websocket: WebSocket, job_id: str) -> None:
|
||||
"""WebSocket 桥接:订阅 Redis `viral_video:{job_id}` 频道并转发给前端。
|
||||
|
||||
认证:通过 ``?token=<jwt>`` query 参数传 JWT(浏览器 WS 握手不支持自定义 header)。
|
||||
事件类型:
|
||||
- viral_video:progress 中间进度(progress: 0-100)
|
||||
- viral_video:wait_user 等待用户确认意图文案
|
||||
- viral_video:completed 任务完成(data.video_url)
|
||||
- viral_video:failed 任务失败(data.error)
|
||||
- viral_video:error 服务端错误(如鉴权失败 / job 不存在 / 无权限)
|
||||
"""
|
||||
|
||||
import redis as redis_lib
|
||||
from app.config import settings
|
||||
|
||||
# ── 1. 鉴权 ──────────────────────────────────────────────────────
|
||||
token = websocket.query_params.get("token", "")
|
||||
user = _ws_authenticate_user(token)
|
||||
if user is None:
|
||||
await websocket.close(code=4401, reason="Unauthorized")
|
||||
return
|
||||
|
||||
# ── 2. 校验 job 归属 ─────────────────────────────────────────────
|
||||
from app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
job_repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
await websocket.close(code=4404, reason="Job not found")
|
||||
return
|
||||
if job.user_id != user.id:
|
||||
await websocket.close(code=4403, reason="Forbidden")
|
||||
return
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
await websocket.accept()
|
||||
|
||||
# ── 3. 发送一条初始状态(前端连接后立即拿到当前进度) ────────────
|
||||
try:
|
||||
session = SessionLocal()
|
||||
job_repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is not None:
|
||||
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
initial = {
|
||||
"type": "viral_video:progress",
|
||||
"job_id": job_id,
|
||||
"stage": _stage_from_status(job),
|
||||
"progress": _estimate_progress(job),
|
||||
"message": _initial_message(job),
|
||||
"data": {"status": status_val},
|
||||
}
|
||||
await websocket.send_json(initial)
|
||||
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
|
||||
if job.is_terminal:
|
||||
is_completed = status_val == "completed"
|
||||
terminal_type = "viral_video:completed" if is_completed else "viral_video:failed"
|
||||
terminal_data = (
|
||||
{"video_url": job.result_video_url or ""} if is_completed else {"error": job.error_msg or ""}
|
||||
)
|
||||
await websocket.send_json(
|
||||
{
|
||||
"type": terminal_type,
|
||||
"job_id": job_id,
|
||||
"stage": "",
|
||||
"progress": 100 if is_completed else 0,
|
||||
"message": "视频生成完成" if is_completed else "任务失败",
|
||||
"data": terminal_data,
|
||||
}
|
||||
)
|
||||
await websocket.close()
|
||||
return
|
||||
session.close()
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频WS] 发送初始状态失败: %s", e)
|
||||
try:
|
||||
session.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ── 4. 订阅 Redis 频道并转发 ─────────────────────────────────────
|
||||
# redis-py 的 pubsub 是同步阻塞的,放到线程里跑,通过 asyncio.Queue 桥接到 event loop。
|
||||
# 该段依赖真实 Redis + 线程调度,属于集成测试范围,单测通过桩替换。
|
||||
await _run_pubsub_forwarder(websocket, redis_lib, settings, job_id)
|
||||
|
||||
|
||||
def _job_status(job) -> str:
|
||||
return job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
|
||||
|
||||
# 初始快照的 stage 推断:领域对象不持久化 stage,
|
||||
# 只能根据 status 给一个占位,后续 worker 推送的真实进度事件会覆盖。
|
||||
_STATUS_STAGE = {
|
||||
"pending": "",
|
||||
"running": "",
|
||||
"wait_user_confirm": "intent_parsing",
|
||||
"completed": "uploading",
|
||||
"failed": "",
|
||||
"cancelled": "",
|
||||
}
|
||||
|
||||
_STATUS_PROGRESS = {
|
||||
"pending": 0.0,
|
||||
"running": 5.0,
|
||||
"wait_user_confirm": 35.0,
|
||||
"completed": 100.0,
|
||||
"failed": 0.0,
|
||||
"cancelled": 0.0,
|
||||
}
|
||||
|
||||
_STATUS_MESSAGE = {
|
||||
"pending": "任务已创建,等待执行",
|
||||
"running": "任务执行中",
|
||||
"wait_user_confirm": "等待用户确认意图文案",
|
||||
"completed": "视频生成完成",
|
||||
"failed": "任务失败",
|
||||
"cancelled": "任务已取消",
|
||||
}
|
||||
|
||||
|
||||
def _stage_from_status(job) -> str:
|
||||
return _STATUS_STAGE.get(_job_status(job), "")
|
||||
|
||||
|
||||
def _estimate_progress(job) -> float:
|
||||
"""根据 status 粗略估算百分比(0-100),用于连接初始快照;
|
||||
连接建立后由 Redis 推送的真实事件持续更新。
|
||||
"""
|
||||
return _STATUS_PROGRESS.get(_job_status(job), 5.0)
|
||||
|
||||
|
||||
def _initial_message(job) -> str:
|
||||
"""给新连接的前端一个可读的初始状态文案。"""
|
||||
status_val = _job_status(job)
|
||||
if status_val == "failed" and job.error_msg:
|
||||
return f"任务失败: {job.error_msg}"
|
||||
return _STATUS_MESSAGE.get(status_val, "任务准备中")
|
||||
@@ -6,7 +6,7 @@ from app.core.celery_app import celery_app
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
|
||||
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
|
||||
USER_PENDING_LIMIT = 20 # 单用户 pending 上限(#2098: 从 3 提到 20,支持批量任务自动排队)
|
||||
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
|
||||
WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值)
|
||||
|
||||
@@ -154,19 +154,18 @@ def check_queue_limits(
|
||||
user_pending_limit: int = USER_PENDING_LIMIT,
|
||||
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
|
||||
) -> None:
|
||||
"""检查队列限流(预检查用,任务创建前调用),超限抛对应异常。
|
||||
"""检查队列限流(预检查用,任务创建前调用)。
|
||||
|
||||
边界语义:>= 上限即拒绝(达到上限就不能再加新任务)。
|
||||
#2098 语义变更:用户级限流改为软提示,不再抛异常拒绝;仅全局硬上限抛 GlobalQueueFull。
|
||||
|
||||
Args:
|
||||
user_id: 用户 ID
|
||||
user_id: 用户 ID(保留参数,当前不做用户级硬拒)
|
||||
generation_task_repository: 任务仓储
|
||||
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
|
||||
user_pending_limit: 单用户 pending 上限(保留,当前未硬拒)
|
||||
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
|
||||
|
||||
Raises:
|
||||
GlobalQueueFull: 全局超限时抛出(优先级更高,先查全局)
|
||||
UserPendingLimitExceeded: 用户超限时抛出
|
||||
GlobalQueueFull: 全局超限时抛出
|
||||
"""
|
||||
# 先查全局(系统级保护优先级更高)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
@@ -179,17 +178,9 @@ def check_queue_limits(
|
||||
)
|
||||
raise GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
|
||||
|
||||
# 再查用户级
|
||||
if user_id:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
if user_pending >= user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
|
||||
user_id,
|
||||
user_pending,
|
||||
user_pending_limit,
|
||||
)
|
||||
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
|
||||
# #2098: 用户级限流改为软提示,不在预检查阶段拒绝(超额任务仍入队排队)。
|
||||
# 真正的系统保护由全局 GLOBAL_PENDING_LIMIT 硬上限承担。
|
||||
# UserPendingLimitExceeded 保留以兼容历史 import/except,但预检查与 safe_enqueue 均不再 raise。
|
||||
|
||||
|
||||
def _mark_task_failed_safely(
|
||||
@@ -246,7 +237,6 @@ def safe_enqueue_generation_task(
|
||||
|
||||
Raises:
|
||||
GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed
|
||||
UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed
|
||||
"""
|
||||
# ── 入队前检查:任务已是 pending,用 > 判断(包含当前任务) ──
|
||||
|
||||
@@ -263,19 +253,18 @@ def safe_enqueue_generation_task(
|
||||
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
||||
raise exc
|
||||
|
||||
# 用户级限流检查(传了 user_id 才做)
|
||||
# Bug B #2098: 用户级限流改为软提示,不再硬拒;所有任务都入队等待 worker 自然消费。
|
||||
# user_pending_limit 作为兜底阈值保留(默认 20),达到时打 warning 日志但仍入队,
|
||||
# 避免极端情况下恶意用户无限堆积任务。真正的系统保护由全局 GLOBAL_PENDING_LIMIT 承担。
|
||||
if user_id:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
if user_pending > user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 任务数超限(入队前): user_id=%s, count=%d/%d",
|
||||
"[队列限流] 用户 pending 任务数超过软上限(入队): user_id=%s, count=%d/%d, 仍允许入队排队",
|
||||
user_id,
|
||||
user_pending,
|
||||
user_pending_limit,
|
||||
)
|
||||
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
|
||||
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
||||
raise exc
|
||||
|
||||
# ── 发送 Celery 任务 ──
|
||||
try:
|
||||
@@ -317,16 +306,18 @@ def safe_enqueue_generation_task(
|
||||
user_after = generation_task_repository.count_pending_by_user(user_id) if user_id else 0
|
||||
|
||||
global_over = global_after > global_pending_limit
|
||||
user_over = bool(user_id and user_after > user_pending_limit)
|
||||
|
||||
if global_over or user_over:
|
||||
if global_over:
|
||||
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
|
||||
exc = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
|
||||
else:
|
||||
reason = f"用户 pending 超限(入队后): {user_after}/{user_pending_limit}"
|
||||
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_after, limit=user_pending_limit)
|
||||
# Bug B #2098: 用户超限仅日志警告,不回滚任务
|
||||
if user_id and user_after > user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 超软上限(入队后): user_id=%s, count=%d/%d",
|
||||
user_id,
|
||||
user_after,
|
||||
user_pending_limit,
|
||||
)
|
||||
|
||||
if global_over:
|
||||
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
|
||||
exc = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
|
||||
logger.warning(
|
||||
"[队列限流] %s, task_id=%s, user_id=%s — 回滚状态为 failed",
|
||||
reason,
|
||||
|
||||
+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):
|
||||
|
||||
Executable
+156
@@ -0,0 +1,156 @@
|
||||
"""爆款视频 API schemas。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
# ── 枚举常量 ─────────────────────────────────────────────────────────────
|
||||
|
||||
VALID_FUSION_LEVELS = ("ai_full", "ai_polish", "user_primary")
|
||||
VALID_STYLE_STRENGTHS = ("light", "medium", "strict")
|
||||
VALID_STAGES = (
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"review",
|
||||
"tts",
|
||||
"bgm_select",
|
||||
"rendering",
|
||||
"musetalk",
|
||||
"uploading",
|
||||
)
|
||||
|
||||
|
||||
# ── Request Schemas ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class CreateViralVideoRequest(BaseModel):
|
||||
"""创建爆款视频任务请求。"""
|
||||
|
||||
images: list[str] = Field(..., min_length=1, max_length=20, description="产品图片 URL 列表")
|
||||
industry: str = Field(default="", description="行业")
|
||||
target_customer: str = Field(default="", description="目标客户描述")
|
||||
persona_id: str = Field(default="", description="人设 ID")
|
||||
viral_structure: str = Field(default="", description="爆款结构类型")
|
||||
marketing_purpose: str = Field(default="", description="营销目的")
|
||||
bgm_preference: str = Field(default="", description="BGM 偏好")
|
||||
duration: int = Field(default=30, ge=5, le=180, description="视频时长(秒)")
|
||||
user_copy_text: str = Field(default="", description="用户原始文案(我说你写)")
|
||||
fusion_level: str = Field(default="ai_polish", description="文案融合级别: ai_full/ai_polish/user_primary")
|
||||
reference_audio_path: str = Field(default="", description="参考音频路径")
|
||||
# v1.3 新增
|
||||
reference_video_url: str = Field(default="", description="参考爆款视频 URL")
|
||||
style_strength: str = Field(default="medium", description="风格强度: light/medium/strict")
|
||||
style_template_id: str = Field(default="", description="风格模板 ID")
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
def _validate_fusion_level(cls, v: str) -> str:
|
||||
if v not in VALID_FUSION_LEVELS:
|
||||
raise ValueError(f"fusion_level 必须是 {VALID_FUSION_LEVELS} 之一")
|
||||
return v
|
||||
|
||||
@field_validator("style_strength")
|
||||
@classmethod
|
||||
def _validate_style_strength(cls, v: str) -> str:
|
||||
if v not in VALID_STYLE_STRENGTHS:
|
||||
raise ValueError(f"style_strength 必须是 {VALID_STYLE_STRENGTHS} 之一")
|
||||
return v
|
||||
|
||||
|
||||
class ConfirmIntentRequest(BaseModel):
|
||||
"""确认意图请求(confirm-intent)。"""
|
||||
|
||||
confirmed_copy: str = Field(default="", description="用户确认/修改后的文案,为空表示使用 AI 生成的文案")
|
||||
adjustments: str = Field(default="", description="用户对 AI 文案的调整意见")
|
||||
|
||||
|
||||
class AnalyzeStyleRequest(BaseModel):
|
||||
"""触发参考视频风格分析请求。"""
|
||||
|
||||
reference_video_url: str = Field(..., description="参考视频 URL")
|
||||
style_template_id: str = Field(default="", description="风格模板 ID(可选覆盖)")
|
||||
|
||||
|
||||
# ── Response Schemas ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ViralVideoJobResponse(BaseModel):
|
||||
"""爆款视频任务响应。"""
|
||||
|
||||
id: str
|
||||
user_id: str
|
||||
images: list[str] = Field(default_factory=list)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
duration: int = 30
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = "ai_polish"
|
||||
reference_audio_path: str = ""
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = "medium"
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
status: str
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class ViralVideoHistoryResponse(BaseModel):
|
||||
"""历史记录列表响应。"""
|
||||
|
||||
items: list[ViralVideoJobResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class StyleTemplateResponse(BaseModel):
|
||||
"""风格模板响应。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
thumbnail_url: str = ""
|
||||
style_config: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class StyleTemplateListResponse(BaseModel):
|
||||
"""风格模板列表响应。"""
|
||||
|
||||
items: list[StyleTemplateResponse]
|
||||
|
||||
|
||||
class AnalyzeStyleResponse(BaseModel):
|
||||
"""风格分析结果响应。"""
|
||||
|
||||
job_id: str
|
||||
status: str
|
||||
style_guide: dict | None = None
|
||||
|
||||
|
||||
# ── WebSocket 事件 Schema ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class WSProgressEvent(BaseModel):
|
||||
"""WebSocket 进度推送事件。"""
|
||||
|
||||
type: str = "viral_video:progress"
|
||||
job_id: str
|
||||
stage: str
|
||||
progress: float = Field(ge=0.0, le=100.0)
|
||||
message: str = ""
|
||||
data: dict = Field(default_factory=dict)
|
||||
@@ -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 = "full_ai" | "polish" | "as_is"
|
||||
export const FUSION_LEVELS: { value: FusionLevel; label: string; desc: string }[] = [
|
||||
{ value: "full_ai", label: "AI 全写", desc: "给我方向,全由AI创作" },
|
||||
{ value: "polish", label: "AI润色", desc: "我写草稿,AI帮我润色" },
|
||||
{ value: "as_is", label: "按我写的来", desc: "几乎不改我的文案" },
|
||||
]
|
||||
|
||||
export type StyleStrength = "light" | "medium" | "strict"
|
||||
export const STYLE_STRENGTHS: { value: StyleStrength; label: string }[] = [
|
||||
{ value: "light", label: "轻度借鉴" },
|
||||
{ value: "medium", label: "中度参考" },
|
||||
{ value: "strict", label: "深度模仿" },
|
||||
]
|
||||
|
||||
export type ViralVideoStatus =
|
||||
"pending" | "running" | "wait_user_confirm" | "completed" | "failed" | "cancelled"
|
||||
|
||||
/**
|
||||
* 后端流水线阶段字符串。前端不展示逐阶段进度列表,仅保留类型
|
||||
* 用于轮询时判断当前在哪个大阶段(分析中 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),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")),
|
||||
|
||||
@@ -30,11 +30,16 @@ EDGE_CROP_MAX_PCT = 0.05
|
||||
# 均以 720p 为基准(见前端 titleCanvas.ts 注释 scale=videoWidth/720,types.ts "px @720p"),
|
||||
# 非 720p 输出时按 video_width / TITLE_SIZE_REF_WIDTH 等比缩放,保证成片位置与前端预览一致。
|
||||
TITLE_SIZE_REF_WIDTH = 720
|
||||
# 与前端 titleCanvas.ts 对齐:top 时文本 top-edge 距视频顶 = PAD(16@720p) + margin_top(默认 24@720p)
|
||||
TITLE_PAD_TOP = 16
|
||||
TITLE_DEFAULT_MARGIN_TOP = 24
|
||||
TITLE_DEFAULT_MARGIN_BOTTOM = 24
|
||||
SUBTITLE_DEFAULT_MARGIN_BOTTOM = 60 # 字幕距底边距(720p 基准,与前端字幕面板默认对齐)
|
||||
# 与 video_filter_builder.build_title_drawtext_filter(CPU 路径)和 ass_subtitle_builder 对齐:
|
||||
# - top/bottom 默认 margin 50@720p(vfb 用 _scale_title_len(50, w),即 y=50 / y=h-th-50)
|
||||
# - margin_top 字段:前端编辑器 marginTop 滑块,叠加在默认 margin 之上(#2095 支持)
|
||||
# - PAD 概念仅用于前端 Canvas 预览;ffmpeg drawtext y 是 baseline,无 font metrics 可用,
|
||||
# 直接用统一 50@720p baseline 位置即可保持三端(GPU/CPU/前端视觉)一致。
|
||||
TITLE_DEFAULT_MARGIN_TOP = 50 # top 位置 baseline 默认距顶 50@720p(与 vfb/CPU 路径一致)
|
||||
TITLE_DEFAULT_MARGIN_BOTTOM = 50 # bottom 位置 baseline 默认距底 50@720p
|
||||
SUBTITLE_DEFAULT_MARGIN_BOTTOM = 50 # 字幕距底边距 50@720p(与 vfb 一致)
|
||||
TITLE_MARGIN_TOP_FROM_CFG_DEFAULT = 24 # 前端 marginTop 滑块默认值(用户未传时叠加 0)
|
||||
TITLE_FAUX_BOLD_WIDTH = 2 # 仿粗黑色描边宽度(与 vfb 一致,2@720p 黑色细描边)
|
||||
|
||||
|
||||
def _scale_title_len(value, video_width: int):
|
||||
@@ -507,7 +512,7 @@ def build_direct_render(
|
||||
except (TypeError, ValueError):
|
||||
t_size_720 = 0
|
||||
if t_size_720 <= 0:
|
||||
t_size_720 = 36 # 与 config_schemas 默认 48 接近;AI Avatar 前端默认 48,剪辑前端默认 28,取中
|
||||
t_size_720 = 48 # 与 config_schemas.DEFAULT_EDIT_PLAN_CONFIG.title.size=48 及 vfb 默认一致
|
||||
t_size = _scale_title_len(t_size_720, output_width)
|
||||
# stroke/shadow 长度字段也需 720p→输出分辨率缩放
|
||||
t_color = str(t_cfg.get("color", t_cfg.get("font_color", "#ffffff")))
|
||||
@@ -525,14 +530,22 @@ def build_direct_render(
|
||||
# baseline 需再下移约 0.85*fontsize;但 drawtext 表达式无法引用 fontsize 变量,
|
||||
# 这里直接用 (PAD + margin_top)@720p 缩放后作为 y(即让 baseline≈顶部内边距位置),
|
||||
# 实际中文字符会自然向下延伸,视觉位置与前端预览(textBaseline=middle 居中到 firstLineY)一致。
|
||||
_t_margin_top_raw = t_cfg.get("margin_top", TITLE_DEFAULT_MARGIN_TOP)
|
||||
# margin_top:前端滑块值(默认 24@720p),叠加在默认 50@720p 基线之上
|
||||
_t_user_margin_top = t_cfg.get("margin_top")
|
||||
try:
|
||||
_t_margin_top_720 = int(_t_margin_top_raw)
|
||||
_t_user_margin_top_720 = int(_t_user_margin_top) if _t_user_margin_top is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
_t_margin_top_720 = TITLE_DEFAULT_MARGIN_TOP
|
||||
t_margin_top = _scale_title_len(TITLE_PAD_TOP + _t_margin_top_720, output_width)
|
||||
# bottom margin(标题放在 bottom 时也支持)
|
||||
t_margin_bottom = _scale_title_len(TITLE_PAD_TOP + TITLE_DEFAULT_MARGIN_BOTTOM, output_width)
|
||||
_t_user_margin_top_720 = 0
|
||||
t_margin_top_720 = TITLE_DEFAULT_MARGIN_TOP + _t_user_margin_top_720
|
||||
t_margin_top = _scale_title_len(t_margin_top_720, output_width)
|
||||
# bottom margin(标题放在 bottom 时):用户 margin_bottom 透传,默认 50@720p
|
||||
_t_user_margin_bottom = t_cfg.get("margin_bottom")
|
||||
try:
|
||||
_t_user_margin_bottom_720 = int(_t_user_margin_bottom) if _t_user_margin_bottom is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
_t_user_margin_bottom_720 = 0
|
||||
t_margin_bottom_720 = TITLE_DEFAULT_MARGIN_BOTTOM + _t_user_margin_bottom_720
|
||||
t_margin_bottom = _scale_title_len(t_margin_bottom_720, output_width)
|
||||
t_borderw = 0
|
||||
t_border_color = "#000000"
|
||||
t_box = False
|
||||
@@ -566,11 +579,11 @@ def build_direct_render(
|
||||
t_shadow_x = _scale_title_len(t_shadow_x_720, output_width)
|
||||
t_shadow_y = _scale_title_len(t_shadow_y_720, output_width)
|
||||
# bold/italic:drawtext 原生无粗斜体选项;通过同色描边模拟粗体
|
||||
t_bold = bool(t_cfg.get("bold", True)) # 与 ASS 路径默认 bold=True 对齐
|
||||
t_bold = bool(t_cfg.get("bold", True)) # 与 ASS/vfb 路径默认 bold=True 对齐
|
||||
if t_bold and t_borderw < 1:
|
||||
# 同色描边宽度 1@720p,按宽度缩放保证各分辨率视觉一致
|
||||
t_borderw = max(1, _scale_title_len(1, output_width))
|
||||
t_border_color = t_color # 用文字色描边模拟加粗
|
||||
# 粗体未配用户描边时:黑色细描边 2@720p(与 vfb 一致,避免同色描边导致重影)
|
||||
t_borderw = _scale_title_len(TITLE_FAUX_BOLD_WIDTH, output_width)
|
||||
t_border_color = "#000000" # 黑色细描边模拟粗体
|
||||
|
||||
# ── 解析 subtitle_config ──
|
||||
s_cfg = dict(subtitle_config) if isinstance(subtitle_config, dict) else {}
|
||||
@@ -595,8 +608,21 @@ def build_direct_render(
|
||||
s_pos_x, s_pos_y = None, None
|
||||
s_margin_top = _scale_title_len(60, output_width) # subtitle top (not commonly used)
|
||||
s_margin_bottom = _scale_title_len(SUBTITLE_DEFAULT_MARGIN_BOTTOM, output_width)
|
||||
s_borderw = max(1, _scale_title_len(2, output_width)) # 字幕默认描边保证可读性
|
||||
# subtitle stroke/bold:先解析用户 stroke,再按 bold 默认补描边
|
||||
s_borderw = 0
|
||||
s_border_color = "#000000"
|
||||
_s_stroke = s_cfg.get("stroke")
|
||||
if isinstance(_s_stroke, dict) and _s_stroke.get("enabled", False):
|
||||
try:
|
||||
s_borderw = _scale_title_len(int(float(_s_stroke.get("width", 2))), output_width)
|
||||
except (TypeError, ValueError):
|
||||
s_borderw = 0
|
||||
s_border_color = str(_s_stroke.get("color", "#000000"))
|
||||
s_bold = bool(s_cfg.get("bold", False))
|
||||
if s_bold and s_borderw < 1:
|
||||
# 粗体默认黑色细描边 2@720p(与 title/CPU vfb 一致)
|
||||
s_borderw = _scale_title_len(TITLE_FAUX_BOLD_WIDTH, output_width)
|
||||
s_border_color = "#000000"
|
||||
|
||||
# 静态字幕:static_subtitle_text 非空时构造全片长 segment(0 → total_duration)
|
||||
static_text = (static_subtitle_text or "").strip()
|
||||
|
||||
@@ -132,6 +132,7 @@ class RenderAdapter:
|
||||
work_dir: Path | None = None,
|
||||
progress_cb: ProgressCallback | None = None,
|
||||
voiceover_audio_path: str | None = None,
|
||||
task_config_override: dict | None = None, # Bug A: task 级 config 覆盖,防并发竞态
|
||||
) -> RenderAdapterResult:
|
||||
"""渲染一个 EditPlan。
|
||||
|
||||
@@ -217,6 +218,7 @@ class RenderAdapter:
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
failed_clip_ids=failed_clip_ids,
|
||||
voiceover_audio_path=voiceover_audio_path,
|
||||
task_config_override=task_config_override,
|
||||
)
|
||||
# 成功时将临时目录所有权转移给调用方,阻止 finally 清理
|
||||
if result.success and temp_dir:
|
||||
@@ -392,7 +394,7 @@ class RenderAdapter:
|
||||
|
||||
return asset_path_map, rendered_clip_ids, failed_clip_ids, asset_storage_map
|
||||
|
||||
def _prepare_bgm(self, plan, work_dir: Path, plan_id: str) -> str | None:
|
||||
def _prepare_bgm(self, plan, work_dir: Path, plan_id: str, *, bgm_override: dict | None = None) -> str | None:
|
||||
"""准备 BGM 音频文件(从 plan.config.bgm 读取配置)。
|
||||
|
||||
支持 3 种来源(按优先级):
|
||||
@@ -405,7 +407,9 @@ class RenderAdapter:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
plan_config = plan.config or {}
|
||||
bgm_config = plan_config.get("bgm", {}) or {}
|
||||
bgm_config = dict(plan_config.get("bgm", {}) or {})
|
||||
if isinstance(bgm_override, dict) and bgm_override:
|
||||
bgm_config.update(bgm_override) # Bug A: 任务级 BGM 覆盖,防并发竞态
|
||||
|
||||
if not bgm_config.get("enabled", False):
|
||||
return None
|
||||
@@ -567,6 +571,7 @@ class RenderAdapter:
|
||||
failed_clip_ids: list[str] | None = None,
|
||||
voiceover_audio_path: str | None = None,
|
||||
asset_storage_map: dict[str, str] | None = None,
|
||||
task_config_override: dict | None = None, # Bug A: task 级 config 覆盖,防并发竞态
|
||||
) -> RenderAdapterResult:
|
||||
"""执行统一渲染核心流程(BGM + ASR + 渲染 + 缩略图 + 上传)。
|
||||
|
||||
@@ -580,8 +585,9 @@ class RenderAdapter:
|
||||
Returns:
|
||||
RenderAdapterResult
|
||||
"""
|
||||
# 1. 准备 BGM
|
||||
bgm_path = self._prepare_bgm(plan, work_dir, plan_id)
|
||||
# 1. 准备 BGM(Bug A: 传 task 级 bgm override)
|
||||
_bgm_override = (task_config_override or {}).get("bgm") if isinstance(task_config_override, dict) else None
|
||||
bgm_path = self._prepare_bgm(plan, work_dir, plan_id, bgm_override=_bgm_override)
|
||||
|
||||
self._report_progress(progress_cb, 40.0, "执行视频渲染")
|
||||
|
||||
@@ -589,8 +595,10 @@ class RenderAdapter:
|
||||
plan_config = plan.config or {}
|
||||
asr_service = self._get_asr_service()
|
||||
|
||||
# 3. 读取输出分辨率
|
||||
export_config = plan_config.get("export", {}) or {}
|
||||
# 3. 读取输出分辨率(Bug A: task override 优先)
|
||||
export_config = dict(plan_config.get("export", {}) or {})
|
||||
if isinstance(task_config_override, dict) and isinstance(task_config_override.get("export"), dict):
|
||||
export_config.update(task_config_override["export"])
|
||||
if not isinstance(export_config, dict):
|
||||
export_config = {}
|
||||
output_width, output_height = _parse_resolution(export_config.get("resolution"))
|
||||
@@ -615,6 +623,7 @@ class RenderAdapter:
|
||||
asr_service=asr_service,
|
||||
voiceover_audio_path=voiceover_audio_path,
|
||||
clip_has_text=clip_has_text,
|
||||
override_config=task_config_override,
|
||||
)
|
||||
# 注入每个视频段对应素材的 storage_key,供全 GPU 直连管线直接签名下载
|
||||
_storage_map = asset_storage_map or {}
|
||||
|
||||
@@ -158,6 +158,7 @@ class UnifiedRenderService:
|
||||
bgm_path: str | None = None, # BGM 本地文件路径
|
||||
voiceover_audio_path: str | None = None, # 配音素材库音频本地路径
|
||||
clip_has_text: list[bool] | None = None, # 源视频片段是否有文字(来自 atom_clip.ai_tags.has_text)
|
||||
override_config: dict | None = None, # Bug A: task 级 config 覆盖(title/bgm/export/subtitle),防并发竞态
|
||||
):
|
||||
self.plan = plan
|
||||
self.clips = clips
|
||||
@@ -170,6 +171,9 @@ class UnifiedRenderService:
|
||||
self.asr_service = asr_service
|
||||
self.bgm_path = bgm_path
|
||||
self.voiceover_audio_path = voiceover_audio_path
|
||||
# Bug A: task 级 config override(深拷贝),优先级高于 plan.config;
|
||||
# 避免同 plan 多任务并发渲染时 _sync_task_config_to_plan 写 plan.config["title"] 互相覆盖。
|
||||
self._override_config = dict(override_config) if isinstance(override_config, dict) else {}
|
||||
# #1970:片段级文字检测(顺序与非 audio 的源视频片段一致);None 表示无可靠检测,保守不翻转
|
||||
self._clip_has_text = clip_has_text
|
||||
self._transition_engine = TransitionEngine(default_duration=transition_duration)
|
||||
@@ -180,6 +184,28 @@ class UnifiedRenderService:
|
||||
self._micro_plan_cache: Any = None
|
||||
self._micro_plan_loaded = False
|
||||
|
||||
def _cfg_section(self, section: str) -> dict:
|
||||
"""读取单个配置段:override_config 优先于 plan.config(Bug A 防并发竞态)。"""
|
||||
base = dict((self.plan.config or {}).get(section, {}) or {})
|
||||
override = self._override_config.get(section)
|
||||
if isinstance(override, dict) and override:
|
||||
base.update(override) # 浅合并,保留 base 中未被覆盖字段
|
||||
return base
|
||||
|
||||
def _effective_config(self) -> dict:
|
||||
"""读取完整 config:override_config 顶层段覆盖 plan.config(Bug A 防并发竞态)。"""
|
||||
import copy
|
||||
|
||||
full = copy.deepcopy(self.plan.config or {})
|
||||
for k, v in self._override_config.items():
|
||||
if isinstance(v, dict):
|
||||
sec = dict(full.get(k, {}) or {})
|
||||
sec.update(v)
|
||||
full[k] = sec
|
||||
else:
|
||||
full[k] = v
|
||||
return full
|
||||
|
||||
# ── #1970 PR2 智能降重:片段级微变换 ───────────────────────────────────
|
||||
def _dedup_enabled(self) -> bool:
|
||||
"""读取 plan.config.dedup_enabled,缺省视为 True(向后兼容)。"""
|
||||
@@ -431,7 +457,7 @@ class UnifiedRenderService:
|
||||
has_audio = pass_through_has_audio
|
||||
# 直通模式下也支持 BGM 混音:提取音频 → 混 BGM → 合并回视频
|
||||
if self.bgm_path and pass_through_has_audio:
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
bgm_config = config.get("bgm", {}) or {}
|
||||
if bgm_config.get("enabled", False):
|
||||
ctx = RenderContext(work_dir=self.work_dir, plan_id=self.plan.id)
|
||||
@@ -471,7 +497,7 @@ class UnifiedRenderService:
|
||||
"[unified-render] pass-through BGM mix failed, skipping: plan_id=%s", self.plan.id
|
||||
)
|
||||
else:
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
bgm_config = config.get("bgm", {}) or {}
|
||||
if not isinstance(bgm_config, dict):
|
||||
bgm_config = {}
|
||||
@@ -696,7 +722,7 @@ class UnifiedRenderService:
|
||||
Returns:
|
||||
ASS 文件路径,没有字幕时返回 None
|
||||
"""
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
# #1901 统一读 "title",兼容老数据 "title_config"
|
||||
title_cfg = config.get("title", {}) or {}
|
||||
if not isinstance(title_cfg, dict) or not (title_cfg.get("text") or "").strip():
|
||||
@@ -881,7 +907,7 @@ class UnifiedRenderService:
|
||||
Returns:
|
||||
是否成功添加了配音音轨
|
||||
"""
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
tts_cfg = config.get("tts", {}) or {}
|
||||
if not isinstance(tts_cfg, dict):
|
||||
tts_cfg = {}
|
||||
@@ -2275,7 +2301,7 @@ class UnifiedRenderService:
|
||||
try:
|
||||
from video_processing import gpu_direct_pipeline as gdp
|
||||
|
||||
cfg = self.plan.config or {}
|
||||
cfg = self._effective_config()
|
||||
video_layer = next(_lyr for _lyr in layers if _lyr.role not in ("audio",))
|
||||
video_clips = [c for c in video_layer.clips if c.clip_type != "audio"]
|
||||
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""爆款视频 Worker 侧模块(#2039/#2040/#2051)。
|
||||
|
||||
video_analyzer(#2051):参考视频风格分析 6 步管线,输出 style_guide + clips 渲染参数映射。
|
||||
#2040 的 prompt 系统(prompts/prompt_store/llm_runner)由 #2040 分支提供,本文件不依赖它。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import (
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_ANALYSIS_TIMEOUT",
|
||||
"MAX_REFERENCE_DURATION_SEC",
|
||||
"MAX_REFERENCE_SIZE_MB",
|
||||
"STYLE_GUIDE_SCHEMA",
|
||||
"analyze_video_style",
|
||||
"build_render_params_for_clip",
|
||||
"map_bgm_bpm",
|
||||
"map_camera_to_ken_burns",
|
||||
"map_color_to_video_filter",
|
||||
"map_transition_to_xfade",
|
||||
]
|
||||
@@ -0,0 +1,961 @@
|
||||
"""参考爆款视频风格分析模块(#2051,v1.3)。
|
||||
|
||||
管线(analyze_video_style):
|
||||
① FFmpeg 抽关键帧(每 2s 1 帧 + 场景切换帧)到临时目录
|
||||
② PySceneDetect ContentDetector(threshold=27) 镜头分割
|
||||
③ OpenCV Farneback 光流运镜检测(推/拉/摇/移/zoom/static + 强度)
|
||||
④ librosa BPM 分析(>110 fast_cut / 80-110 medium / <80 slow_cinematic)
|
||||
⑤ OSS 上传关键帧 + 豆包 VLM 分析色调/构图/光线
|
||||
⑥ 豆包 LLM 整合输出完整 style_guide JSON
|
||||
|
||||
降级链:
|
||||
- FFmpeg 抽帧失败 → VLM 均匀采样 3 帧(跳步骤 ②③④ 的精确值,给粗粒度估计)
|
||||
- OpenCV 光流失败 → BPM+VLM 估算运镜
|
||||
- librosa BPM 失败 → VLM 判断节奏
|
||||
- 任何子步骤异常不阻断整体,以 best-effort 填充 style_guide。
|
||||
|
||||
资源约束:
|
||||
- 参考视频 ≤60s 且 ≤100MB;分析总超时 ≤60s;临时帧 try/finally 清理。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
import tempfile
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 资源约束 ─────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_REFERENCE_DURATION_SEC = 60
|
||||
MAX_REFERENCE_SIZE_MB = 100
|
||||
DEFAULT_ANALYSIS_TIMEOUT = 60 # 秒
|
||||
KEYFRAME_INTERVAL_SEC = 2
|
||||
SCENEDETECT_THRESHOLD = 27
|
||||
VLM_SAMPLE_FRAMES = 5 # 上传给 VLM 的关键帧上限
|
||||
FARNEBACK_PARAMS = dict(pyr_scale=0.5, levels=3, winsize=15, iterations=3, poly_n=5, poly_sigma=1.2, flags=0)
|
||||
|
||||
# ── style_guide 输出 schema(最小校验参考,不强制 jsonschema 依赖) ───────
|
||||
|
||||
STYLE_GUIDE_SCHEMA: dict[str, Any] = {
|
||||
"style_name": str,
|
||||
"avg_shot_duration": float,
|
||||
"shot_count": int,
|
||||
"pace": str, # fast_cut | medium | slow_cinematic
|
||||
"bpm": int,
|
||||
"camera_movements": list,
|
||||
"transitions": list,
|
||||
"color_palette": list,
|
||||
"color_tone": str, # warm | cool | high_sat | low_sat | vintage | fresh | dramatic | bright
|
||||
"color_filter": str, # none | warm_vintage | cool_fresh | high_contrast | soft_pastel | dramatic_cinematic
|
||||
"composition": dict,
|
||||
"lighting": str,
|
||||
"mood": str,
|
||||
"visual_keywords": list,
|
||||
"ken_burns_params": dict,
|
||||
"transition_map": dict,
|
||||
"video_filter_eq_params": dict,
|
||||
"ken_burns_direction_hint": str,
|
||||
}
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ShotBoundary:
|
||||
"""一段镜头(帧号区间)。"""
|
||||
|
||||
index: int
|
||||
start_sec: float
|
||||
end_sec: float
|
||||
movement: str = (
|
||||
"static" # push_in | pull_out | pan_left | pan_right | tilt_up | tilt_down | static | zoom_in | zoom_out
|
||||
)
|
||||
intensity: str = "low" # low | medium | high
|
||||
transition: str = "hard_cut" # 到下一个镜头的转场
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnalysisArtifacts:
|
||||
"""中间产物(降级路径用)。"""
|
||||
|
||||
frames_dir: Path
|
||||
frame_paths: list[Path] = field(default_factory=list)
|
||||
shots: list[ShotBoundary] = field(default_factory=list)
|
||||
bpm: int = 0
|
||||
vlm_descriptions: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _strip_code_fence(text: str) -> str:
|
||||
"""移除 markdown 代码块围栏,返回纯文本。"""
|
||||
t = text.strip()
|
||||
for fence in ("```json", "```JSON", "```"):
|
||||
if t.startswith(fence):
|
||||
t = t[len(fence) :].lstrip()
|
||||
if t.endswith("```"):
|
||||
t = t[:-3].rstrip()
|
||||
return t
|
||||
|
||||
|
||||
# ── FFmpeg / ffprobe ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ffmpeg_bin() -> str:
|
||||
return shutil.which("ffmpeg") or "ffmpeg"
|
||||
|
||||
|
||||
def _ffprobe_bin() -> str:
|
||||
return shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
def _probe_duration(video_path: str | Path) -> float:
|
||||
"""用 ffprobe 取视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
_ffprobe_bin(),
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(video_path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
text=True,
|
||||
) # nosec B603
|
||||
return float(out.strip() or 0)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("ffprobe 时长探测失败 %s: %s", video_path, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _extract_keyframes(video_path: Path, out_dir: Path, interval: int = KEYFRAME_INTERVAL_SEC) -> list[Path]:
|
||||
"""按固定间隔抽帧;同时检测场景切换帧(select='gt(scene,...)')。"""
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
# 固定间隔
|
||||
fixed_tpl = str(out_dir / "f_%04d.jpg")
|
||||
cmd_fixed = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
f"fps=1/{interval}",
|
||||
"-q:v",
|
||||
"3",
|
||||
fixed_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_fixed, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
# 场景切换帧(独立命名,scene_ 前缀)
|
||||
scene_tpl = str(out_dir / "scene_%04d.jpg")
|
||||
cmd_scene = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
"select='gt(scene,0.35)',showinfo",
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
"3",
|
||||
scene_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_scene, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
frames = sorted(out_dir.glob("f_*.jpg")) + sorted(out_dir.glob("scene_*.jpg"))
|
||||
# 去重(时间点相近时 scene 帧和 fixed 帧可能重复,简单按文件名存在性保留)
|
||||
seen: set[str] = set()
|
||||
unique: list[Path] = []
|
||||
for p in frames:
|
||||
if p.name not in seen:
|
||||
seen.add(p.name)
|
||||
unique.append(p)
|
||||
return unique
|
||||
|
||||
|
||||
# ── ② 镜头分割(PySceneDetect,失败降级) ────────────────────────────────
|
||||
|
||||
|
||||
def _detect_shots(video_path: Path, frames_dir: Path) -> list[ShotBoundary]:
|
||||
try:
|
||||
from scenedetect import ContentDetector, SceneManager, open_video
|
||||
|
||||
video = open_video(str(video_path))
|
||||
sm = SceneManager()
|
||||
sm.add_detector(ContentDetector(threshold=SCENEDETECT_THRESHOLD))
|
||||
sm.detect_scenes(video)
|
||||
scenes = sm.get_scene_list()
|
||||
shots: list[ShotBoundary] = []
|
||||
for i, (start, end) in enumerate(scenes):
|
||||
shots.append(
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=start.get_seconds(),
|
||||
end_sec=end.get_seconds(),
|
||||
)
|
||||
)
|
||||
if shots:
|
||||
return shots
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("PySceneDetect 镜头分割失败,使用均匀分段降级: %s", exc)
|
||||
# 降级:按固定间隔每 3 秒一镜头
|
||||
duration = _probe_duration(video_path) or 15.0
|
||||
dur = max(3.0, min(duration, float(MAX_REFERENCE_DURATION_SEC)))
|
||||
shots = []
|
||||
seg = 3.0
|
||||
i = 0
|
||||
t = 0.0
|
||||
while t < dur - 0.1:
|
||||
shots.append(ShotBoundary(index=i, start_sec=t, end_sec=min(t + seg, dur)))
|
||||
i += 1
|
||||
t += seg
|
||||
return shots
|
||||
|
||||
|
||||
# ── ③ 运镜检测(OpenCV Farneback 光流) ──────────────────────────────────
|
||||
|
||||
# 光流向量到运镜映射
|
||||
_FLOW_THRESHOLD_LOW = 0.3
|
||||
_FLOW_THRESHOLD_HIGH = 1.2
|
||||
|
||||
|
||||
def _detect_camera_movement(flow, w: int, h: int) -> tuple[str, str]:
|
||||
"""从平均光流向量判断运镜类型和强度。"""
|
||||
import numpy as np # noqa: PLC0415 - numpy 已在 requirements 中
|
||||
|
||||
fx = float(np.median(flow[..., 0]))
|
||||
fy = float(np.median(flow[..., 1]))
|
||||
trans_mag = math.hypot(fx, fy)
|
||||
# 发散/收敛判断 zoom:比较边缘流沿径向外指的平均分量(稳健版)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
rx, ry = (xs - cx) / max(cx, 1.0), (ys - cy) / max(cy, 1.0)
|
||||
rmag = np.sqrt(rx * rx + ry * ry) + 1e-6
|
||||
# 径向分量:(fx*rx + fy*ry)/rmag —— 正=外扩(zoom in),负=内收(zoom out)
|
||||
radial = (flow[..., 0] * rx + flow[..., 1] * ry) / rmag
|
||||
# 只看边缘带(|r|>0.5),且减去平移贡献:径向减去平均平移投影
|
||||
edge_mask = (rmag > 0.5).astype(np.float32)
|
||||
if edge_mask.sum() > 10:
|
||||
trans_radial = (fx * rx + fy * ry) / rmag
|
||||
zoom_signal = float(np.mean((radial - trans_radial)[edge_mask > 0]))
|
||||
else:
|
||||
zoom_signal = 0.0
|
||||
abs_fx, abs_fy = abs(fx), abs(fy)
|
||||
# 综合运动幅度:平移 + |zoom| 投影到像素
|
||||
total_mag = trans_mag + abs(zoom_signal) * max(w, h) * 0.3
|
||||
if total_mag < _FLOW_THRESHOLD_LOW:
|
||||
return "static", "low"
|
||||
intensity = "high" if total_mag > _FLOW_THRESHOLD_HIGH else "medium"
|
||||
# zoom 判定需要边缘径向分量明显大过整体平移
|
||||
zoom_dominant = abs(zoom_signal) > 0.6 and abs(zoom_signal) * max(w, h) * 0.3 > trans_mag * 1.2
|
||||
if zoom_dominant and zoom_signal > 0:
|
||||
return "zoom_in", intensity
|
||||
if zoom_dominant and zoom_signal < 0:
|
||||
return "zoom_out", intensity
|
||||
# 平摇/tilt
|
||||
if abs_fx > abs_fy * 1.5:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
if abs_fy > abs_fx * 1.5:
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
# 轨道/跟拍:以主轴为主
|
||||
if abs_fx >= abs_fy:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
|
||||
|
||||
def _analyze_movements(video_path: Path, shots: list[ShotBoundary]) -> None:
|
||||
"""对每个 shot 的首尾帧算光流,填充 movement/intensity。失败时静默降级为 static/low。"""
|
||||
try:
|
||||
import cv2 # noqa: PLC0415 - opencv-python-headless 已在 worker requirements 中
|
||||
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OpenCV 不可用,运镜检测降级为 static/low: %s", exc)
|
||||
return
|
||||
try:
|
||||
cap = cv2.VideoCapture(str(video_path))
|
||||
for shot in shots:
|
||||
mid_t = (shot.start_sec + shot.end_sec) / 2.0
|
||||
dt = max(0.2, min(0.5, (shot.end_sec - shot.start_sec) / 4.0))
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, max(0.0, (mid_t - dt)) * 1000)
|
||||
ok1, f1 = cap.read()
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, min(mid_t + dt, shot.end_sec - 0.05) * 1000)
|
||||
ok2, f2 = cap.read()
|
||||
if not (ok1 and ok2):
|
||||
continue
|
||||
g1 = cv2.cvtColor(f1, cv2.COLOR_BGR2GRAY)
|
||||
g2 = cv2.cvtColor(f2, cv2.COLOR_BGR2GRAY)
|
||||
h, w = g1.shape
|
||||
# 降采样加速
|
||||
scale = 360.0 / h if h > 360 else 1.0
|
||||
if scale < 1.0:
|
||||
g1 = cv2.resize(g1, (int(w * scale), int(h * scale)))
|
||||
g2 = cv2.resize(g2, (int(w * scale), int(h * scale)))
|
||||
flow = cv2.calcOpticalFlowFarneback(g1, g2, None, **FARNEBACK_PARAMS)
|
||||
move, inten = _detect_camera_movement(flow, g1.shape[1], g1.shape[0])
|
||||
shot.movement = move
|
||||
shot.intensity = inten
|
||||
cap.release()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("运镜检测异常,已降级: %s", exc)
|
||||
|
||||
|
||||
# ── ④ librosa BPM ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _detect_bpm(video_path: Path) -> int:
|
||||
"""提取音轨并估算 BPM;失败返回 0。"""
|
||||
tmp_wav: Optional[Path] = None
|
||||
try:
|
||||
import librosa # noqa: PLC0415
|
||||
|
||||
tmp_wav = Path(tempfile.mkstemp(suffix=".wav")[1])
|
||||
# ffmpeg 抽 22050Hz 单声道 wav
|
||||
subprocess.run(
|
||||
[_ffmpeg_bin(), "-y", "-i", str(video_path), "-vn", "-ac", "1", "-ar", "22050", "-f", "wav", str(tmp_wav)],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=20,
|
||||
check=False,
|
||||
) # nosec B603
|
||||
if not tmp_wav.exists() or tmp_wav.stat().st_size < 1024:
|
||||
return 0
|
||||
y, sr = librosa.load(str(tmp_wav), sr=22050, mono=True)
|
||||
if len(y) < sr * 2:
|
||||
return 0
|
||||
tempo, _ = librosa.beat.beat_track(y=y, sr=sr)
|
||||
try:
|
||||
bpm = int(round(float(tempo)))
|
||||
except Exception: # noqa: BLE001
|
||||
bpm = int(round(float(tempo[0]))) if len(tempo) else 0
|
||||
return max(40, min(bpm, 220))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("librosa BPM 分析失败: %s", exc)
|
||||
return 0
|
||||
finally:
|
||||
if tmp_wav and tmp_wav.exists():
|
||||
try:
|
||||
tmp_wav.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _pace_from_bpm(bpm: int) -> str:
|
||||
if bpm >= 110:
|
||||
return "fast_cut"
|
||||
if bpm >= 80:
|
||||
return "medium"
|
||||
if bpm > 0:
|
||||
return "slow_cinematic"
|
||||
return "medium"
|
||||
|
||||
|
||||
# ── ⑤ VLM 帧分析 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_frames(frame_paths: list[Path], shots: list[ShotBoundary], k: int = VLM_SAMPLE_FRAMES) -> list[Path]:
|
||||
"""从全量帧中均匀选 k 张代表性帧(优先场景帧)。"""
|
||||
if not frame_paths:
|
||||
return []
|
||||
scene_frames = sorted(p for p in frame_paths if p.name.startswith("scene_"))
|
||||
fixed_frames = sorted(p for p in frame_paths if p.name.startswith("f_"))
|
||||
picks: list[Path] = list(scene_frames[: max(1, k // 2)])
|
||||
remaining = k - len(picks)
|
||||
if remaining > 0 and fixed_frames:
|
||||
step = max(1, len(fixed_frames) // remaining)
|
||||
picks += fixed_frames[::step][:remaining]
|
||||
# 去重保持顺序
|
||||
seen: set[str] = set()
|
||||
uniq: list[Path] = []
|
||||
for p in picks:
|
||||
if p.name not in seen and p.exists():
|
||||
seen.add(p.name)
|
||||
uniq.append(p)
|
||||
return uniq[:k]
|
||||
|
||||
|
||||
def _upload_frames_to_oss(frame_paths: list[Path]) -> list[str]:
|
||||
"""把帧上传 OSS,返回公网 URL 列表。失败时降级为 data URI。"""
|
||||
urls: list[str] = []
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
for p in frame_paths:
|
||||
try:
|
||||
key = f"viral-video/analysis/{uuid.uuid4().hex}/{p.name}"
|
||||
url = upload_to_oss(p, key)
|
||||
if url:
|
||||
urls.append(url)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("单帧 OSS 上传失败 %s: %s", p.name, exc)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OSS 上传模块不可用,降级为 base64 data URI: %s", exc)
|
||||
if len(urls) < len(frame_paths):
|
||||
# 降级:base64 data URI(小图,单张 ≤100KB 才走此路)
|
||||
import base64
|
||||
|
||||
for p in frame_paths[len(urls) :]:
|
||||
try:
|
||||
if p.stat().st_size > 120_000:
|
||||
continue
|
||||
b64 = base64.b64encode(p.read_bytes()).decode("ascii")
|
||||
urls.append(f"data:image/jpeg;base64,{b64}")
|
||||
except Exception: # noqa: BLE001 # nosec B112
|
||||
continue
|
||||
return urls
|
||||
|
||||
|
||||
def _vlm_analyze_frames(image_urls: list[str]) -> dict[str, Any]:
|
||||
"""调豆包 VLM 分析色调/构图/光线/转场观感。"""
|
||||
if not image_urls:
|
||||
return {}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演和调色师。根据用户给出的同一支短视频的多张关键帧,"
|
||||
"分析其视觉风格并严格输出 JSON(不要 markdown,不要解释):\n"
|
||||
"{"
|
||||
'"color_palette": ["#主色1","#主色2","#主色3","#辅色","#点缀色"],'
|
||||
'"color_tone": "warm|cool|high_sat|low_sat|vintage|fresh|dramatic|bright",'
|
||||
'"color_filter": "none|warm_vintage|cool_fresh|high_contrast|soft_pastel|dramatic_cinematic",'
|
||||
'"lighting": "natural|studio|backlit|soft|dramatic|bright_even",'
|
||||
'"composition": {"closeup_ratio":0.0,"medium_ratio":0.0,"wide_ratio":0.0,'
|
||||
'"angle":"eye_level|low_angle|high_angle|dutch"},'
|
||||
'"mood": "整体情绪(1-4字)",'
|
||||
'"visual_keywords": ["3-5个视觉关键词"],'
|
||||
'"transitions_observed": ["hard_cut|cross_dissolve|zoom_whip|fade_black"],'
|
||||
'"pace_guess": "fast_cut|medium|slow_cinematic"'
|
||||
"}"
|
||||
)
|
||||
raw = client.vision_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": "请分析这支参考视频的风格。"},
|
||||
],
|
||||
images=image_urls,
|
||||
temperature=0.2,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return {}
|
||||
raw = _strip_code_fence(raw)
|
||||
# 容忍模型可能前后加文本
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
return json.loads(raw[i : j + 1])
|
||||
return {}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("VLM 帧分析失败: %s", exc)
|
||||
return {}
|
||||
|
||||
|
||||
# ── ⑥ LLM 整合 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _llm_synthesize(
|
||||
shots: list[ShotBoundary],
|
||||
bpm: int,
|
||||
vlm: dict[str, Any],
|
||||
style_strength: str,
|
||||
) -> dict[str, Any]:
|
||||
"""把结构化信号整合成 style_guide;LLM 不可用时走规则合成。"""
|
||||
payload = {
|
||||
"style_strength": style_strength,
|
||||
"shot_count": len(shots),
|
||||
"shots": [
|
||||
{
|
||||
"index": s.index,
|
||||
"start_sec": round(s.start_sec, 2),
|
||||
"end_sec": round(s.end_sec, 2),
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"transition": s.transition,
|
||||
}
|
||||
for s in shots
|
||||
],
|
||||
"bpm": bpm,
|
||||
"pace_guess": _pace_from_bpm(bpm),
|
||||
"vlm": vlm,
|
||||
}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演。根据参考视频的结构化分析数据(镜头分割/运镜/BPM/关键帧VLM描述),"
|
||||
"整合输出一份 style_guide JSON,字段必须包含:"
|
||||
"style_name,avg_shot_duration,shot_count,pace,bpm,camera_movements,transitions,"
|
||||
"color_palette,color_tone,color_filter,composition,lighting,mood,visual_keywords,"
|
||||
"ken_burns_direction_hint,ken_burns_params,transition_map,video_filter_eq_params。"
|
||||
"严格输出一个合法 JSON 对象,不要 markdown/解释。"
|
||||
)
|
||||
user_text = "分析数据:\n" + json.dumps(payload, ensure_ascii=False)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": sys_prompt}, {"role": "user", "content": user_text}],
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
if raw:
|
||||
raw = _strip_code_fence(raw)
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
result = json.loads(raw[i : j + 1])
|
||||
if isinstance(result, dict) and result.get("style_name"):
|
||||
return result
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("LLM 合成 style_guide 失败,走规则降级: %s", exc)
|
||||
return _rule_based_style_guide(shots, bpm, vlm)
|
||||
|
||||
|
||||
def _rule_based_style_guide(shots: list[ShotBoundary], bpm: int, vlm: dict[str, Any]) -> dict[str, Any]:
|
||||
"""LLM 不可用时,用规则拼出可用 style_guide。"""
|
||||
durations = [s.end_sec - s.start_sec for s in shots] or [3.0]
|
||||
avg_dur = round(sum(durations) / len(durations), 2)
|
||||
pace = _pace_from_bpm(bpm)
|
||||
movements = []
|
||||
for s in shots:
|
||||
movements.append(
|
||||
{
|
||||
"shot_index": s.index + 1,
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"duration": round(s.end_sec - s.start_sec, 2),
|
||||
"subject_hint": _default_subject_hint(s.movement),
|
||||
}
|
||||
)
|
||||
transitions = []
|
||||
for i in range(len(shots) - 1):
|
||||
transitions.append({"between_shot": [i + 1, i + 2], "type": shots[i].transition})
|
||||
color_palette = vlm.get("color_palette") or ["#E0E0E0", "#333333", "#F5F5F5", "#888888", "#FF6B35"]
|
||||
color_tone = vlm.get("color_tone") or "bright"
|
||||
color_filter = vlm.get("color_filter") or "none"
|
||||
lighting = vlm.get("lighting") or "bright_even"
|
||||
composition = vlm.get("composition") or {
|
||||
"closeup_ratio": 0.4,
|
||||
"medium_ratio": 0.4,
|
||||
"wide_ratio": 0.2,
|
||||
"angle": "eye_level",
|
||||
}
|
||||
mood = vlm.get("mood") or "明快"
|
||||
vk = vlm.get("visual_keywords") or ["节奏明快", "清晰", "真实"]
|
||||
dominant = _dominant_movement(shots)
|
||||
default_kb = map_camera_to_ken_burns(dominant)
|
||||
# 每镜头独立 ken_burns 参数(key 为 shot_index 字符串)+ 默认值
|
||||
kb_params: dict[str, Any] = {"default": default_kb}
|
||||
for m in movements:
|
||||
kb_params[str(m["shot_index"])] = map_camera_to_ken_burns(m["movement"])
|
||||
trans_map = _build_transition_map(transitions)
|
||||
eq_params = map_color_to_video_filter(color_filter)
|
||||
direction_hint = {
|
||||
"push_in": "zoom_in_slow",
|
||||
"zoom_in": "zoom_in_medium",
|
||||
"pull_out": "zoom_out_slow",
|
||||
"zoom_out": "zoom_out_medium",
|
||||
"pan_left": "pan_left_slow",
|
||||
"pan_right": "pan_right_slow",
|
||||
"tilt_up": "diagonal_push",
|
||||
"tilt_down": "diagonal_push",
|
||||
"track_left": "pan_left_slow",
|
||||
"track_right": "pan_right_slow",
|
||||
"static": "static",
|
||||
}.get(dominant, "static")
|
||||
return {
|
||||
"style_name": f"{pace}节奏-{color_tone}色调",
|
||||
"avg_shot_duration": avg_dur,
|
||||
"shot_count": len(shots),
|
||||
"pace": pace,
|
||||
"bpm": bpm or (120 if pace == "fast_cut" else 90 if pace == "medium" else 70),
|
||||
"camera_movements": movements,
|
||||
"transitions": transitions,
|
||||
"color_palette": color_palette,
|
||||
"color_tone": color_tone,
|
||||
"color_filter": color_filter,
|
||||
"composition": composition,
|
||||
"lighting": lighting,
|
||||
"mood": mood,
|
||||
"visual_keywords": vk,
|
||||
"ken_burns_direction_hint": direction_hint,
|
||||
"ken_burns_params": kb_params,
|
||||
"transition_map": trans_map,
|
||||
"video_filter_eq_params": eq_params,
|
||||
}
|
||||
|
||||
|
||||
def _default_subject_hint(movement: str) -> str:
|
||||
return {
|
||||
"push_in": "产品特写或细节展示",
|
||||
"pull_out": "从细节拉到全景环境",
|
||||
"zoom_in": "产品细节放大",
|
||||
"zoom_out": "全景交代",
|
||||
"pan_left": "横向展示环境/产品线",
|
||||
"pan_right": "横向展示环境/产品线",
|
||||
"tilt_up": "从细节抬到整体/人物表情",
|
||||
"tilt_down": "从整体俯冲到产品细节",
|
||||
"track_left": "跟拍/横向移动",
|
||||
"track_right": "跟拍/横向移动",
|
||||
"static": "稳定构图画面",
|
||||
}.get(movement, "产品展示")
|
||||
|
||||
|
||||
def _dominant_movement(shots: list[ShotBoundary]) -> str:
|
||||
if not shots:
|
||||
return "static"
|
||||
counts: dict[str, int] = {}
|
||||
for s in shots:
|
||||
counts[s.movement] = counts.get(s.movement, 0) + 1
|
||||
return max(counts, key=counts.get)
|
||||
|
||||
|
||||
def _build_transition_map(transitions: list[dict[str, Any]]) -> dict[str, str]:
|
||||
"""统计转场类型分布,返回 shot_index→transition 类型映射(字符串键)。"""
|
||||
m: dict[str, str] = {}
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or [0, 0]
|
||||
if len(pair) >= 2:
|
||||
m[f"{pair[0]}-{pair[1]}"] = t.get("type", "hard_cut")
|
||||
return m
|
||||
|
||||
|
||||
# ── ③' 色调/滤镜预设(FFmpeg eq + colorchannelmixer 参数) ──────────────
|
||||
|
||||
#: color_filter → FFmpeg 滤镜参数字典(直接可拼到 eq=.../colorchannelmixer=...)
|
||||
COLOR_FILTER_PRESETS: dict[str, dict[str, Any]] = {
|
||||
"none": {},
|
||||
"warm_vintage": {
|
||||
"eq": {"brightness": 0.02, "contrast": 1.05, "saturation": 0.9, "gamma": 1.05},
|
||||
"colorchannelmixer": {"rr": 1.1, "gg": 0.98, "bb": 0.82, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"cool_fresh": {
|
||||
"eq": {"brightness": 0.03, "contrast": 1.08, "saturation": 1.05},
|
||||
"colorchannelmixer": {"rr": 0.9, "gg": 1.0, "bb": 1.12, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"high_contrast": {
|
||||
"eq": {"brightness": 0.0, "contrast": 1.3, "saturation": 1.2},
|
||||
"colorchannelmixer": {},
|
||||
},
|
||||
"soft_pastel": {
|
||||
"eq": {"brightness": 0.05, "contrast": 0.92, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 1.03, "bb": 1.05, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"dramatic_cinematic": {
|
||||
"eq": {"brightness": -0.03, "contrast": 1.2, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 0.98, "bb": 0.9, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_color_to_video_filter(color_filter: str) -> dict[str, Any]:
|
||||
"""color_filter 枚举 → FFmpeg eq/colorchannelmixer 参数字典(渲染端直接使用)。"""
|
||||
preset = COLOR_FILTER_PRESETS.get(color_filter) or COLOR_FILTER_PRESETS["none"]
|
||||
# 返回深拷贝防污染
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── ③'' 运镜 → ken_burns 参数映射 ───────────────────────────────────────
|
||||
|
||||
#: 运镜类型 → URS 可直接消费的 ken_burns 参数字典
|
||||
CAMERA_TO_KEN_BURNS: dict[str, dict[str, Any]] = {
|
||||
"static": {
|
||||
"type": "static",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"push_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.12,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.18,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pull_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.12,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.18,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_up": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.08,
|
||||
"zoom_end": 1.14,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": -0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_down": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.14,
|
||||
"zoom_end": 1.08,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_camera_to_ken_burns(movement: str) -> dict[str, Any]:
|
||||
"""运镜类型 → URS ken_burns 参数字典。未知类型回退 static。"""
|
||||
preset = CAMERA_TO_KEN_BURNS.get(movement) or CAMERA_TO_KEN_BURNS["static"]
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── 转场 → xfade transition 名称 ────────────────────────────────────────
|
||||
|
||||
TRANSITION_TO_XFADE: dict[str, str] = {
|
||||
"hard_cut": "cut",
|
||||
"cross_dissolve": "dissolve",
|
||||
"fade_black": "fadeblack",
|
||||
"fade": "fade",
|
||||
"zoom_whip": "zoom",
|
||||
"slide_left": "slideright", # 画面左移 = 新画面从右滑入
|
||||
"slide_right": "slideleft",
|
||||
"wipe_left": "wipeleft",
|
||||
"wipe_right": "wiperight",
|
||||
}
|
||||
|
||||
|
||||
def map_transition_to_xfade(transition_type: str) -> str:
|
||||
"""转场枚举 → TransitionEngine 支持的 xfade 名称;未知回退 cut。"""
|
||||
return TRANSITION_TO_XFADE.get(transition_type, "cut")
|
||||
|
||||
|
||||
# ── BPM → BGM 推荐 BPM ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def map_bgm_bpm(bpm: int) -> int:
|
||||
"""BGM 选曲 BPM:参考视频 BPM ±5。bpm=0 返回 90(默认 medium)。"""
|
||||
if bpm <= 0:
|
||||
return 90
|
||||
return max(60, min(bpm, 180))
|
||||
|
||||
|
||||
# ── 单 clip 渲染参数聚合(给 URS build_render_plan 使用) ───────────────
|
||||
|
||||
|
||||
def build_render_params_for_clip(
|
||||
clip_index: int,
|
||||
style_guide: dict[str, Any],
|
||||
*,
|
||||
duration_sec: Optional[float] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""根据 style_guide 为第 clip_index 个 clip 生成可直接喂给 URS 的渲染参数。"""
|
||||
shot_idx = clip_index + 1
|
||||
movements = style_guide.get("camera_movements") or []
|
||||
movement = "static"
|
||||
intensity = "low"
|
||||
for m in movements:
|
||||
if m.get("shot_index") == shot_idx:
|
||||
movement = m.get("movement", "static")
|
||||
intensity = m.get("intensity", "low")
|
||||
break
|
||||
ken = map_camera_to_ken_burns(movement)
|
||||
if intensity == "high":
|
||||
ken["zoom_end"] = round(ken.get("zoom_end", 1.0) * 1.08, 3)
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 1.3, 3)
|
||||
elif intensity == "low":
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 0.6, 3)
|
||||
transitions = style_guide.get("transitions") or []
|
||||
trans_type = "hard_cut"
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or []
|
||||
if len(pair) >= 2 and pair[0] == shot_idx:
|
||||
trans_type = t.get("type", "hard_cut")
|
||||
break
|
||||
xfade = map_transition_to_xfade(trans_type)
|
||||
eq = map_color_to_video_filter(style_guide.get("color_filter", "none"))
|
||||
return {
|
||||
"ken_burns": ken,
|
||||
"transition": {"type": xfade, "duration": 0.3 if xfade != "cut" else 0.0},
|
||||
"video_filter": eq,
|
||||
"bgm_bpm_hint": map_bgm_bpm(int(style_guide.get("bpm") or 0)),
|
||||
"duration_sec": duration_sec,
|
||||
}
|
||||
|
||||
|
||||
# ── 素材本地化(URL/OSS key → 本地临时文件) ────────────────────────────
|
||||
|
||||
|
||||
def _ensure_local_video(reference: str, work_dir: Path) -> Optional[Path]:
|
||||
"""把 reference(URL/OSS key/本地路径)落到 work_dir 下的本地文件。"""
|
||||
p = Path(reference)
|
||||
if p.exists() and p.is_file():
|
||||
return p
|
||||
try:
|
||||
from video_processing.oss_helpers import download_asset
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
ok = download_asset(reference, target)
|
||||
if ok and target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("download_asset 失败,尝试 http 直连: %s", exc)
|
||||
if reference.startswith(("http://", "https://")):
|
||||
try:
|
||||
import httpx # noqa: PLC0415 - 项目依赖,延迟导入
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
with httpx.Client(timeout=20.0, follow_redirects=True) as client:
|
||||
with client.stream("GET", reference) as resp:
|
||||
resp.raise_for_status()
|
||||
with open(target, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=64 * 1024):
|
||||
f.write(chunk)
|
||||
if target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("HTTP 下载参考视频失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── 入口 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def analyze_video_style(
|
||||
reference_video_path: str | Path,
|
||||
style_strength: str = "medium",
|
||||
*,
|
||||
timeout_sec: int = DEFAULT_ANALYSIS_TIMEOUT,
|
||||
) -> dict[str, Any]:
|
||||
"""分析参考视频风格,返回 style_guide dict。
|
||||
|
||||
Args:
|
||||
reference_video_path: 本地路径、HTTP(S) URL 或 OSS storage key。
|
||||
style_strength: light | medium | strict。
|
||||
timeout_sec: 单步超时(秒),默认 60。
|
||||
|
||||
Returns:
|
||||
style_guide dict,详见 STYLE_GUIDE_SCHEMA。任何子步骤失败都会降级,不抛异常。
|
||||
"""
|
||||
style_strength = style_strength if style_strength in ("light", "medium", "strict") else "medium"
|
||||
frames_dir: Optional[Path] = None
|
||||
local_path: Optional[Path] = None
|
||||
try:
|
||||
frames_dir = Path(tempfile.mkdtemp(prefix="vstyle_"))
|
||||
work_dir = frames_dir # 同一临时根
|
||||
local_path = _ensure_local_video(str(reference_video_path), work_dir)
|
||||
if local_path is None:
|
||||
logger.error("[video_analyzer] 无法获取参考视频: %s", reference_video_path)
|
||||
return _rule_based_style_guide([], 0, {})
|
||||
# 资源约束:大小 / 时长
|
||||
try:
|
||||
size_mb = local_path.stat().st_size / (1024 * 1024)
|
||||
if size_mb > MAX_REFERENCE_SIZE_MB:
|
||||
logger.warning(
|
||||
"[video_analyzer] 参考视频 %.1fMB 超上限,按前 %ds 分析", size_mb, MAX_REFERENCE_DURATION_SEC
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
duration = _probe_duration(local_path)
|
||||
if duration > MAX_REFERENCE_DURATION_SEC:
|
||||
duration = MAX_REFERENCE_DURATION_SEC
|
||||
# ① 抽帧
|
||||
try:
|
||||
frame_paths = _extract_keyframes(local_path, frames_dir / "frames")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("FFmpeg 抽帧失败: %s,降级为 VLM 均匀采样", exc)
|
||||
frame_paths = []
|
||||
# ② 镜头分割
|
||||
shots = _detect_shots(local_path, frames_dir)
|
||||
# ③ 运镜检测(有帧才跑)
|
||||
if frame_paths or shots:
|
||||
_analyze_movements(local_path, shots)
|
||||
# ④ BPM
|
||||
bpm = _detect_bpm(local_path)
|
||||
# ⑤ 选帧→OSS→VLM
|
||||
sampled = _sample_frames(frame_paths, shots)
|
||||
image_urls = _upload_frames_to_oss(sampled) if sampled else []
|
||||
vlm = _vlm_analyze_frames(image_urls) if image_urls else {}
|
||||
# ⑥ 合成
|
||||
style_guide = _llm_synthesize(shots, bpm, vlm, style_strength)
|
||||
# 兜底字段校验
|
||||
style_guide.setdefault("style_strength", style_strength)
|
||||
style_guide.setdefault("pace", _pace_from_bpm(bpm))
|
||||
style_guide.setdefault("bpm", bpm)
|
||||
style_guide.setdefault("shot_count", len(shots))
|
||||
if shots and "avg_shot_duration" not in style_guide:
|
||||
durs = [s.end_sec - s.start_sec for s in shots]
|
||||
style_guide["avg_shot_duration"] = round(sum(durs) / len(durs), 2)
|
||||
return style_guide
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("[video_analyzer] 整体分析异常,返回最小占位 style_guide: %s", exc)
|
||||
return _rule_based_style_guide([], 0, {"mood": "未知"})
|
||||
finally:
|
||||
# 临时帧清理
|
||||
if frames_dir and frames_dir.exists():
|
||||
shutil.rmtree(frames_dir, ignore_errors=True)
|
||||
@@ -38,6 +38,7 @@ celery_app.conf.imports = (
|
||||
"worker_app.tasks.voice_extraction",
|
||||
"worker_app.tasks.voice_clone",
|
||||
"worker_app.tasks.tts_synthesis",
|
||||
"worker_app.tasks.viral_video", # #2039 爆款视频编排器(10步流水线)
|
||||
"worker_app.tasks.batch_download",
|
||||
"worker_app.tasks.duplication_check",
|
||||
# #1798 AI 数字人渲染:必须在 Worker 实例上注册同名任务,否则消息无人消费(渲染卡 0%)
|
||||
|
||||
@@ -616,69 +616,49 @@ def _precompute_render_metadata(
|
||||
# ── Celery Task ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sync_task_config_to_plan(source_edit_plan_id: str, task_info: dict, db) -> str | None:
|
||||
"""将 GenerationTask 的配置同步到 EditPlan.config,返回配音本地路径(如果有)。
|
||||
def _build_task_config_override(task_info: dict) -> dict:
|
||||
"""Bug A: 从 task_info 构建任务级 config override 深拷贝,供渲染时覆盖 plan.config。
|
||||
|
||||
包括:title_config、BGM、输出分辨率。配音单独处理(需下载到本地)。
|
||||
所有渲染相关配置(title/bgm/export)从任务自身读取,不再依赖共享 plan.config,
|
||||
彻底消除同 plan 多任务并发渲染时的竞态覆盖问题。
|
||||
"""
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
|
||||
SQLAlchemyEditPlanRepository,
|
||||
)
|
||||
import copy
|
||||
|
||||
plan_repo = SQLAlchemyEditPlanRepository(db)
|
||||
plan = plan_repo.get(source_edit_plan_id)
|
||||
if plan is None:
|
||||
logger.error("[task] EditPlan not found: %s", source_edit_plan_id)
|
||||
return None
|
||||
override: dict = {}
|
||||
|
||||
plan_config = dict(plan.config or {})
|
||||
changed = False
|
||||
|
||||
# 标题配置
|
||||
# 标题配置(字段名归一化)
|
||||
title_config = task_info.get("title_config") or {}
|
||||
if title_config and isinstance(title_config, dict) and title_config.get("text", "").strip():
|
||||
cfg = dict(title_config)
|
||||
# 字段名归一化
|
||||
if isinstance(title_config, dict) and title_config:
|
||||
cfg = copy.deepcopy(title_config)
|
||||
if "font_size" in cfg and "size" not in cfg:
|
||||
cfg["size"] = cfg["font_size"]
|
||||
if "font_color" in cfg and "color" not in cfg:
|
||||
cfg["color"] = cfg["font_color"]
|
||||
plan_config["title"] = cfg
|
||||
changed = True
|
||||
logger.info("[task] title_config synced to plan: %s", cfg.get("text", "")[:30])
|
||||
override["title"] = cfg
|
||||
|
||||
# BGM 配置
|
||||
bgm_config = task_info.get("bgm_config") or {}
|
||||
if bgm_config:
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
existing_bgm = plan_config.get("bgm", {}) or {}
|
||||
plan_config["bgm"] = merge_bgm_config(existing_bgm, bgm_config)
|
||||
changed = True
|
||||
if isinstance(bgm_config, dict) and bgm_config:
|
||||
override["bgm"] = copy.deepcopy(bgm_config)
|
||||
|
||||
# 输出分辨率
|
||||
ow = task_info.get("output_width") or OUTPUT_WIDTH
|
||||
oh = task_info.get("output_height") or OUTPUT_HEIGHT
|
||||
if ow >= 100 and oh >= 100:
|
||||
export_cfg = dict(plan_config.get("export", {}) or {})
|
||||
export_cfg["resolution"] = f"{ow}x{oh}"
|
||||
plan_config["export"] = export_cfg
|
||||
changed = True
|
||||
override["export"] = {"resolution": f"{ow}x{oh}"}
|
||||
|
||||
if changed:
|
||||
plan.config = plan_config
|
||||
plan_repo.update(plan)
|
||||
logger.info("[task] plan.config synced: plan_id=%s", source_edit_plan_id)
|
||||
return override
|
||||
|
||||
|
||||
def _download_voice_for_task(task_info: dict, source_edit_plan_id: str) -> str | None:
|
||||
"""下载任务配音到本地临时文件,返回路径(不读写 plan.config)。"""
|
||||
import tempfile
|
||||
|
||||
# 配音下载
|
||||
voiceover_path: str | None = None
|
||||
voice_library_id = task_info.get("voice_library_id", "")
|
||||
# #1749:voice_ids 冗余字段已移除;配音一律以 voice_library_id 为准(独立配音每变体各自绑定)
|
||||
effective_voice_id = voice_library_id or ""
|
||||
|
||||
if effective_voice_id:
|
||||
import tempfile
|
||||
|
||||
voice_tmp = Path(tempfile.gettempdir()) / f"voice_{source_edit_plan_id}_{id(task_info)}.mp3"
|
||||
try:
|
||||
if _download_voice_asset(effective_voice_id, voice_tmp):
|
||||
@@ -697,6 +677,9 @@ def _render_from_edit_plan(
|
||||
) -> tuple[Path, float, list[dict] | None, str | None, str | None, str, bool]:
|
||||
"""从 EditPlan 数据库记录直接渲染(不再内存重建clips)。
|
||||
|
||||
Bug A 修复:不再通过 _sync_task_config_to_plan 写共享 plan.config;
|
||||
渲染配置通过 task_config_override 参数直接传入渲染层,彻底消除并发竞态。
|
||||
|
||||
Returns:
|
||||
(output_path, render_duration, cover_candidates, voiceover_path, temp_dir, thumbnail_url, edge_crop_applied)
|
||||
"""
|
||||
@@ -705,8 +688,10 @@ def _render_from_edit_plan(
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
# 同步配置到 plan.config + 下载配音
|
||||
voiceover_path = _sync_task_config_to_plan(source_edit_plan_id, task_info, db)
|
||||
# Bug A: 构建任务级 config override(深拷贝自 task_info),不写 plan.config,避免并发竞态
|
||||
task_override = _build_task_config_override(task_info)
|
||||
# 下载配音到本地临时文件(不依赖 plan.config)
|
||||
voiceover_path = _download_voice_for_task(task_info, source_edit_plan_id)
|
||||
|
||||
# 进度回调
|
||||
def _progress_cb(progress: float, stage: str):
|
||||
@@ -722,6 +707,7 @@ def _render_from_edit_plan(
|
||||
job_id=task_id,
|
||||
progress_cb=_progress_cb,
|
||||
voiceover_audio_path=voiceover_path,
|
||||
task_config_override=task_override,
|
||||
)
|
||||
|
||||
if not result.success:
|
||||
|
||||
@@ -0,0 +1,841 @@
|
||||
"""爆款视频 Celery 编排器 — ViralVideoOrchestrator.
|
||||
|
||||
9 步流水线(Seedance 2.5 直生口型,不再走 MuseTalk):
|
||||
1. 图片 VLM 分析
|
||||
1.5 [v1.3] 视频风格分析(如用户上传参考视频)
|
||||
2. 用户文案意图解析
|
||||
3. 文案融合生成
|
||||
4. 分镜脚本生成
|
||||
5. 合规审核(6 维度,不通过自动重写 1 次)
|
||||
6. CosyVoice 配音
|
||||
7. BGM 选择(素材未就绪时跳过)
|
||||
8. Seedance 逐分镜生成 + ffmpeg concat + 混 TTS
|
||||
9. OSS 上传 + 通知 + 扣点
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
from worker_app.celery_app import celery_app # noqa: F401 - 加载 app 以注册任务
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
ViralVideoJob,
|
||||
ViralVideoStage,
|
||||
ViralVideoStatus,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── WS 进度推送 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _emit_progress(
|
||||
job_id: str,
|
||||
stage: str,
|
||||
progress: float,
|
||||
message: str = "",
|
||||
data: dict | None = None,
|
||||
event_type: str = "viral_video:progress",
|
||||
):
|
||||
"""通过 Redis 发布进度事件,供 WebSocket 消费。
|
||||
|
||||
event_type 取值:
|
||||
- viral_video:progress 中间进度(默认)
|
||||
- viral_video:completed 任务完成
|
||||
- viral_video:failed 任务失败
|
||||
- viral_video:wait_user 等待用户确认
|
||||
所有事件 payload 均为合法 JSON,前端 JSON.parse 即可。
|
||||
"""
|
||||
try:
|
||||
import redis as redis_lib
|
||||
|
||||
redis_url = os.environ.get("REDIS_URL", "redis://localhost:6379/0")
|
||||
r = redis_lib.from_url(redis_url)
|
||||
event = {
|
||||
"type": event_type,
|
||||
"job_id": job_id,
|
||||
"stage": stage,
|
||||
"progress": progress,
|
||||
"message": message or STAGE_LABELS.get(stage, stage),
|
||||
"data": data or {},
|
||||
}
|
||||
r.publish(f"viral_video:{job_id}", json.dumps(event, ensure_ascii=False))
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] WS 进度推送失败: %s", e)
|
||||
|
||||
|
||||
# ── 仓储辅助 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_repo_and_job(job_id: str):
|
||||
"""获取 session, repo, job 三元组。"""
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
return session, repo, job
|
||||
|
||||
|
||||
def _save_job(repo, job, session):
|
||||
"""持久化并关闭 session。"""
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── 流水线各步骤 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _step_image_analysis(job: ViralVideoJob) -> dict:
|
||||
"""步骤 1: 图片 VLM 分析 — 识别产品特征、场景、卖点。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_vision
|
||||
except ImportError:
|
||||
logger.warning("[爆款视频] ai_service.call_vision 不可用,使用占位结果")
|
||||
return {"products": [{"name": "产品", "features": ["特征1", "特征2"], "scene": "通用场景"}]}
|
||||
|
||||
results = []
|
||||
for img_url in job.images:
|
||||
try:
|
||||
result = call_vision(
|
||||
image_url=img_url,
|
||||
prompt="请分析这张产品图片,识别:1)产品名称和类别 2)主要特征和卖点 3)适用场景 4)视觉风格。以JSON格式返回。",
|
||||
)
|
||||
results.append(result)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 图片分析失败 img=%s: %s", img_url, e)
|
||||
results.append({"name": "未识别", "features": [], "scene": "通用"})
|
||||
|
||||
return {"products": results}
|
||||
|
||||
|
||||
def _step_video_analysis(job: ViralVideoJob) -> dict | None:
|
||||
"""步骤 1.5 [v1.3]: 参考视频风格分析。"""
|
||||
if not job.reference_video_url:
|
||||
return None
|
||||
|
||||
try:
|
||||
# P0-2: 修正 import 路径(video_analyzer.py 在 apps/worker/viral_video/ 下,worker PYTHONPATH 含 apps/worker)
|
||||
from viral_video.video_analyzer import analyze_video_style
|
||||
|
||||
style_guide = analyze_video_style(job.reference_video_url)
|
||||
return style_guide
|
||||
except ImportError as e:
|
||||
logger.info("[爆款视频] video_analyzer 模块未就绪(%s),使用占位风格分析", e)
|
||||
return {
|
||||
"cut_speed": "medium",
|
||||
"transition": "cross_dissolve",
|
||||
"energy": "medium",
|
||||
"color_grade": "neutral",
|
||||
"narrative": False,
|
||||
"source": "placeholder",
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 视频风格分析失败: %s", e)
|
||||
return {"error": str(e), "source": "failed"}
|
||||
|
||||
|
||||
def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
"""步骤 2: 用户文案意图解析 — 理解用户想表达什么。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业"}
|
||||
|
||||
products_summary = ""
|
||||
for p in image_analysis.get("products", []):
|
||||
products_summary += f"- {p.get('name', '产品')}: {', '.join(p.get('features', []))}\n"
|
||||
|
||||
prompt = f"""你是一个营销文案策略师。请分析以下信息,理解用户的营销意图:
|
||||
|
||||
用户原始文案:{job.user_copy_text or "(未提供)"}
|
||||
行业:{job.industry or "未指定"}
|
||||
目标客户:{job.target_customer or "未指定"}
|
||||
营销目的:{job.marketing_purpose or "未指定"}
|
||||
产品信息:
|
||||
{products_summary}
|
||||
|
||||
请分析并返回JSON格式:
|
||||
1. intent: 核心营销意图(一句话)
|
||||
2. key_messages: 要传达的3-5个关键信息
|
||||
3. tone: 文案调性(如:专业/亲切/高端/活力)
|
||||
4. target_emotion: 希望触发的用户情感
|
||||
5. call_to_action: 行动号召建议"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, dict) else {"raw": result}
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 意图解析失败: %s", e)
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业"}
|
||||
|
||||
|
||||
def _step_copy_fusion(job: ViralVideoJob, intent: dict, image_analysis: dict) -> str:
|
||||
"""步骤 3: 文案融合生成 — 根据 fusion_level 融合用户文案和 AI 文案。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return f"【{job.industry or '行业'}】优质产品,{job.target_customer or '您'}的不二之选!"
|
||||
|
||||
products_desc = ""
|
||||
for p in image_analysis.get("products", []):
|
||||
products_desc += f"{p.get('name', '产品')}({','.join(p.get('features', []))})\n"
|
||||
|
||||
if job.fusion_level == "ai_full":
|
||||
prompt = f"""请为以下产品撰写一段爆款短视频文案({job.duration}秒):
|
||||
产品:{products_desc}
|
||||
行业:{job.industry}
|
||||
目标客户:{job.target_customer}
|
||||
营销目的:{job.marketing_purpose}
|
||||
调性:{intent.get("tone", "专业")}
|
||||
关键信息:{", ".join(intent.get("key_messages", []))}
|
||||
|
||||
要求:吸引眼球、节奏紧凑、有行动号召。直接输出文案内容。"""
|
||||
elif job.fusion_level == "user_primary":
|
||||
prompt = f"""请基于用户原始文案进行润色优化,保留用户原意和风格:
|
||||
用户原文:{job.user_copy_text}
|
||||
产品信息:{products_desc}
|
||||
|
||||
要求:保留用户原意,仅修正表达和节奏。直接输出文案内容。"""
|
||||
else: # ai_polish (default)
|
||||
prompt = f"""请将用户文案与AI分析融合,生成一段优化后的爆款短视频文案({job.duration}秒):
|
||||
用户原文:{job.user_copy_text or "(未提供)"}
|
||||
产品分析:{products_desc}
|
||||
行业:{job.industry}
|
||||
目标客户:{job.target_customer}
|
||||
营销目的:{job.marketing_purpose}
|
||||
意图分析:{intent.get("intent", "")}
|
||||
调性:{intent.get("tone", "专业")}
|
||||
|
||||
要求:融合用户意图和产品卖点,节奏紧凑,适合短视频。直接输出文案内容。"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, str) else str(result)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 文案融合失败: %s", e)
|
||||
return job.user_copy_text or f"精选{job.industry or '行业'}好物,值得关注!"
|
||||
|
||||
|
||||
def _step_storyboard(job: ViralVideoJob, copy_text: str, image_analysis: dict) -> list[dict]:
|
||||
"""步骤 4: 分镜脚本生成。每个分镜独立一段视频,段内时长建议 3~6 秒。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return [
|
||||
{
|
||||
"order": 0,
|
||||
"type": "product_shot",
|
||||
"text": copy_text[:50],
|
||||
"duration": min(5, job.duration),
|
||||
"description": "产品展示",
|
||||
"ken_burns": "zoom_in",
|
||||
"transition": "cut",
|
||||
}
|
||||
]
|
||||
|
||||
products_hint = ""
|
||||
products = image_analysis.get("products", []) if image_analysis else []
|
||||
if products:
|
||||
p0 = products[0] if isinstance(products[0], dict) else {}
|
||||
feats = p0.get("features", []) if isinstance(p0, dict) else []
|
||||
products_hint = f"\n首帧参考产品特征:{p0.get('name','')} - {', '.join(feats[:3])}"
|
||||
|
||||
seg_seconds = 5
|
||||
n_segments = max(2, min(6, max(1, job.duration // seg_seconds)))
|
||||
ratio = "9:16"
|
||||
|
||||
prompt = f"""请根据以下文案生成爆款短视频分镜脚本,共 {n_segments} 个分镜:
|
||||
|
||||
文案内容:{copy_text}
|
||||
视频总时长:{job.duration}秒(每个分镜 3~6 秒,总和约等于总时长)
|
||||
风格强度:{job.style_strength}
|
||||
输出宽高比:{ratio}{products_hint}
|
||||
|
||||
请以 JSON 数组格式返回分镜列表,每个分镜包含:
|
||||
- order: 序号(从0开始)
|
||||
- type: 镜头类型(product_shot/close_up/scene/action/text_card/closing)
|
||||
- description: 画面详细描述(中文,含主体、动作、场景、运镜、光影,用于AI视频生成prompt)
|
||||
- text: 该分镜配音/字幕文本
|
||||
- duration: 时长(秒,3~6秒的整数)
|
||||
- ken_burns: 运镜方式(zoom_in/zoom_out/pan_left/pan_right/static)
|
||||
- transition: 与下一分镜的转场(cut/dissolve/fade)"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
if isinstance(result, list):
|
||||
return _normalize_storyboard(result, job.duration, n_segments, copy_text)
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(result) if isinstance(result, str) else result
|
||||
if isinstance(parsed, list):
|
||||
return _normalize_storyboard(parsed, job.duration, n_segments, copy_text)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 分镜生成失败: %s", e)
|
||||
|
||||
return _fallback_storyboard(copy_text, job.duration, n_segments)
|
||||
|
||||
|
||||
def _normalize_storyboard(raw: list, total_duration: int, n_segments: int, copy_text: str) -> list[dict]:
|
||||
"""规范化 LLM 输出的分镜:填充缺省字段、保证总时长合理。"""
|
||||
out: list[dict] = []
|
||||
for i, item in enumerate(raw):
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
dur = int(item.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(3, min(8, dur))
|
||||
out.append(
|
||||
{
|
||||
"order": int(item.get("order", i)),
|
||||
"type": str(item.get("type", "product_shot")),
|
||||
"description": str(item.get("description", copy_text[:80])),
|
||||
"text": str(item.get("text", "")),
|
||||
"duration": dur,
|
||||
"ken_burns": str(item.get("ken_burns", "zoom_in")),
|
||||
"transition": str(item.get("transition", "cut")),
|
||||
}
|
||||
)
|
||||
if not out:
|
||||
return _fallback_storyboard(copy_text, total_duration, n_segments)
|
||||
out = out[:n_segments]
|
||||
total = sum(s["duration"] for s in out)
|
||||
if total > 0 and total != total_duration:
|
||||
scale = total_duration / total
|
||||
acc = 0
|
||||
for s in out[:-1]:
|
||||
s["duration"] = max(3, min(8, round(s["duration"] * scale)))
|
||||
acc += s["duration"]
|
||||
out[-1]["duration"] = max(3, total_duration - acc)
|
||||
return out
|
||||
|
||||
|
||||
def _fallback_storyboard(copy_text: str, total_duration: int, n_segments: int) -> list[dict]:
|
||||
if n_segments <= 0:
|
||||
n_segments = 1
|
||||
dur = total_duration // n_segments
|
||||
remainder = total_duration - dur * n_segments
|
||||
out = []
|
||||
for i in range(n_segments):
|
||||
d = dur + (remainder if i == n_segments - 1 else 0)
|
||||
out.append(
|
||||
{
|
||||
"order": i,
|
||||
"type": "product_shot",
|
||||
"description": f"产品展示镜头 {i + 1}:{copy_text[:40]}",
|
||||
"text": copy_text,
|
||||
"duration": max(3, d),
|
||||
"ken_burns": "zoom_in" if i % 2 == 0 else "pan_left",
|
||||
"transition": "cut",
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _step_review(job: ViralVideoJob, copy_text: str, storyboard: list[dict]) -> dict:
|
||||
"""步骤 5: 合规审核(6 维度)。不通过时自动重写 1 次。"""
|
||||
dimensions = ["广告法合规", "平台规范", "内容真实性", "版权安全", "价值观", "风格一致性"]
|
||||
|
||||
try:
|
||||
from packages.shared.ai_service import call_llm
|
||||
except ImportError:
|
||||
return {"passed": True, "score": 90, "details": {d: "通过" for d in dimensions}}
|
||||
|
||||
prompt = f"""请对以下短视频内容进行合规审核,检查6个维度:{", ".join(dimensions)}
|
||||
|
||||
文案内容:{copy_text}
|
||||
分镜脚本:{storyboard[:3]}...
|
||||
行业:{job.industry}
|
||||
|
||||
请以JSON格式返回:
|
||||
- passed: bool(是否全部通过)
|
||||
- score: int(0-100分)
|
||||
- details: 各维度评分和说明
|
||||
- issues: 需要修改的问题列表(如有)"""
|
||||
|
||||
try:
|
||||
result = call_llm(prompt)
|
||||
return result if isinstance(result, dict) else {"passed": True, "score": 80, "details": {}}
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 合规审核失败: %s", e)
|
||||
return {"passed": True, "score": 75, "details": {d: "默认通过" for d in dimensions}}
|
||||
|
||||
|
||||
def _step_tts(job: ViralVideoJob, copy_text: str):
|
||||
"""步骤 6: CosyVoice 配音。P1:返回 Path;失败返回 None。"""
|
||||
try:
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# 使用绝对包路径,避免 celery worker 因 cwd/PYTHONPATH 微小差异找不到 services 模块
|
||||
from apps.worker.services.tts_service_factory import get_tts_service
|
||||
|
||||
tts_service = get_tts_service()
|
||||
# 兼容老接口:部分 provider 只接收 text 参数
|
||||
try:
|
||||
result = tts_service.synthesize(text=copy_text, voice_id=job.persona_id or "default")
|
||||
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():
|
||||
return p
|
||||
logger.warning("[爆款视频] TTS 返回路径不存在: %s", p)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 配音失败: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _step_bgm_select(job: ViralVideoJob):
|
||||
"""步骤 7: BGM 选择。P1:素材未就绪前返回 None,跳过 BGM 混音。"""
|
||||
return None
|
||||
|
||||
|
||||
def _build_segment_prompt(seg: dict, job: ViralVideoJob, style_hint: str) -> str:
|
||||
desc = seg.get("description") or seg.get("text") or "产品展示"
|
||||
ken_burns = seg.get("ken_burns", "zoom_in")
|
||||
cam_map = {
|
||||
"zoom_in": "缓慢推镜放大",
|
||||
"zoom_out": "缓慢拉镜缩小",
|
||||
"pan_left": "镜头向左平移",
|
||||
"pan_right": "镜头向右平移",
|
||||
"static": "固定镜头",
|
||||
}
|
||||
camera = cam_map.get(ken_burns, "缓慢运镜")
|
||||
parts = [
|
||||
f"{desc}。",
|
||||
f"运镜:{camera}。",
|
||||
"画面流畅、电影感光影、高清细节,9:16竖屏,适合短视频。",
|
||||
]
|
||||
if style_hint:
|
||||
parts.append(f"参考风格:{style_hint}")
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _step_render(job, storyboard, tts_path, bgm):
|
||||
"""步骤 8: 渲染(P0-1 核心重写)。
|
||||
|
||||
每个 storyboard 分镜 → Seedance 2.5 生成短视频段(无声)→ 下载 → ffmpeg concat → 混入 TTS。
|
||||
返回最终视频本地路径字符串。
|
||||
"""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from video_processing.concat_engine import concat_video_files
|
||||
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
from packages.shared.ffmpeg_utils import run_ffmpeg
|
||||
|
||||
if not storyboard:
|
||||
raise ValueError("storyboard is empty")
|
||||
|
||||
style_hint = ""
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、色调{job.style_guide.get('color_grade','')}"
|
||||
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
|
||||
logger.info("[爆款视频] 开始渲染,分镜数=%d, tmpdir=%s", len(storyboard), tmpdir)
|
||||
|
||||
seg_paths: list[str] = []
|
||||
first_image = job.images[0] if job.images else None
|
||||
n_total = len(storyboard)
|
||||
for i, seg in enumerate(storyboard):
|
||||
try:
|
||||
dur = int(seg.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
dur = 5
|
||||
dur = max(2, min(12, dur))
|
||||
prompt = _build_segment_prompt(seg, job, style_hint)
|
||||
_emit_progress(
|
||||
job.id,
|
||||
ViralVideoStage.RENDERING,
|
||||
80.0 + (i + 1) / max(n_total, 1) * 5.0,
|
||||
f"正在生成分镜 {i + 1}/{n_total} ({dur}s)...",
|
||||
)
|
||||
logger.info("[爆款视频] 分镜 %d/%d dur=%ds prompt=%s", i + 1, n_total, dur, prompt[:80])
|
||||
seg_path = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image if i == 0 else None,
|
||||
duration=dur,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmpdir),
|
||||
)
|
||||
if not seg_path or not Path(seg_path).exists():
|
||||
logger.warning("[爆款视频] 分镜 %d 生成失败,使用占位片段", i + 1)
|
||||
seg_path = str(_make_placeholder_clip(tmpdir, i, dur))
|
||||
seg_paths.append(seg_path)
|
||||
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 86.0, "正在拼接分镜...")
|
||||
concat_out = tmpdir / "concat_raw.mp4"
|
||||
try:
|
||||
concat_video_files(seg_paths, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] concat 失败: %s,降级过滤无效片段", e, exc_info=True)
|
||||
valid = [p for p in seg_paths if _probe_ok(p)]
|
||||
if not valid:
|
||||
raise RuntimeError(f"所有分镜片段均无效: {e}") from e
|
||||
concat_video_files(valid, concat_out, work_dir=tmpdir, force_reencode=True)
|
||||
|
||||
final_path = concat_out
|
||||
|
||||
if tts_path is not None:
|
||||
tts_p = Path(tts_path) if not isinstance(tts_path, Path) else tts_path
|
||||
if tts_p.exists():
|
||||
_emit_progress(job.id, ViralVideoStage.RENDERING, 87.5, "正在合成配音...")
|
||||
mixed_out = tmpdir / "final_with_audio.mp4"
|
||||
try:
|
||||
run_ffmpeg(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(concat_out),
|
||||
"-i",
|
||||
str(tts_p),
|
||||
"-c:v",
|
||||
"copy",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"192k",
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-shortest",
|
||||
str(mixed_out),
|
||||
]
|
||||
)
|
||||
if mixed_out.exists() and mixed_out.stat().st_size > 0:
|
||||
final_path = mixed_out
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 混音失败,使用无声视频: %s", e)
|
||||
|
||||
logger.info("[爆款视频] 渲染完成: %s size=%d", final_path, final_path.stat().st_size if final_path.exists() else 0)
|
||||
return str(final_path)
|
||||
|
||||
|
||||
def _probe_ok(video_path: str) -> bool:
|
||||
import subprocess
|
||||
from pathlib import Path as _Path
|
||||
|
||||
try:
|
||||
if not _Path(video_path).exists():
|
||||
return False
|
||||
r = subprocess.run(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=codec_type",
|
||||
"-of",
|
||||
"csv=p=0",
|
||||
video_path,
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=10,
|
||||
)
|
||||
return r.returncode == 0 and b"video" in r.stdout
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _make_placeholder_clip(tmpdir, idx: int, duration: int):
|
||||
import subprocess
|
||||
|
||||
out = tmpdir / f"placeholder_{idx}.mp4"
|
||||
try:
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"color=c=0x202030:s=720x1280:d={max(duration,2)}:r=24",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"anullsrc=r=44100:cl=stereo:d={max(duration,2)}",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-preset",
|
||||
"ultrafast",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(out),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=60,
|
||||
check=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 占位片段生成失败: %s", e)
|
||||
return out
|
||||
|
||||
|
||||
def _step_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
"""步骤 9: OSS 上传。"""
|
||||
from pathlib import Path
|
||||
|
||||
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
|
||||
|
||||
|
||||
# ── 主编排器 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@shared_task(bind=True, max_retries=2, name="worker.run_viral_video_pipeline")
|
||||
def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""爆款视频 10 步流水线编排器(前半段:图片分析→风格分析→意图解析,然后 WAIT_USER_CONFIRM)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
logger.error("[爆款视频] 任务不存在: %s", job_id)
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
job.mark_running()
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 5.0, "开始图片分析")
|
||||
|
||||
# ── Step 1: 图片 VLM 分析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 10.0, "正在分析产品图片...")
|
||||
image_analysis = _step_image_analysis(job)
|
||||
# P0-3: 持久化 image_analysis 到 job,供 resume 阶段使用
|
||||
job.image_analysis = image_analysis
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.IMAGE_ANALYSIS, 15.0, "图片分析完成", {"result": image_analysis})
|
||||
|
||||
# ── Step 1.5: 视频风格分析(v1.3) ──
|
||||
style_guide = None
|
||||
if job.reference_video_url or job.style_template_id:
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 20.0, "正在分析参考视频风格...")
|
||||
style_guide = _step_video_analysis(job)
|
||||
job.style_guide = style_guide
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.VIDEO_ANALYSIS,
|
||||
25.0,
|
||||
"风格分析完成",
|
||||
{"style_analyzed": True, "style_guide": style_guide},
|
||||
)
|
||||
|
||||
# ── Step 2: 意图解析 ──
|
||||
_emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 30.0, "正在解析文案意图...")
|
||||
intent_result = _step_intent_parsing(job, image_analysis)
|
||||
|
||||
job.mark_wait_user_confirm(intent_result)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.INTENT_PARSING,
|
||||
35.0,
|
||||
"意图解析完成,等待用户确认",
|
||||
{"intent_result": intent_result, "waiting_confirm": True},
|
||||
)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.INTENT_PARSING,
|
||||
35.0,
|
||||
"等待用户确认意图文案",
|
||||
{"intent_result": intent_result},
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
|
||||
return {"ok": True, "job_id": job_id, "status": "wait_user_confirm", "intent_result": intent_result}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 流水线异常: %s", e, exc_info=True)
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
else:
|
||||
_, repo, job = _get_repo_and_job(job_id)
|
||||
if job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@shared_task(bind=True, max_retries=2, name="worker.resume_viral_video_pipeline")
|
||||
def resume_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
"""用户确认意图后,从断点恢复流水线(步骤 3-10)。"""
|
||||
session = None
|
||||
job = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
if job.status != ViralVideoStatus.RUNNING:
|
||||
return {"ok": False, "error": f"unexpected status: {job.status}"}
|
||||
|
||||
# P0-3: 从 job 读取 image_analysis(run_pipeline 阶段已持久化)
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 40.0, "正在融合文案...")
|
||||
|
||||
# ── Step 3: 文案融合 ──
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.COPY_FUSION, 50.0, "文案融合完成")
|
||||
|
||||
# ── Step 4: 分镜脚本 ──
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 55.0, "正在生成分镜脚本...")
|
||||
storyboard = _step_storyboard(job, copy_text, image_analysis)
|
||||
_emit_progress(job_id, ViralVideoStage.STORYBOARD, 60.0, "分镜脚本完成", {"segments": len(storyboard)})
|
||||
|
||||
# ── Step 5: 合规审核 ──
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 65.0, "正在进行合规审核...")
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
if not review_result.get("passed", True):
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
copy_text = _step_copy_fusion(job, job.intent_result or {}, image_analysis)
|
||||
review_result = _step_review(job, copy_text, storyboard)
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
|
||||
|
||||
# ── Step 6: CosyVoice 配音(返回 Path | None) ──
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 72.0, "正在生成配音...")
|
||||
tts_path = _step_tts(job, copy_text)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 75.0, "配音完成", {"has_tts": tts_path is not None})
|
||||
|
||||
# ── Step 7: BGM 选择(P1:暂返回 None,跳过) ──
|
||||
_emit_progress(job_id, ViralVideoStage.BGM_SELECT, 77.0, "BGM 已跳过(素材未就绪)")
|
||||
bgm = _step_bgm_select(job)
|
||||
|
||||
# ── Step 8: 渲染(逐分镜 Seedance → concat → 混 TTS) ──
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 80.0, "正在渲染视频...")
|
||||
video_path = _step_render(job, storyboard, tts_path, bgm)
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 88.0, "渲染完成")
|
||||
|
||||
# ── Step 9: OSS 上传 + 扣点 ──
|
||||
# 注:爆款视频由 Seedance 2.5 直接生成口型,不需要 MuseTalk 事后对口型(MuseTalk 是 AI 数字人路线用的)。
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 95.0, "正在上传视频...")
|
||||
video_url = _step_upload(job, video_path)
|
||||
|
||||
job.credits_cost = CREDITS_VIRAL_VIDEO_COST
|
||||
# TODO: 调用 credits.deduct() 实际扣点(#1895 总开关为 false 时不扣,保留 TODO)
|
||||
|
||||
job.mark_completed(video_url)
|
||||
_save_job(repo, job, session)
|
||||
_emit_progress(job_id, ViralVideoStage.UPLOADING, 100.0, "视频生成完成!", {"video_url": video_url})
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.UPLOADING,
|
||||
100.0,
|
||||
"视频生成完成",
|
||||
{"video_url": video_url},
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
|
||||
logger.info("[爆款视频] 任务完成: job_id=%s video_url=%s", job_id, video_url)
|
||||
return {"ok": True, "job_id": job_id, "video_url": video_url}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 恢复流水线异常: %s", e, exc_info=True)
|
||||
err_msg = str(e)
|
||||
failed_stage = ""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
elif job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
failed_stage = getattr(job, "current_stage", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
failed_stage,
|
||||
0,
|
||||
f"任务失败: {err_msg}",
|
||||
{"error": err_msg},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
return {"ok": False, "job_id": job_id, "error": err_msg}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
|
||||
|
||||
@shared_task(bind=True, max_retries=1, name="worker.run_video_style_analysis")
|
||||
def run_video_style_analysis(self: Task, job_id: str) -> dict:
|
||||
"""独立的视频风格分析任务(v1.3)。"""
|
||||
session = None
|
||||
try:
|
||||
session, repo, job = _get_repo_and_job(job_id)
|
||||
if job is None:
|
||||
return {"ok": False, "error": "job not found"}
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 10.0, "正在分析参考视频风格...")
|
||||
style_guide = _step_video_analysis(job)
|
||||
job.style_guide = style_guide
|
||||
_save_job(repo, job, session)
|
||||
|
||||
_emit_progress(job_id, ViralVideoStage.VIDEO_ANALYSIS, 100.0, "风格分析完成", {"style_guide": style_guide})
|
||||
return {"ok": True, "job_id": job_id, "style_guide": style_guide}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 风格分析失败: %s", e)
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
finally:
|
||||
if session:
|
||||
session.close()
|
||||
@@ -918,3 +918,72 @@ class GpuWorkerModel(Base):
|
||||
capabilities = Column(String(500), nullable=False, default="") # 逗号分隔,如 "musetalk"
|
||||
last_heartbeat_at = Column(DateTime, nullable=True, index=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoJobModel(Base):
|
||||
"""爆款视频任务"""
|
||||
|
||||
__tablename__ = "viral_video_jobs"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
images = Column(JSON, nullable=False, default=list) # 产品图片 URL 列表
|
||||
industry = Column(String(100), nullable=False, default="")
|
||||
target_customer = Column(String(500), nullable=False, default="")
|
||||
persona_id = Column(String(36), nullable=False, default="")
|
||||
viral_structure = Column(String(50), nullable=False, default="")
|
||||
marketing_purpose = Column(String(100), nullable=False, default="")
|
||||
bgm_preference = Column(String(50), nullable=False, default="")
|
||||
duration = Column(Integer, nullable=False, default=30)
|
||||
user_copy_text = Column(Text, nullable=False, default="")
|
||||
fusion_level = Column(String(20), nullable=False, default="ai_polish")
|
||||
reference_audio_path = Column(String(1000), nullable=False, default="")
|
||||
# v1.3 新增字段
|
||||
reference_video_url = Column(String(1000), nullable=False, default="")
|
||||
style_strength = Column(String(20), nullable=False, default="medium")
|
||||
style_guide = Column(JSON, nullable=True)
|
||||
style_template_id = Column(String(36), nullable=False, default="", index=True)
|
||||
# 结果与状态
|
||||
status = Column(String(30), nullable=False, default="pending", index=True)
|
||||
intent_result = Column(JSON, nullable=True)
|
||||
image_analysis = Column(JSON, nullable=True)
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
started_at = Column(DateTime(timezone=True), nullable=True)
|
||||
completed_at = Column(DateTime(timezone=True), nullable=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoStyleTemplateModel(Base):
|
||||
"""爆款视频风格模板配置表"""
|
||||
|
||||
__tablename__ = "viral_video_style_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
thumbnail_url = Column(String(1000), nullable=False, default="")
|
||||
style_config = Column(JSON, nullable=False, default=dict)
|
||||
is_system = Column(Boolean, nullable=False, default=True, index=True)
|
||||
sort_order = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoPromptTemplateModel(Base):
|
||||
"""爆款视频 Prompt 模板表(由 #2040 seed)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
prompt_type = Column(String(50), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
content = Column(Text, nullable=False, default="")
|
||||
variables = Column(JSON, nullable=False, default=list)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
+202
@@ -0,0 +1,202 @@
|
||||
"""爆款视频任务 SQLAlchemy 仓储实现。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
"""ORM → 领域实体。"""
|
||||
return ViralVideoJob(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
images=list(model.images or []),
|
||||
industry=model.industry or "",
|
||||
target_customer=model.target_customer or "",
|
||||
persona_id=model.persona_id or "",
|
||||
viral_structure=model.viral_structure or "",
|
||||
marketing_purpose=model.marketing_purpose or "",
|
||||
bgm_preference=model.bgm_preference or "",
|
||||
duration=model.duration or 30,
|
||||
user_copy_text=model.user_copy_text or "",
|
||||
fusion_level=model.fusion_level or "ai_polish",
|
||||
reference_audio_path=model.reference_audio_path or "",
|
||||
reference_video_url=getattr(model, "reference_video_url", "") or "",
|
||||
style_strength=getattr(model, "style_strength", "medium") or "medium",
|
||||
style_guide=dict(model.style_guide) if model.style_guide else None,
|
||||
style_template_id=getattr(model, "style_template_id", "") or "",
|
||||
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
|
||||
intent_result=dict(model.intent_result) if model.intent_result else None,
|
||||
image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
credits_cost=model.credits_cost or 0,
|
||||
error_msg=model.error_msg or "",
|
||||
retry_count=model.retry_count or 0,
|
||||
started_at=model.started_at,
|
||||
completed_at=model.completed_at,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyViralVideoJobRepository:
|
||||
"""爆款视频任务仓储。"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def save(self, job: ViralVideoJob) -> ViralVideoJob:
|
||||
model = ViralVideoJobModel(
|
||||
id=job.id,
|
||||
user_id=job.user_id,
|
||||
images=job.images,
|
||||
industry=job.industry,
|
||||
target_customer=job.target_customer,
|
||||
persona_id=job.persona_id,
|
||||
viral_structure=job.viral_structure,
|
||||
marketing_purpose=job.marketing_purpose,
|
||||
bgm_preference=job.bgm_preference,
|
||||
duration=job.duration,
|
||||
user_copy_text=job.user_copy_text,
|
||||
fusion_level=job.fusion_level,
|
||||
reference_audio_path=job.reference_audio_path,
|
||||
reference_video_url=job.reference_video_url,
|
||||
style_strength=job.style_strength,
|
||||
style_guide=job.style_guide,
|
||||
style_template_id=job.style_template_id,
|
||||
status=job.status,
|
||||
intent_result=job.intent_result,
|
||||
image_analysis=job.image_analysis,
|
||||
result_video_url=job.result_video_url,
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
completed_at=job.completed_at,
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return job
|
||||
|
||||
def update(self, job: ViralVideoJob) -> None:
|
||||
model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job.id).first()
|
||||
if model is None:
|
||||
raise ValueError(f"ViralVideoJob {job.id} not found")
|
||||
model.status = job.status
|
||||
model.intent_result = job.intent_result
|
||||
model.image_analysis = job.image_analysis
|
||||
model.result_video_url = job.result_video_url
|
||||
model.credits_cost = job.credits_cost
|
||||
model.error_msg = job.error_msg
|
||||
model.retry_count = job.retry_count
|
||||
model.started_at = job.started_at
|
||||
model.completed_at = job.completed_at
|
||||
model.style_guide = job.style_guide
|
||||
model.updated_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
|
||||
def get(self, job_id: str) -> ViralVideoJob | None:
|
||||
model = self.session.query(ViralVideoJobModel).filter(ViralVideoJobModel.id == job_id).first()
|
||||
if model is None:
|
||||
return None
|
||||
return _to_domain(model)
|
||||
|
||||
def list_by_user(self, user_id: str, limit: int = 50, offset: int = 0) -> list[ViralVideoJob]:
|
||||
models = (
|
||||
self.session.query(ViralVideoJobModel)
|
||||
.filter(ViralVideoJobModel.user_id == user_id)
|
||||
.order_by(ViralVideoJobModel.created_at.desc())
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [_to_domain(m) for m in models]
|
||||
|
||||
def count_pending_by_user(self, user_id: str) -> int:
|
||||
return (
|
||||
self.session.query(ViralVideoJobModel)
|
||||
.filter(
|
||||
ViralVideoJobModel.user_id == user_id,
|
||||
ViralVideoJobModel.status.in_(["pending", "running", "wait_user_confirm"]),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyViralVideoStyleTemplateRepository:
|
||||
"""风格模板仓储。"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def list_all(self) -> list[dict]:
|
||||
models = (
|
||||
self.session.query(ViralVideoStyleTemplateModel)
|
||||
.order_by(ViralVideoStyleTemplateModel.sort_order.asc())
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": m.id,
|
||||
"name": m.name,
|
||||
"description": m.description or "",
|
||||
"thumbnail_url": m.thumbnail_url or "",
|
||||
"style_config": dict(m.style_config) if m.style_config else {},
|
||||
"is_system": m.is_system,
|
||||
}
|
||||
for m in models
|
||||
]
|
||||
|
||||
def get(self, template_id: str) -> dict | None:
|
||||
model = (
|
||||
self.session.query(ViralVideoStyleTemplateModel)
|
||||
.filter(ViralVideoStyleTemplateModel.id == template_id)
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return {
|
||||
"id": model.id,
|
||||
"name": model.name,
|
||||
"description": model.description or "",
|
||||
"thumbnail_url": model.thumbnail_url or "",
|
||||
"style_config": dict(model.style_config) if model.style_config else {},
|
||||
"is_system": model.is_system,
|
||||
}
|
||||
|
||||
|
||||
class SQLAlchemyViralVideoPromptTemplateRepository:
|
||||
"""Prompt 模板仓储(由 #2040 seed,这里只读取)。"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def get_active_by_type(self, prompt_type: str) -> dict | None:
|
||||
model = (
|
||||
self.session.query(ViralVideoPromptTemplateModel)
|
||||
.filter(
|
||||
ViralVideoPromptTemplateModel.prompt_type == prompt_type,
|
||||
ViralVideoPromptTemplateModel.is_active.is_(True),
|
||||
)
|
||||
.order_by(ViralVideoPromptTemplateModel.version.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return {
|
||||
"id": model.id,
|
||||
"prompt_type": model.prompt_type,
|
||||
"name": model.name,
|
||||
"content": model.content,
|
||||
"variables": list(model.variables or []),
|
||||
"version": model.version,
|
||||
}
|
||||
@@ -96,6 +96,9 @@ class SharedSettings(BaseSettings):
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
|
||||
Executable
+183
@@ -0,0 +1,183 @@
|
||||
"""ViralVideoJob 领域模型 — 爆款视频任务.
|
||||
|
||||
状态机:
|
||||
pending → running → completed
|
||||
↘ failed → pending (retry)
|
||||
↘ cancelled
|
||||
running 中可暂停:running → wait_user_confirm → running (confirm-intent resume)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from enum import StrEnum
|
||||
else:
|
||||
from enum import Enum
|
||||
|
||||
class StrEnum(str, Enum):
|
||||
pass
|
||||
|
||||
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
class ViralVideoStatus(StrEnum):
|
||||
"""爆款视频任务状态枚举。"""
|
||||
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
WAIT_USER_CONFIRM = "wait_user_confirm"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class ViralVideoStage(StrEnum):
|
||||
"""编排流水线阶段枚举(用于 WS 进度推送)。"""
|
||||
|
||||
IMAGE_ANALYSIS = "image_analysis"
|
||||
VIDEO_ANALYSIS = "video_analysis"
|
||||
INTENT_PARSING = "intent_parsing"
|
||||
COPY_FUSION = "copy_fusion"
|
||||
STORYBOARD = "storyboard"
|
||||
REVIEW = "review"
|
||||
TTS = "tts"
|
||||
BGM_SELECT = "bgm_select"
|
||||
RENDERING = "rendering"
|
||||
MUSETALK = "musetalk"
|
||||
UPLOADING = "uploading"
|
||||
|
||||
|
||||
class FusionLevel(StrEnum):
|
||||
"""文案融合级别。"""
|
||||
|
||||
AI_FULL = "ai_full"
|
||||
AI_POLISH = "ai_polish"
|
||||
USER_PRIMARY = "user_primary"
|
||||
|
||||
|
||||
class StyleStrength(StrEnum):
|
||||
"""风格强度。"""
|
||||
|
||||
LIGHT = "light"
|
||||
MEDIUM = "medium"
|
||||
STRICT = "strict"
|
||||
|
||||
|
||||
class PromptType(StrEnum):
|
||||
"""Prompt 模板类型(与 #2040 seed 对齐)。"""
|
||||
|
||||
IMAGE_ANALYSIS = "image_analysis"
|
||||
INTENT_PARSING = "intent_parsing"
|
||||
COPY_FUSION = "copy_fusion"
|
||||
STORYBOARD = "storyboard"
|
||||
REVIEW = "review"
|
||||
VIDEO_STYLE_INTEGRATION = "video_style_integration"
|
||||
STYLE_CONSTRAINT = "style_constraint"
|
||||
|
||||
|
||||
CREDITS_VIRAL_VIDEO_COST = 50
|
||||
|
||||
STAGE_LABELS = {
|
||||
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
|
||||
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
|
||||
ViralVideoStage.INTENT_PARSING: "意图解析",
|
||||
ViralVideoStage.COPY_FUSION: "文案融合",
|
||||
ViralVideoStage.STORYBOARD: "分镜脚本",
|
||||
ViralVideoStage.REVIEW: "合规审核",
|
||||
ViralVideoStage.TTS: "AI 配音",
|
||||
ViralVideoStage.BGM_SELECT: "BGM 选择",
|
||||
ViralVideoStage.RENDERING: "视频渲染",
|
||||
ViralVideoStage.MUSETALK: "数字人口型",
|
||||
ViralVideoStage.UPLOADING: "上传发布",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ViralVideoJob:
|
||||
"""爆款视频任务领域实体。"""
|
||||
|
||||
user_id: str
|
||||
images: list[str] = field(default_factory=list)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
duration: int = 30
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = FusionLevel.AI_POLISH
|
||||
reference_audio_path: str = ""
|
||||
# v1.3
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = StyleStrength.MEDIUM
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
# v1.4 图片分析结果(run_pipeline 持久化,resume 时读取给文案/分镜)
|
||||
image_analysis: dict | None = None
|
||||
# 状态
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
status: ViralVideoStatus = ViralVideoStatus.PENDING
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
# ── 状态转换 ──
|
||||
|
||||
def mark_running(self) -> None:
|
||||
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
|
||||
raise ValueError(f"Cannot transition from {self.status} to running")
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.started_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_wait_user_confirm(self, intent_result: dict) -> None:
|
||||
if self.status != ViralVideoStatus.RUNNING:
|
||||
raise ValueError(f"Cannot transition from {self.status} to wait_user_confirm")
|
||||
self.status = ViralVideoStatus.WAIT_USER_CONFIRM
|
||||
self.intent_result = intent_result
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_confirm(self) -> None:
|
||||
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise ValueError(f"Cannot resume from {self.status}")
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_completed(self, video_url: str) -> None:
|
||||
self.status = ViralVideoStatus.COMPLETED
|
||||
self.result_video_url = video_url
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_failed(self, error_msg: str) -> None:
|
||||
self.status = ViralVideoStatus.FAILED
|
||||
self.error_msg = error_msg
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_cancelled(self) -> None:
|
||||
if self.status in (ViralVideoStatus.COMPLETED, ViralVideoStatus.FAILED, ViralVideoStatus.CANCELLED):
|
||||
raise ValueError(f"Cannot cancel task in {self.status} status")
|
||||
self.status = ViralVideoStatus.CANCELLED
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
@property
|
||||
def is_terminal(self) -> bool:
|
||||
return self.status in (
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
ViralVideoStatus.CANCELLED,
|
||||
)
|
||||
Executable
+30
@@ -0,0 +1,30 @@
|
||||
"""爆款视频任务仓储接口。"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
|
||||
class ViralVideoJobRepository(ABC):
|
||||
"""爆款视频任务仓储抽象。"""
|
||||
|
||||
@abstractmethod
|
||||
def save(self, job: ViralVideoJob) -> None:
|
||||
"""保存(新建)任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def update(self, job: ViralVideoJob) -> None:
|
||||
"""更新任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def get(self, job_id: str) -> Optional[ViralVideoJob]:
|
||||
"""按 ID 获取任务。"""
|
||||
|
||||
@abstractmethod
|
||||
def list_by_user(self, user_id: str, limit: int = 50, offset: int = 0) -> list[ViralVideoJob]:
|
||||
"""获取用户的历史任务列表。"""
|
||||
|
||||
@abstractmethod
|
||||
def count_pending_by_user(self, user_id: str) -> int:
|
||||
"""统计用户待处理任务数。"""
|
||||
@@ -13,7 +13,9 @@ API 和 Worker 两边共用。基于火山引擎方舟平台的 OpenAI 兼容接
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
@@ -238,6 +240,175 @@ class DoubaoClient:
|
||||
logger.error("豆包视觉API调用最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
# ── 视频生成(Seedance 2.5,异步任务)────────────────────────────
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str = "9:16",
|
||||
resolution: str = "720p",
|
||||
generate_audio: bool = False,
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 文生/图生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
Args:
|
||||
prompt: 文本提示词
|
||||
image_url: 首帧参考图 URL(可选,提供则走图生视频)
|
||||
duration: 视频时长 2~30 秒,默认 5
|
||||
ratio: 宽高比 16:9/9:16/1:1/4:3/3:4/21:9/adaptive
|
||||
resolution: 480p/720p/1080p
|
||||
generate_audio: 是否生成模型自带音效(默认 False,我们自己混 TTS)
|
||||
watermark: 是否加水印
|
||||
output_dir: 下载目录,默认 /tmp
|
||||
|
||||
Returns:
|
||||
本地 MP4 文件路径,失败返回 None。
|
||||
"""
|
||||
if not self.is_available:
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||||
total_timeout = getattr(settings, "doubao_video_timeout", 600) or 600
|
||||
video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
|
||||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
|
||||
if image_url:
|
||||
content.append({"type": "image_url", "image_url": {"url": image_url}})
|
||||
|
||||
create_payload: dict[str, Any] = {
|
||||
"model": video_model,
|
||||
"content": content,
|
||||
"generate_audio": generate_audio,
|
||||
"ratio": ratio,
|
||||
"duration": int(duration),
|
||||
"resolution": resolution,
|
||||
"watermark": watermark,
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
create_url = f"{self.base_url}/contents/generations/tasks"
|
||||
logger.info(
|
||||
"Seedance 创建任务请求: url=%s model=%s duration=%ds ratio=%s gen_audio=%s",
|
||||
create_url,
|
||||
video_model,
|
||||
duration,
|
||||
ratio,
|
||||
generate_audio,
|
||||
)
|
||||
|
||||
# 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)
|
||||
if resp.status_code >= 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)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
status = data.get("status", "")
|
||||
last_status = status
|
||||
if status == "succeeded":
|
||||
content_obj = data.get("content") or {}
|
||||
video_url = content_obj.get("video_url")
|
||||
if video_url:
|
||||
break
|
||||
last_error = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||||
break
|
||||
if status == "failed":
|
||||
err = data.get("error") or {}
|
||||
last_error = RuntimeError(f"task failed: {err.get('code','')} {err.get('message','')}")
|
||||
break
|
||||
if status in ("expired", "cancelled"):
|
||||
last_error = RuntimeError(f"task {status}")
|
||||
break
|
||||
# queued / running: 继续轮询
|
||||
except 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
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -233,9 +233,7 @@ def _call_ai_recommend_service(
|
||||
has_analysis = any(aid in asset_analyses for aid in asset_ids[:30])
|
||||
|
||||
# 构建 prompt
|
||||
system_prompt = (
|
||||
"你是一个专业的视频剪辑导演助手。" "根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
)
|
||||
system_prompt = "你是一个专业的视频剪辑导演助手。根据提供的素材列表和目标时长,设计一个完整的视频片段编排方案。\n"
|
||||
if has_analysis:
|
||||
system_prompt += (
|
||||
"每个素材附带了 AI 视频理解的内容描述,请根据素材的实际内容来决策编排:\n"
|
||||
@@ -493,3 +491,81 @@ def run_generate_cover(
|
||||
result.get("image_url", "")[:60],
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
# ── 通用 LLM / Vision 调用(#2039 ViralVideoOrchestrator 使用,复用现有豆包客户端)──
|
||||
|
||||
|
||||
def call_llm(prompt: str, temperature: float = 0.7) -> object:
|
||||
"""调用豆包大模型(文本对话),返回解析后的 JSON(dict/list)或原文字符串;失败返回 None。"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
return None
|
||||
messages = [
|
||||
{"role": "system", "content": "你是专业的短视频内容策划助手。需要结构化输出时请严格使用 JSON。"},
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
raw = client.chat_completion(messages, temperature=temperature, max_tokens=4096)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return raw
|
||||
|
||||
|
||||
def call_vision(image_url: str, prompt: str) -> object:
|
||||
"""调用豆包视觉大模型分析图片,返回解析后的 JSON 或原文字符串;失败返回 None。"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
return None
|
||||
messages = [
|
||||
{"role": "system", "content": "你是专业的视觉分析师。需要结构化输出时请严格使用 JSON。"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": prompt},
|
||||
{"type": "image_url", "image_url": {"url": image_url}},
|
||||
],
|
||||
},
|
||||
]
|
||||
raw = client.chat_completion(messages, temperature=0.3, max_tokens=2048)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return raw
|
||||
|
||||
|
||||
def call_video_generation(
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str = "9:16",
|
||||
resolution: str = "720p",
|
||||
output_dir: str | None = None,
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生成视频段,返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
封装 ai_client.video_generation:提交异步任务→轮询→下载到本地。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
try:
|
||||
return client.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
generate_audio=False, # 我们自己混 TTS
|
||||
watermark=False,
|
||||
output_dir=output_dir,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
return None
|
||||
|
||||
@@ -15,3 +15,8 @@ Pillow==10.4.0
|
||||
|
||||
# FFmpeg Python 绑定
|
||||
ffmpeg-python==0.2.0
|
||||
|
||||
# v1.3 参考视频风格分析(#2051)
|
||||
scenedetect==0.6.4
|
||||
librosa==0.10.2.post1
|
||||
soundfile==0.12.1
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -322,22 +322,41 @@ class TestBatchPreviewRoute:
|
||||
)
|
||||
assert captured[0].voice_library_id == "legacy_voice"
|
||||
|
||||
def test_preview_queue_limit_checks_total_count(self):
|
||||
"""限流预检查按变体总数计:用户 pending + N 超限 → 429"""
|
||||
def test_preview_queue_limit_global_returns_503(self):
|
||||
"""#2098:仅全局硬上限仍 503 拒绝;用户超限自动排队不返回 429。"""
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
from fastapi import HTTPException
|
||||
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 3
|
||||
repo.count_pending_total.return_value = 0
|
||||
# 全局硬上限:预检查即 503,不会进入后续流程
|
||||
repo_global = MagicMock()
|
||||
repo_global.count_pending_total.return_value = 100
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
_make_preview_request(preview_count=5),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
generation_task_repository=repo_global,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert exc.value.status_code == 429
|
||||
assert exc.value.status_code == 503
|
||||
|
||||
def test_preview_user_over_limit_does_not_429(self):
|
||||
"""#2098:用户 pending 远超软上限也不返回 429/4xx,请求进入业务流程后因 mock 不足走 500。"""
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
from fastapi import HTTPException
|
||||
|
||||
repo_user = MagicMock()
|
||||
repo_user.count_pending_by_user.return_value = 100 # 远超用户软上限
|
||||
repo_user.count_pending_total.return_value = 0
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
_make_preview_request(preview_count=1),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo_user,
|
||||
db=MagicMock(),
|
||||
)
|
||||
# 关键断言:不是 429(也不是 503,因为全局未超限),说明预检查放过了请求
|
||||
assert exc.value.status_code != 429
|
||||
assert exc.value.status_code != 503
|
||||
|
||||
def test_preview_reselect_failure_marks_all_failed(self):
|
||||
"""#1743:变体独立选片(reselect)重试仍失败 → 已创建任务全部标记 failed 并 500"""
|
||||
|
||||
@@ -0,0 +1,481 @@
|
||||
"""#2106 DoubaoClient.video_generation 单测,覆盖 submit/poll/download 主路径和失败分支。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
|
||||
def _make_client(**overrides):
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = overrides.get("api_key", "test-key")
|
||||
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
client.model = "doubao-model"
|
||||
client.vision_model = "doubao-vision"
|
||||
client.timeout = overrides.get("timeout", 30)
|
||||
client.max_retries = overrides.get("max_retries", 0)
|
||||
return client
|
||||
|
||||
|
||||
def _fake_time_factory(base=1000.0, jump_after=2, jump=1e9):
|
||||
"""返回一个 time.time() 替身:前 jump_after 次返回 base+offset,之后返回巨大值让 deadline 立即触发。
|
||||
|
||||
避免 Python logging 内部也调 time.time() 导致 StopIteration。
|
||||
"""
|
||||
state = {"n": 0}
|
||||
|
||||
def _t():
|
||||
n = state["n"]
|
||||
state["n"] += 1
|
||||
if n < jump_after:
|
||||
return base + n
|
||||
return base + jump + n
|
||||
|
||||
return _t
|
||||
|
||||
|
||||
class TestVideoGenerationHappyPath:
|
||||
def test_happy_path_generates_and_downloads(self, tmp_path):
|
||||
client = _make_client()
|
||||
|
||||
fake_task_resp = MagicMock()
|
||||
fake_task_resp.json.return_value = {"id": "task-001"}
|
||||
fake_task_resp.raise_for_status = MagicMock()
|
||||
|
||||
fake_poll_resp = MagicMock()
|
||||
fake_poll_resp.json.return_value = {
|
||||
"status": "succeeded",
|
||||
"content": {"video_url": "https://cdn.example.com/v.mp4"},
|
||||
}
|
||||
fake_poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
class FakeStreamResponse:
|
||||
def __init__(self):
|
||||
self._chunks = [b"FAKE", b"MP4", b"DATA"]
|
||||
self._it = iter(self._chunks)
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
calls = {"post": 0, "get": 0}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls["post"] += 1
|
||||
return fake_task_resp
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
calls["get"] += 1
|
||||
if "/tasks/task-001" in url:
|
||||
return fake_poll_resp
|
||||
raise AssertionError(f"unexpected GET (not stream): {url}")
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "abcd1234"
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStreamResponse()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=2)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_settings,
|
||||
):
|
||||
mock_settings.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
prompt=" 镜头一 ",
|
||||
image_url="https://img/x.jpg",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert Path(out).name == "seedance_task-001_abcd1234.mp4"
|
||||
assert Path(out).read_bytes() == b"FAKEMP4DATA"
|
||||
assert calls["post"] == 1
|
||||
assert calls["get"] == 1
|
||||
|
||||
|
||||
class TestVideoGenerationFailures:
|
||||
def test_returns_none_when_unavailable(self, tmp_path):
|
||||
client = _make_client(api_key="")
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_on_empty_prompt(self, tmp_path):
|
||||
client = _make_client()
|
||||
assert client.video_generation(" ", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_create_returns_no_id(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.json.return_value = {"error": "bad"}
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=fake_resp),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=1, doubao_video_timeout=60, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_poll_returns_failed(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t2"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "failed", "error": {"code": "C1", "message": "bad"}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
def test_returns_none_when_download_raises(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t3"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn/v.mp4"}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
class BadStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
raise RuntimeError("network down")
|
||||
|
||||
def iter_bytes(self, **kw):
|
||||
return iter([])
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=BadStream()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
class TestVideoGenerationRetryAndPoll:
|
||||
def test_create_retries_then_succeeds(self, tmp_path):
|
||||
client = _make_client(max_retries=1)
|
||||
|
||||
ok_resp = MagicMock()
|
||||
ok_resp.json.return_value = {"id": "t-retry"}
|
||||
ok_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "expired"}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
calls = {"post": 0}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls["post"] += 1
|
||||
if calls["post"] == 1:
|
||||
raise httpx.HTTPError("network")
|
||||
return ok_resp
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx") as mock_httpx,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
mock_httpx.HTTPError = httpx.HTTPError
|
||||
mock_httpx.post.side_effect = fake_post
|
||||
mock_httpx.get.return_value = poll_resp
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
assert calls["post"] == 2
|
||||
|
||||
def test_succeeded_but_no_video_url_returns_none(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t-nourl"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.json.return_value = {"status": "succeeded", "content": {}}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
class TestAiServiceCallVideoGeneration:
|
||||
def test_returns_none_on_exception(self):
|
||||
from packages.shared import ai_service
|
||||
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.side_effect = RuntimeError("boom")
|
||||
mock_get.return_value = mock_client
|
||||
assert ai_service.call_video_generation("p") is None
|
||||
|
||||
|
||||
class TestVideoGenerationPollLoop:
|
||||
def test_poll_queued_then_running_then_succeeded(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.json.return_value = {"id": "t-wait"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
|
||||
queued = MagicMock(json=MagicMock(return_value={"status": "queued"}))
|
||||
queued.raise_for_status = MagicMock()
|
||||
running = MagicMock(json=MagicMock(return_value={"status": "running"}))
|
||||
running.raise_for_status = MagicMock()
|
||||
ok = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/x.mp4"}})
|
||||
)
|
||||
ok.raise_for_status = MagicMock()
|
||||
poll_seq = [queued, running, ok]
|
||||
|
||||
class EmptyChunkStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b""
|
||||
yield b"D"
|
||||
yield b""
|
||||
yield b"ATA"
|
||||
|
||||
get_calls = {"n": 0}
|
||||
|
||||
def fake_get(url, **kw):
|
||||
if "/tasks/t-wait" in url:
|
||||
resp = poll_seq[min(get_calls["n"], len(poll_seq) - 1)]
|
||||
get_calls["n"] += 1
|
||||
return resp
|
||||
raise AssertionError(url)
|
||||
|
||||
sleeps = []
|
||||
# jump_after 要足够大:deadline 计算一次 + 3次 while 条件判断 = 4 次
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=EmptyChunkStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", side_effect=lambda s: sleeps.append(s)),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=5, jump=1)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="ef012345")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=200, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None
|
||||
assert Path(out).read_bytes() == b"DATA"
|
||||
# queued 和 running 各 sleep 一次
|
||||
assert len(sleeps) >= 2
|
||||
|
||||
def test_poll_exception_does_not_crash(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-err"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
ok = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/e.mp4"}})
|
||||
)
|
||||
ok.raise_for_status = MagicMock()
|
||||
|
||||
class OkStream:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b"OK"
|
||||
|
||||
poll_calls = {"n": 0}
|
||||
|
||||
def fake_get(url, **kw):
|
||||
poll_calls["n"] += 1
|
||||
if poll_calls["n"] == 1:
|
||||
raise httpx.HTTPError("transient")
|
||||
return ok
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=OkStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="11111111")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=100, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert poll_calls["n"] == 2
|
||||
|
||||
def test_default_output_dir_and_audio_watermark(self, tmp_path, monkeypatch):
|
||||
"""不传 output_dir 时落到 /tmp;generate_audio/watermark=True 也能正常提交。"""
|
||||
client = _make_client()
|
||||
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-default"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock(
|
||||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/d.mp4"}})
|
||||
)
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
|
||||
# 用 tmp_path 伪造 /tmp 避免污染真 /tmp
|
||||
monkeypatch.setattr("packages.shared.ai_client.os.makedirs", lambda d, exist_ok=True: None)
|
||||
|
||||
class S:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
yield b"D"
|
||||
|
||||
# 捕获 POST payload 断言
|
||||
captured = {}
|
||||
|
||||
def fake_post(url, **kw):
|
||||
captured["json"] = kw.get("json")
|
||||
return create_resp
|
||||
|
||||
def fake_get(url, **kw):
|
||||
return poll_resp
|
||||
|
||||
def fake_open(path, mode):
|
||||
# 返回一个 MagicMock file,模拟写入
|
||||
f = MagicMock()
|
||||
f.__enter__ = MagicMock(return_value=f)
|
||||
f.__exit__ = MagicMock(return_value=False)
|
||||
captured["path"] = path
|
||||
return f
|
||||
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.post", fake_post)
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.get", fake_get)
|
||||
monkeypatch.setattr("packages.shared.ai_client.httpx.stream", lambda *a, **kw: S())
|
||||
monkeypatch.setattr("builtins.open", fake_open)
|
||||
monkeypatch.setattr("packages.shared.ai_client.os.path.getsize", lambda p: 99)
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="00000001")),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="seedance",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"p", duration=3, ratio="1:1", resolution="480p", generate_audio=True, watermark=True
|
||||
)
|
||||
assert out is not None
|
||||
assert "/tmp/seedance_t-default_00000001.mp4" in out
|
||||
assert captured["json"]["generate_audio"] is True
|
||||
assert captured["json"]["watermark"] is True
|
||||
assert captured["json"]["ratio"] == "1:1"
|
||||
assert captured["json"]["resolution"] == "480p"
|
||||
|
||||
|
||||
class TestGetDoubaoClientSingleton:
|
||||
def test_singleton_lazy_init(self):
|
||||
from packages.shared import ai_client
|
||||
|
||||
prev = ai_client._client
|
||||
try:
|
||||
ai_client._client = None
|
||||
c1 = ai_client.get_doubao_client()
|
||||
c2 = ai_client.get_doubao_client()
|
||||
assert c1 is c2
|
||||
assert isinstance(c1, ai_client.DoubaoClient)
|
||||
finally:
|
||||
ai_client._client = prev
|
||||
|
||||
|
||||
class TestVideoGenerationCancelled:
|
||||
def test_poll_cancelled_returns_none(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-can"}))
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock(json=MagicMock(return_value={"status": "cancelled"}))
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
@@ -741,22 +741,23 @@ class TestCreatePreviewRoute:
|
||||
assert resp.items[0].status == "pending"
|
||||
assert resp.items[0].variant_index == 0
|
||||
|
||||
def test_user_pending_limit_exceeded(self):
|
||||
"""用户待处理任务超限 → 429"""
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 3
|
||||
repo.count_pending_total.return_value = 5
|
||||
|
||||
def test_user_pending_limit_no_longer_rejects(self):
|
||||
"""#2098: 用户待处理任务超限不再 429 拒绝(预检查仅全局 503)。"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 100
|
||||
repo.count_pending_total.return_value = 5
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert exc_info.value.status_code == 429
|
||||
# 不是 429(用户级硬拒已移除)也不是 503(全局未超限)
|
||||
assert exc.value.status_code != 429
|
||||
assert exc.value.status_code != 503
|
||||
|
||||
def test_global_queue_full(self):
|
||||
"""全局队列满 → 503"""
|
||||
@@ -844,36 +845,32 @@ class TestCreatePreviewRoute:
|
||||
assert resp.total == 1
|
||||
assert resp.items[0].status == "failed"
|
||||
|
||||
def test_enqueue_raises_user_limit(self):
|
||||
"""safe_enqueue 抛出 UserPendingLimitExceeded → 429"""
|
||||
def test_enqueue_user_limit_exception_no_longer_returns_429(self):
|
||||
"""#2098: 即使 safe_enqueue 模拟抛 UserPendingLimitExceeded,也不再触发 429 整体拒绝,
|
||||
变体被标记 failed 后正常返回响应。"""
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 0
|
||||
repo.count_pending_total.return_value = 0
|
||||
|
||||
task = _make_task()
|
||||
from fastapi import HTTPException
|
||||
|
||||
def _set_failed_limit(error_message="", **_kwargs):
|
||||
def _set_failed(reason="", **_kw):
|
||||
task.status = GenerationTaskStatus.FAILED
|
||||
task.error_message = error_message or "待处理任务超限"
|
||||
task.error_message = reason
|
||||
|
||||
task.mark_failed.side_effect = _set_failed_limit
|
||||
|
||||
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
|
||||
MockUC.return_value.execute.return_value = task
|
||||
with patch(
|
||||
"app.api.routes.generation_preview.safe_enqueue_generation_task",
|
||||
side_effect=UserPendingLimitExceeded(user_id="u1", pending_count=4, limit=3),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
# 全部变体入队失败且错误消息含"待处理任务" → 429
|
||||
assert exc_info.value.status_code == 429
|
||||
task.mark_failed.side_effect = _set_failed
|
||||
repo.create.return_value = task
|
||||
with patch(
|
||||
"app.api.routes.generation_preview.safe_enqueue_generation_task",
|
||||
side_effect=UserPendingLimitExceeded(user_id="u1", pending_count=4, limit=20),
|
||||
):
|
||||
resp = create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert resp.total == 1
|
||||
assert resp.items[0].status == "failed"
|
||||
|
||||
def test_enqueue_raises_global_queue_full(self):
|
||||
"""safe_enqueue 抛出 GlobalQueueFull → 503"""
|
||||
|
||||
@@ -85,7 +85,7 @@ def _patch_pipeline_helpers(monkeypatch):
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestNoConfigBackwardCompat:
|
||||
def test_default_title_drawtext_white_top(self, monkeypatch):
|
||||
"""不传 title_config 时:白字、top、36@720 按 width 缩放到 64、y=71(PAD16+margin24)。"""
|
||||
"""不传 title_config 时:白字、top、48@720 按 width 缩放到 85、y=89(50@720p baseline)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
@@ -101,15 +101,15 @@ class TestNoConfigBackwardCompat:
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "fontcolor=0xffffff" in fc
|
||||
# 默认 position=top → y=71(scale(16+24)=71),不含 h-th
|
||||
assert "y=71" in fc
|
||||
# 默认 position=top → y=89(scale(50)=89,与 CPU/vfb 一致),不含 h-th
|
||||
assert "y=89" in fc
|
||||
assert "h-th" not in fc
|
||||
# 默认字号 36@720 经 1280/720 缩放 = 64
|
||||
assert "fontsize=64" in fc
|
||||
# 默认字号 48@720 经 1280/720 缩放 = 85
|
||||
assert "fontsize=85" in fc
|
||||
assert "text='默认标题'" in fc
|
||||
# 默认粗体:drawtext 无原生粗体,用同色描边 borderw=1@720(1280/720→scale=2)模拟
|
||||
assert "borderw=2" in fc
|
||||
assert "bordercolor=0xffffff" in fc
|
||||
# 默认粗体:黑色细描边 borderw=2@720(scale=4),与 CPU vfb 一致避免重影
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
|
||||
def test_no_subtitle_when_none(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
@@ -153,8 +153,8 @@ class TestTitleStylePassthrough:
|
||||
assert "fontcolor=0xff0000" in fc
|
||||
# 60@720 按 1280/720 缩放 = 107
|
||||
assert "fontsize=107" in fc
|
||||
# top 位置 y=71(scale(PAD16+margin_top24)=71)
|
||||
assert "y=71" in fc
|
||||
# top 位置默认 margin 50@720p → scale=89(未传 margin_top 用默认)
|
||||
assert "y=89" in fc
|
||||
assert "h-th" not in fc
|
||||
|
||||
def test_title_position_center(self, monkeypatch):
|
||||
@@ -218,9 +218,9 @@ class TestTitleStylePassthrough:
|
||||
# 阴影层 + 主字层 = 2 条 drawtext
|
||||
assert drawtext_count == 2
|
||||
joined = " ".join(plan.filter_complex)
|
||||
# shadow offset 3@720 经 1280/720 缩放 = 5;默认 position=top → y=71
|
||||
# shadow offset 3@720 经 1280/720 缩放 = 5;默认 position=top → y=89
|
||||
assert "x=(w-text_w)/2+5" in joined
|
||||
assert "y=71+5" in joined
|
||||
assert "y=89+5" in joined
|
||||
assert "h-th" not in joined
|
||||
|
||||
def test_title_font_override(self, monkeypatch):
|
||||
@@ -282,10 +282,10 @@ class TestTitleStylePassthrough:
|
||||
assert "(w-text_w)*0.3000" in fc
|
||||
assert "(h-text_h)*0.6000" in fc
|
||||
assert "h-th-" not in fc
|
||||
assert "y=71" not in fc
|
||||
assert "y=89" not in fc
|
||||
|
||||
def test_title_margin_top_respected(self, monkeypatch):
|
||||
"""margin_top 透传:100@720 → scale(16+100)=206@1280。"""
|
||||
"""margin_top 透传:100@720,叠加默认 50 → 150@720 → scale=267@1280。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
@@ -300,8 +300,8 @@ class TestTitleStylePassthrough:
|
||||
title_config={"text": "远离顶部", "position": "top", "margin_top": 100},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# scale(PAD16 + margin_top100) = 116*1280/720 = 206
|
||||
assert "y=206" in fc
|
||||
# 默认 50 + margin_top100 = 150@720 → scale=267
|
||||
assert "y=267" in fc
|
||||
|
||||
def test_title_default_bold_true(self, monkeypatch):
|
||||
"""不传 bold 时默认粗体:drawtext 用同色描边 borderw 模拟。"""
|
||||
@@ -319,9 +319,9 @@ class TestTitleStylePassthrough:
|
||||
title_config={"text": "粗体"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# bold=True 默认:同色描边 width=1@720 → scale=2,颜色与文字相同(白色)
|
||||
assert "borderw=2" in fc
|
||||
assert "bordercolor=0xffffff" in fc
|
||||
# bold=True 默认:黑色细描边 width=2@720 → scale=4,与 CPU vfb 一致
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
assert "text='粗体'" in fc
|
||||
|
||||
def test_title_bold_false_disables_faux_bold(self, monkeypatch):
|
||||
@@ -359,8 +359,8 @@ class TestTitleStylePassthrough:
|
||||
title_config={"text": "底部", "position": "bottom"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# bottom margin PAD16+24=40 → scale=71
|
||||
assert "y=h-th-71" in fc
|
||||
# bottom margin 50@720 → scale=89
|
||||
assert "y=h-th-89" in fc
|
||||
|
||||
def test_title_size_scales_by_width_not_height(self, monkeypatch):
|
||||
"""不同分辨率下同 @720 基准的 size 等比缩放:1080x1920 下 36→54。"""
|
||||
@@ -378,10 +378,10 @@ class TestTitleStylePassthrough:
|
||||
clip_volumes=[1.0],
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# default size 36@720 → 1080w = 54
|
||||
assert "fontsize=54" in fc
|
||||
# top margin 40@720 → 60
|
||||
assert "y=60" in fc
|
||||
# default size 48@720 → 1080w = 72
|
||||
assert "fontsize=72" in fc
|
||||
# top margin 50@720 → 75
|
||||
assert "y=75" in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -434,8 +434,8 @@ class TestSubtitlePassthrough:
|
||||
assert "fontcolor=0x0000ff" in fc
|
||||
# sub size 28@720 经 1280/720 缩放 = 50
|
||||
assert "fontsize=50" in fc
|
||||
# bottom margin 60@720 → 107
|
||||
assert "y=h-th-107" in fc
|
||||
# bottom margin 50@720 → 89
|
||||
assert "y=h-th-89" in fc
|
||||
|
||||
def test_subtitle_disabled_hides_subs(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
@@ -644,18 +644,18 @@ class TestCombinedConfig:
|
||||
bgm_config={"volume": 0.2, "fade_in": 0.5, "fade_out": 1.0},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 标题:size 50@720 → 89,top margin 40@720 → 71,stroke 2@720 → 4
|
||||
# 标题:size 50@720 → 89,top margin 50@720 → 89,stroke 2@720 → 4
|
||||
assert "text='主标题'" in fc
|
||||
assert "fontcolor=0xffff00" in fc
|
||||
assert "fontsize=89" in fc
|
||||
assert "y=71" in fc
|
||||
assert "y=89" in fc
|
||||
assert "borderw=4" in fc
|
||||
# 字幕:size 24@720 → 43,bottom margin 60@720 → 107
|
||||
# 字幕:size 24@720 → 43,bottom margin 50@720 → 89(默认不加粗)
|
||||
assert "text='成片全字幕'" in fc
|
||||
assert "fontsize=43" in fc
|
||||
assert "y=h-th-107" in fc
|
||||
assert "y=h-th-89" in fc
|
||||
assert "fontcolor=0xffffff" in fc
|
||||
# subtitle 默认 bold=True 也会加同色描边(width=1→2);但若主标题已有 stroke 不影响字幕独立 borderw=4(s 默认 2 像素@720→scale=4)
|
||||
# subtitle 默认 bold=False,无额外描边(用户未开 stroke)
|
||||
# BGM
|
||||
assert "volume=0.200" in fc
|
||||
assert "afade=t=in:st=0:d=0.50" in fc
|
||||
@@ -688,7 +688,7 @@ class TestEdgeCases:
|
||||
assert "drawtext=" in fc
|
||||
|
||||
def test_bold_title_increases_borderw(self, monkeypatch):
|
||||
"""bold=True 时若原无描边,自动加 borderw=1 用同色描边模拟加粗。"""
|
||||
"""bold=True 时若原无描边,自动加 borderw 黑色细描边模拟加粗(与 CPU vfb 一致,避免重影)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
@@ -703,9 +703,9 @@ class TestEdgeCases:
|
||||
title_config={"text": "粗体", "bold": True, "color": "#ff0000"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 加粗模拟 borderw>=1,且 bordercolor 跟字体色一致(0xff0000)
|
||||
assert "borderw=" in fc
|
||||
assert "bordercolor=0xff0000" in fc
|
||||
# 仿粗用黑色细描边 2@720→scale=4,不跟文字色(避免同色描边重影)
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
|
||||
def test_asr_and_static_subtitle_asr_wins(self, monkeypatch):
|
||||
"""同时传 static_subtitle_text 和 ASR segments 时,ASR 优先(不插入静态全文)。"""
|
||||
|
||||
+46
-119
@@ -1,4 +1,4 @@
|
||||
"""task_enqueue 单测 — 队列限流 + 安全入队逻辑."""
|
||||
"""task_enqueue 单测 — 队列限流 + 安全入队逻辑 (#2098)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -14,12 +14,8 @@ from app.core.task_enqueue import (
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
|
||||
# ── Fixtures / Helpers ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class MockRepository:
|
||||
"""Mock 任务仓储,用计数器模拟 pending 数量."""
|
||||
|
||||
def __init__(self, global_count: int = 0, user_count: int = 0):
|
||||
self._global = global_count
|
||||
self._user = user_count
|
||||
@@ -43,19 +39,12 @@ def make_mock_task(task_id: str = "task-1"):
|
||||
return task
|
||||
|
||||
|
||||
# ── check_queue_limits ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCheckQueueLimits:
|
||||
"""check_queue_limits 预检查限流."""
|
||||
|
||||
def test_below_limits_passes(self):
|
||||
repo = MockRepository(global_count=5, user_count=1)
|
||||
# 不抛异常就是通过
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_at_limit_raises(self):
|
||||
"""达到全局上限即拒绝."""
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=1)
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
@@ -67,210 +56,145 @@ class TestCheckQueueLimits:
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_at_limit_raises(self):
|
||||
"""达到用户上限即拒绝."""
|
||||
def test_user_at_limit_no_longer_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == USER_PENDING_LIMIT
|
||||
assert exc_info.value.limit == USER_PENDING_LIMIT
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_over_limit_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 1)
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_priority_over_user(self):
|
||||
"""全局和用户都超限时,优先抛全局异常."""
|
||||
repo = MockRepository(
|
||||
global_count=GLOBAL_PENDING_LIMIT + 1,
|
||||
user_count=USER_PENDING_LIMIT + 1,
|
||||
)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
def test_user_over_limit_no_longer_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_empty_user_id_skips_user_check(self):
|
||||
"""user_id 为空时跳过用户级检查."""
|
||||
repo = MockRepository(global_count=5, user_count=999)
|
||||
# 不抛异常 = 通过(只检查全局)
|
||||
check_queue_limits("", repo)
|
||||
|
||||
def test_custom_limits(self):
|
||||
"""支持自定义限流阈值."""
|
||||
repo = MockRepository(global_count=5, user_count=5)
|
||||
# 默认阈值下 user 5 > 3 会被拒
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("u1", repo)
|
||||
|
||||
# 自定义更高阈值就能通过
|
||||
check_queue_limits("u1", repo, user_pending_limit=10, global_pending_limit=10)
|
||||
|
||||
|
||||
# ── safe_enqueue_generation_task ──────────────────────────────────────────
|
||||
def test_custom_global_limit_still_honored(self):
|
||||
repo = MockRepository(global_count=15, user_count=999)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=10)
|
||||
check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=20)
|
||||
|
||||
|
||||
class TestSafeEnqueueGenerationTask:
|
||||
"""safe_enqueue_generation_task 安全入队."""
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_success_path(self, mock_celery):
|
||||
"""正常路径:入队前检查通过 → 发送Celery → 入队后检查通过."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once_with("worker.generate_video", args=[task.id])
|
||||
task.mark_failed.assert_not_called()
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_no_user_id_skips_user_check(self, mock_celery):
|
||||
"""不传 user_id 跳过用户级限流."""
|
||||
repo = MockRepository(global_count=1, user_count=999)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
assert result is True
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_precheck_global_over_marks_failed(self, mock_celery):
|
||||
"""入队前全局超限:标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT + 1, user_count=0)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
mock_celery.send_task.assert_not_called()
|
||||
assert repo.update_called == 1
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_precheck_user_over_marks_failed(self, mock_celery):
|
||||
"""入队前用户超限:标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 1)
|
||||
def test_precheck_user_over_still_enqueues(self, mock_celery, caplog):
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
mock_celery.send_task.assert_not_called()
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
task.mark_failed.assert_not_called()
|
||||
assert any("超过软上限" in r.message for r in caplog.records)
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_celery_send_false_returns_false(self, mock_celery):
|
||||
"""Celery 发送失败:返回 False,任务标记 failed."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
mock_celery.send_task.side_effect = Exception("celery down")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is False
|
||||
task.mark_failed.assert_called_once()
|
||||
assert "入队失败" in task.mark_failed.call_args[0][0]
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_celery_send_failure_update_also_fails(self, mock_celery):
|
||||
"""Celery 发送失败 + mark_failed 更新也失败:不崩溃."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
repo.update = MagicMock(side_effect=Exception("db down"))
|
||||
task = make_mock_task()
|
||||
mock_celery.send_task.side_effect = Exception("celery down")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is False
|
||||
# 不抛异常就是胜利
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_postcheck_global_over_rollback(self, mock_celery):
|
||||
"""入队后全局超限(并发竞态):回滚标记 failed,抛异常."""
|
||||
# 入队前刚好通过,但入队后再查发现超限
|
||||
call_count = [0]
|
||||
|
||||
def count_pending_total_side_effect():
|
||||
call_count[0] += 1
|
||||
if call_count[0] == 1: # 入队前检查
|
||||
return GLOBAL_PENDING_LIMIT # 等于上限,用 > 判断所以通过
|
||||
return GLOBAL_PENDING_LIMIT + 1 # 入队后再查,超限
|
||||
if call_count[0] == 1:
|
||||
return GLOBAL_PENDING_LIMIT
|
||||
return GLOBAL_PENDING_LIMIT + 1
|
||||
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=0)
|
||||
repo.count_pending_total = MagicMock(side_effect=count_pending_total_side_effect)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# 异常是 GlobalQueueFull 类型,且任务已被标记为 failed(含"入队后"原因)
|
||||
task.mark_failed.assert_called_once()
|
||||
assert "入队后" in task.mark_failed.call_args[0][0]
|
||||
mock_celery.send_task.assert_called_once()
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_postcheck_user_over_rollback(self, mock_celery):
|
||||
"""入队后用户超限:回滚标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
# 入队前用 > 判断,等于上限通过;入队后模拟并发超限
|
||||
original_user_count = repo.count_pending_by_user
|
||||
def test_postcheck_user_over_does_not_rollback(self, mock_celery, caplog):
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
call_count = [0]
|
||||
|
||||
def count_by_user_side_effect(user_id):
|
||||
call_count[0] += 1
|
||||
if call_count[0] <= 1: # 入队前
|
||||
return USER_PENDING_LIMIT # 用 > 判断,等于时通过
|
||||
return USER_PENDING_LIMIT + 1 # 入队后,超限
|
||||
if call_count[0] <= 1:
|
||||
return USER_PENDING_LIMIT
|
||||
return USER_PENDING_LIMIT + 1
|
||||
|
||||
repo.count_pending_by_user = MagicMock(side_effect=count_by_user_side_effect)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
task.mark_failed.assert_not_called()
|
||||
assert any("超软上限(入队后)" in r.message for r in caplog.records)
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_log_task_status_enabled(self, mock_celery):
|
||||
"""log_task_status=True 时日志中包含状态."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1", log_task_status=True)
|
||||
assert result is True
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_custom_limits_in_enqueue(self, mock_celery):
|
||||
"""自定义限流阈值用于入队检查."""
|
||||
repo = MockRepository(global_count=5, user_count=5)
|
||||
def test_custom_global_limit_in_enqueue(self, mock_celery):
|
||||
repo = MockRepository(global_count=15, user_count=999)
|
||||
task = make_mock_task()
|
||||
|
||||
# 默认阈值下用户 5 > 3 会被拒
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# 重置 mock 计数
|
||||
task.mark_failed.reset_mock()
|
||||
|
||||
# 调大阈值后通过
|
||||
result = safe_enqueue_generation_task(
|
||||
task,
|
||||
repo,
|
||||
user_id="user-1",
|
||||
user_pending_limit=10,
|
||||
global_pending_limit=10,
|
||||
)
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
|
||||
|
||||
# ── 异常类 ────────────────────────────────────────────────────────────────
|
||||
task2 = make_mock_task("task-2")
|
||||
mock_celery.reset_mock()
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task2, repo, user_id="user-1", user_pending_limit=999, global_pending_limit=10)
|
||||
|
||||
|
||||
class TestExceptionClasses:
|
||||
"""异常类消息格式."""
|
||||
|
||||
def test_user_pending_limit_message(self):
|
||||
exc = UserPendingLimitExceeded("u1", 5, 3)
|
||||
assert "u1" in str(exc)
|
||||
@@ -281,3 +205,6 @@ class TestExceptionClasses:
|
||||
exc = GlobalQueueFull(25, 20)
|
||||
assert "25" in str(exc)
|
||||
assert "20" in str(exc)
|
||||
|
||||
def test_user_pending_limit_constant(self):
|
||||
assert USER_PENDING_LIMIT == 20
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""任务队列限流防护单元测试。"""
|
||||
"""任务队列限流防护单元测试 (#2098: 用户级改为软上限,仅全局硬拒)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -19,17 +19,8 @@ from app.core.task_enqueue import (
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MockRepository:
|
||||
"""支持 pending 计数的 mock repository。
|
||||
|
||||
支持通过 set_pending 动态修改计数,用于模拟入队后计数变化的并发场景。
|
||||
"""
|
||||
|
||||
def __init__(self, user_pending: int = 0, global_pending: int = 0):
|
||||
self._user_pending = user_pending
|
||||
self._global_pending = global_pending
|
||||
@@ -47,7 +38,6 @@ class MockRepository:
|
||||
return task
|
||||
|
||||
def set_pending(self, *, user_pending: int | None = None, global_pending: int | None = None):
|
||||
"""动态修改 pending 计数,模拟并发场景。"""
|
||||
if user_pending is not None:
|
||||
self._user_pending = user_pending
|
||||
if global_pending is not None:
|
||||
@@ -67,58 +57,39 @@ class MockTask:
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_celery(monkeypatch):
|
||||
"""mock 掉 celery_app.send_task,避免真实发送。"""
|
||||
mock_send = MagicMock()
|
||||
monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send)
|
||||
return mock_send
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 常量导出测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_limit_constants_are_exported():
|
||||
"""限流阈值常量已导出,供业务代码引用。"""
|
||||
assert USER_PENDING_LIMIT == 3
|
||||
"""#2098: USER_PENDING_LIMIT 从 3 提到 20 作为软上限;GLOBAL_PENDING_LIMIT 保持 20 为硬上限。"""
|
||||
assert USER_PENDING_LIMIT == 20
|
||||
assert GLOBAL_PENDING_LIMIT == 20
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# check_queue_limits 单元测试(预检查用,>= 边界)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCheckQueueLimits:
|
||||
"""队列限流检查函数测试(预检查语义,>= 上限即拒绝)。"""
|
||||
"""check_queue_limits 预检查:仅全局硬上限拒绝,用户级改为软提示。"""
|
||||
|
||||
def test_normal_passes_through(self):
|
||||
"""正常范围内的任务不受限制。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_limit_exceeded_raises(self):
|
||||
"""用户 pending 超过上限抛 UserPendingLimitExceeded。"""
|
||||
repo = MockRepository(user_pending=4, global_pending=5)
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == 4
|
||||
assert exc_info.value.limit == 3
|
||||
def test_user_limit_exceeded_no_longer_raises(self):
|
||||
"""#2098: 用户 pending 超过软上限不再抛异常。"""
|
||||
repo = MockRepository(user_pending=100, global_pending=5)
|
||||
check_queue_limits("user-1", repo) # 不抛即通过
|
||||
|
||||
def test_user_at_limit_also_raises(self):
|
||||
"""用户 pending 刚好等于上限也拒绝(>= 边界)。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5)
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("user-1", repo)
|
||||
def test_user_at_limit_no_longer_raises(self):
|
||||
"""#2098: 用户 pending 等于软上限也不拒绝。"""
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_below_limit_passes(self):
|
||||
"""用户 pending 比上限少 1,通过。"""
|
||||
repo = MockRepository(user_pending=2, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_limit_exceeded_raises(self):
|
||||
"""全局 pending 超过上限抛 GlobalQueueFull。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=21)
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
@@ -126,82 +97,52 @@ class TestCheckQueueLimits:
|
||||
assert exc_info.value.limit == 20
|
||||
|
||||
def test_global_at_limit_also_raises(self):
|
||||
"""全局 pending 刚好等于上限也拒绝(>= 边界)。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_below_limit_passes(self):
|
||||
"""全局 pending 比上限少 1,通过。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=19)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_takes_priority_over_user(self):
|
||||
"""全局和用户都超限时,优先抛全局异常。"""
|
||||
repo = MockRepository(user_pending=5, global_pending=25)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_empty_user_id_skips_user_check(self):
|
||||
"""不传 user_id 时跳过用户级检查,只做全局检查。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
# 用户超限但不传 user_id → 全局未超限,应该通过
|
||||
check_queue_limits("", repo)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# safe_enqueue_generation_task 限流集成测试(入队前用 >,包含当前任务)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSafeEnqueueWithLimits:
|
||||
"""安全入队函数的限流功能测试。"""
|
||||
"""safe_enqueue_generation_task:用户超限仅 warning 仍入队;全局超限硬拒。"""
|
||||
|
||||
def test_normal_task_enqueues_successfully(self, mock_celery):
|
||||
"""正常任务入队成功,返回 True。"""
|
||||
repo = MockRepository(user_pending=0, global_pending=0)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
|
||||
# 成功入队后持久化 celery 消息 ID(#1714:孤儿清理据此 revoke/清队列)
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_user_limit_rejected_with_failed_status(self, mock_celery):
|
||||
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
|
||||
repo = MockRepository(user_pending=5, global_pending=5)
|
||||
def test_user_limit_exceeded_still_enqueues(self, mock_celery, caplog):
|
||||
"""#2098: 用户远超软上限仍入队,任务不被标记 failed。"""
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(user_pending=100, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
assert task.status == "failed"
|
||||
assert "限流" in task.error_message
|
||||
assert len(repo.updated_tasks) == 1
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending" # 没被标记 failed
|
||||
assert any("超过软上限" in r.message for r in caplog.records)
|
||||
|
||||
def test_user_at_limit_still_passes(self, mock_celery):
|
||||
"""用户 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。
|
||||
|
||||
与预检查的 >= 语义一致:预检查时 pending=3 拒绝(不能再加新的),
|
||||
但 safe_enqueue 被调用时任务已是 pending(就是第3个),
|
||||
pending=3 不满足 >3,所以通过。
|
||||
"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5)
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_user_one_over_limit_rejected(self, mock_celery):
|
||||
"""用户 pending = limit + 1:超限被拒。"""
|
||||
repo = MockRepository(user_pending=4, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
|
||||
def test_global_limit_rejected_with_failed_status(self, mock_celery):
|
||||
"""全局超限:任务标记为 failed,抛 GlobalQueueFull。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=21)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
@@ -211,7 +152,6 @@ class TestSafeEnqueueWithLimits:
|
||||
assert len(repo.updated_tasks) == 1
|
||||
|
||||
def test_global_at_limit_still_passes(self, mock_celery):
|
||||
"""全局 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
@@ -219,7 +159,6 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_no_user_id_skips_user_limit(self, mock_celery):
|
||||
"""不传 user_id 时跳过用户级限流,只做全局检查。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
@@ -227,7 +166,6 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_no_user_id_still_checks_global(self, mock_celery):
|
||||
"""不传 user_id 时全局超限仍然被拦。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=25)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
@@ -235,121 +173,74 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_not_called()
|
||||
|
||||
def test_default_limits_match_constants(self, mock_celery):
|
||||
"""默认配置与导出常量一致。"""
|
||||
# 刚好在默认限制内(limit - 1)
|
||||
repo = MockRepository(user_pending=2, global_pending=19)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
|
||||
def test_update_failure_does_not_crash(self, mock_celery):
|
||||
"""repository.update 失败也不崩溃,异常继续向上抛。"""
|
||||
|
||||
class BadRepo(MockRepository):
|
||||
def update(self, task):
|
||||
raise RuntimeError("db down")
|
||||
|
||||
repo = BadRepo(user_pending=5, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
# 仍然抛 UserPendingLimitExceeded,不会被 update 失败掩盖
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
# 任务状态还是变了(内存里改了)
|
||||
assert task.status == "failed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 入队后最终校验(并发竞态兜底)测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPostEnqueueFinalCheck:
|
||||
"""入队后最终校验:模拟并发场景,Celery发送后计数增加被兜住。"""
|
||||
"""入队后校验:仅全局超限回滚;用户超限仅 warning。"""
|
||||
|
||||
def test_post_enqueue_global_overflow_rollback(self, mock_celery):
|
||||
"""并发场景:入队前检查通过,但发送Celery后全局计数超限 → 回滚为failed。
|
||||
|
||||
模拟两个请求同时通过入队前检查(都查到 global=19),
|
||||
都创建了任务(DB里变成 21),先发送Celery的那个在最终校验时被兜住。
|
||||
"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20) # 入队前:20 > 20?否
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
|
||||
# 模拟发送Celery后,另一个并发请求也创建了任务,全局变成21
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(global_pending=21)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# Celery 确实发出去了(兜底不撤销 Celery,只回滚 DB 状态)
|
||||
mock_celery.assert_called_once()
|
||||
# 任务被标记为 failed
|
||||
assert task.status == "failed"
|
||||
assert "入队后" in task.error_message
|
||||
assert exc_info.value.pending_count == 21
|
||||
assert len(repo.updated_tasks) == 1
|
||||
|
||||
def test_post_enqueue_user_overflow_rollback(self, mock_celery):
|
||||
"""并发场景:入队前检查通过,但发送Celery后用户计数超限 → 回滚为failed。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5) # 入队前:3 > 3?否
|
||||
def test_post_enqueue_user_overflow_no_rollback(self, mock_celery, caplog):
|
||||
"""#2098: 入队后用户超软上限仅 warning,不回滚。"""
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=4)
|
||||
repo.set_pending(user_pending=USER_PENDING_LIMIT + 1)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "failed"
|
||||
assert "入队后" in task.error_message
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == 4
|
||||
|
||||
def test_post_enqueue_global_priority_over_user(self, mock_celery):
|
||||
"""入队后校验:全局和用户都超限时,优先抛全局异常。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=5, global_pending=22)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert task.status == "failed"
|
||||
|
||||
def test_post_enqueue_no_change_still_passes(self, mock_celery):
|
||||
"""入队后计数没变 → 正常通过,不回滚。"""
|
||||
repo = MockRepository(user_pending=2, global_pending=10)
|
||||
task = MockTask("task-1")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending" # 状态没变
|
||||
# 入队成功后持久化 celery_task_id(#1714),业务状态不变
|
||||
assert task.status == "pending" # 不回滚
|
||||
assert any("超软上限(入队后)" in r.message for r in caplog.records)
|
||||
|
||||
def test_post_enqueue_no_change_still_passes(self, mock_celery):
|
||||
repo = MockRepository(user_pending=2, global_pending=10)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending"
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
|
||||
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=15, global_pending=5) # 用户超限但全局没超
|
||||
repo.set_pending(user_pending=15, global_pending=5)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
assert result is True # 用户级不检查,全局没超限 → 通过
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_user_pending_limit_exceeded_class_still_exists():
|
||||
"""UserPendingLimitExceeded 保留用于兼容历史 import/except(#2098 后不再主动 raise)。"""
|
||||
exc = UserPendingLimitExceeded("u1", 5, 3)
|
||||
assert exc.user_id == "u1"
|
||||
assert exc.pending_count == 5
|
||||
assert exc.limit == 3
|
||||
|
||||
@@ -0,0 +1,360 @@
|
||||
"""video_analyzer(#2051)单元测试。
|
||||
|
||||
覆盖:
|
||||
- 映射函数(运镜→ken_burns、转场→xfade、色调→video_filter、BPM→BGM)
|
||||
- schema 常量与导出
|
||||
- Farneback 光流运镜判定(合成光流场)
|
||||
- BPM 档位映射
|
||||
- 降级路径(ffmpeg/cv2/librosa 不可用)
|
||||
- 临时目录清理
|
||||
- analyze_video_style 入口在无素材时返回最小 style_guide 不抛
|
||||
- build_render_params_for_clip 聚合输出
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import ( # noqa: E402
|
||||
COLOR_FILTER_PRESETS,
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
TRANSITION_TO_XFADE,
|
||||
ShotBoundary,
|
||||
_detect_camera_movement,
|
||||
_pace_from_bpm,
|
||||
_rule_based_style_guide,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
# ── 映射函数 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_constants_exported():
|
||||
assert MAX_REFERENCE_DURATION_SEC == 60
|
||||
assert MAX_REFERENCE_SIZE_MB == 100
|
||||
assert DEFAULT_ANALYSIS_TIMEOUT == 60
|
||||
assert "style_name" in STYLE_GUIDE_SCHEMA
|
||||
assert "ken_burns_params" in STYLE_GUIDE_SCHEMA
|
||||
assert "video_filter_eq_params" in STYLE_GUIDE_SCHEMA
|
||||
|
||||
|
||||
# ── 运镜→ken_burns ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"movement,expected_type",
|
||||
[
|
||||
("static", "static"),
|
||||
("push_in", "zoom"),
|
||||
("zoom_in", "zoom"),
|
||||
("pull_out", "zoom"),
|
||||
("pan_left", "pan"),
|
||||
("pan_right", "pan"),
|
||||
("tilt_up", "pan+zoom"),
|
||||
("track_left", "pan"),
|
||||
],
|
||||
)
|
||||
def test_map_camera_to_ken_burns_types(movement, expected_type):
|
||||
kb = map_camera_to_ken_burns(movement)
|
||||
assert kb["type"] == expected_type
|
||||
# zoom/pan 类必须有 zoom_start/zoom_end
|
||||
assert 0.8 <= kb["zoom_start"] <= 1.3
|
||||
assert 0.8 <= kb["zoom_end"] <= 1.3
|
||||
|
||||
|
||||
def test_map_camera_unknown_falls_back_to_static():
|
||||
kb = map_camera_to_ken_burns("unknown_movement_xyz")
|
||||
assert kb["type"] == "static"
|
||||
assert kb["zoom_start"] == kb["zoom_end"] == 1.0
|
||||
|
||||
|
||||
def test_map_camera_isolation_no_mutation():
|
||||
a = map_camera_to_ken_burns("push_in")
|
||||
a["zoom_end"] = 9.99
|
||||
b = map_camera_to_ken_burns("push_in")
|
||||
assert b["zoom_end"] != 9.99
|
||||
|
||||
|
||||
# ── 转场→xfade ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ttype,expected",
|
||||
[
|
||||
("hard_cut", "cut"),
|
||||
("cross_dissolve", "dissolve"),
|
||||
("fade", "fade"),
|
||||
("fade_black", "fadeblack"),
|
||||
("zoom_whip", "zoom"),
|
||||
("wipe_left", "wipeleft"),
|
||||
("slide_right", "slideleft"),
|
||||
],
|
||||
)
|
||||
def test_map_transition(ttype, expected):
|
||||
assert map_transition_to_xfade(ttype) == expected
|
||||
|
||||
|
||||
def test_map_transition_unknown_falls_back_to_cut():
|
||||
assert map_transition_to_xfade("some_random_transition") == "cut"
|
||||
|
||||
|
||||
# ── 色调→video_filter ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name", ["none", "warm_vintage", "cool_fresh", "high_contrast", "soft_pastel", "dramatic_cinematic"]
|
||||
)
|
||||
def test_map_color_presets_available(name):
|
||||
p = map_color_to_video_filter(name)
|
||||
assert isinstance(p, dict)
|
||||
# 所有预设必须能被 FFmpeg eq/colorchannelmixer 消费:eq 是 dict,ccm 是 dict
|
||||
assert "eq" in p or p == {} or "colorchannelmixer" in p
|
||||
|
||||
|
||||
def test_map_color_unknown_is_none_preset():
|
||||
p = map_color_to_video_filter("not_a_real_filter")
|
||||
assert p == {}
|
||||
|
||||
|
||||
def test_map_color_isolation():
|
||||
a = map_color_to_video_filter("warm_vintage")
|
||||
a["eq"]["brightness"] = 9.99
|
||||
b = map_color_to_video_filter("warm_vintage")
|
||||
assert b["eq"]["brightness"] != 9.99
|
||||
|
||||
|
||||
# ── BPM → BGM ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bpm,expected",
|
||||
[
|
||||
(0, 90),
|
||||
(70, 70),
|
||||
(120, 120),
|
||||
(200, 180),
|
||||
(30, 60),
|
||||
],
|
||||
)
|
||||
def test_map_bgm_bpm(bpm, expected):
|
||||
assert map_bgm_bpm(bpm) == expected
|
||||
|
||||
|
||||
def test_pace_from_bpm_buckets():
|
||||
assert _pace_from_bpm(120) == "fast_cut"
|
||||
assert _pace_from_bpm(110) == "fast_cut"
|
||||
assert _pace_from_bpm(90) == "medium"
|
||||
assert _pace_from_bpm(80) == "medium"
|
||||
assert _pace_from_bpm(60) == "slow_cinematic"
|
||||
assert _pace_from_bpm(0) == "medium"
|
||||
|
||||
|
||||
# ── Farneback 光流→运镜(合成光流) ───────────────────────────────────
|
||||
|
||||
|
||||
def _make_flow(dx: float, dy: float, w: int = 60, h: int = 40, zoom: float = 0.0):
|
||||
"""构造一个合成光流场:整体平移(dx,dy)+径向发散(zoom>0=zoom in,<0=out)。"""
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
fx = dx + (xs - cx) * zoom
|
||||
fy = dy + (ys - cy) * zoom
|
||||
return np.stack([fx, fy], axis=-1).astype(np.float32)
|
||||
|
||||
|
||||
def test_detect_movement_static():
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "static"
|
||||
assert i == "low"
|
||||
|
||||
|
||||
def test_detect_movement_pan_right():
|
||||
flow = _make_flow(2.0, 0.0)
|
||||
m, i = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_right"
|
||||
assert i in ("medium", "high")
|
||||
|
||||
|
||||
def test_detect_movement_pan_left():
|
||||
flow = _make_flow(-2.0, 0.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "pan_left"
|
||||
|
||||
|
||||
def test_detect_movement_tilt_down():
|
||||
flow = _make_flow(0.0, 2.0)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "tilt_down"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_in_radial():
|
||||
# 径向向外发散 = zoom in
|
||||
flow = _make_flow(0.0, 0.0, zoom=0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_in"
|
||||
|
||||
|
||||
def test_detect_movement_zoom_out_radial():
|
||||
flow = _make_flow(0.0, 0.0, zoom=-0.08)
|
||||
m, _ = _detect_camera_movement(flow, 60, 40)
|
||||
assert m == "zoom_out"
|
||||
|
||||
|
||||
# ── 规则合成 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_shots(n=4):
|
||||
return [
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=float(i * 3),
|
||||
end_sec=float((i + 1) * 3),
|
||||
movement=["static", "push_in", "pan_left", "zoom_in"][i],
|
||||
intensity=["low", "medium", "low", "high"][i],
|
||||
transition="hard_cut",
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
|
||||
CAMERA_TO_KEN_BURNS_DIRS = {
|
||||
"zoom_in_slow",
|
||||
"zoom_out_slow",
|
||||
"pan_left_slow",
|
||||
"pan_right_slow",
|
||||
"zoom_in_medium",
|
||||
"zoom_out_medium",
|
||||
"diagonal_push",
|
||||
"static",
|
||||
}
|
||||
|
||||
|
||||
def test_rule_based_style_guide_structure():
|
||||
shots = _sample_shots()
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={"color_filter": "warm_vintage"})
|
||||
# 关键字段存在且类型正确
|
||||
assert sg["shot_count"] == 4
|
||||
assert sg["pace"] == "fast_cut"
|
||||
assert sg["bpm"] == 120
|
||||
assert sg["avg_shot_duration"] == 3.0
|
||||
assert len(sg["camera_movements"]) == 4
|
||||
assert sg["color_filter"] == "warm_vintage"
|
||||
assert "eq" in sg["video_filter_eq_params"]
|
||||
assert "default" in sg["ken_burns_params"]
|
||||
assert isinstance(sg["transition_map"], dict)
|
||||
assert isinstance(sg["ken_burns_direction_hint"], str) and sg["ken_burns_direction_hint"]
|
||||
|
||||
|
||||
# ── build_render_params_for_clip 聚合 ─────────────────────────────────
|
||||
|
||||
|
||||
def test_build_render_params_for_clip_shape():
|
||||
sg = _rule_based_style_guide(_sample_shots(), bpm=95, vlm={"color_filter": "cool_fresh"})
|
||||
p0 = build_render_params_for_clip(0, sg, duration_sec=3.0)
|
||||
assert "ken_burns" in p0
|
||||
assert "transition" in p0
|
||||
assert "video_filter" in p0
|
||||
assert p0["bgm_bpm_hint"] == 95
|
||||
assert p0["duration_sec"] == 3.0
|
||||
# clip 1 是 push_in → zoom
|
||||
p1 = build_render_params_for_clip(1, sg)
|
||||
assert p1["ken_burns"]["type"] == "zoom"
|
||||
|
||||
|
||||
def test_build_render_params_high_intensity_amplifies():
|
||||
shots = _sample_shots() # shot 3 = zoom_in/high
|
||||
sg = _rule_based_style_guide(shots, bpm=120, vlm={})
|
||||
p3 = build_render_params_for_clip(3, sg)
|
||||
base = map_camera_to_ken_burns("zoom_in")
|
||||
assert p3["ken_burns"]["zoom_end"] > base["zoom_end"]
|
||||
|
||||
|
||||
# ── 降级与容错 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_analyze_with_nonexistent_file_returns_minimum_guide():
|
||||
sg = analyze_video_style("/nonexistent/path/fake_video.mp4")
|
||||
assert isinstance(sg, dict)
|
||||
assert "style_name" in sg
|
||||
assert sg["shot_count"] == 0
|
||||
# 不抛异常且字段完整
|
||||
|
||||
|
||||
def test_analyze_invalid_style_strength_defaults_to_medium():
|
||||
# 即使视频不存在,也应被规范化为 medium 并写入返回值
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("fake", style_strength="banana")
|
||||
assert sg.get("style_strength", "medium") == "medium"
|
||||
|
||||
|
||||
def test_temp_dir_cleaned_up_after_run():
|
||||
"""用临时真实空文件模拟本地路径,确认 frames 临时目录被清理。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "fake.mp4"
|
||||
fake.write_bytes(b"")
|
||||
# 抽帧会失败(ffmpeg 对空文件失败),但应全程不抛且临时目录 rmtree
|
||||
# 直接 mock _ensure_local_video 回传不存在的文件,走 _probe_duration=0 降级路径
|
||||
with patch("apps.worker.viral_video.video_analyzer._ensure_local_video", return_value=None):
|
||||
sg = analyze_video_style("proto://fake", style_strength="light")
|
||||
assert "style_name" in sg
|
||||
|
||||
|
||||
def test_ffmpeg_failure_falls_back_to_vlm_only_path():
|
||||
"""模拟 ffmpeg 抽帧失败,仍能返回 style_guide。"""
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
fake = Path(td) / "ref.mp4"
|
||||
fake.write_bytes(b"not a real video")
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._extract_keyframes", side_effect=RuntimeError("ffmpeg exploded")
|
||||
):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_shots") as mock_shots:
|
||||
mock_shots.return_value = [ShotBoundary(0, 0.0, 3.0)]
|
||||
with patch("apps.worker.viral_video.video_analyzer._analyze_movements"):
|
||||
with patch("apps.worker.viral_video.video_analyzer._detect_bpm", return_value=90):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._vlm_analyze_frames",
|
||||
return_value={"color_filter": "none"},
|
||||
):
|
||||
with patch(
|
||||
"apps.worker.viral_video.video_analyzer._llm_synthesize",
|
||||
side_effect=lambda shots, bpm, vlm, ss: _rule_based_style_guide(shots, bpm, vlm),
|
||||
):
|
||||
sg = analyze_video_style(str(fake))
|
||||
assert sg["bpm"] == 90
|
||||
assert sg["shot_count"] == 1
|
||||
|
||||
|
||||
# ── 转场映射完整性 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_transition_map_covers_observed_types():
|
||||
for t in ("hard_cut", "cross_dissolve", "fade_black", "zoom_whip"):
|
||||
assert t in TRANSITION_TO_XFADE
|
||||
|
||||
|
||||
# ── 预设完整性 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_color_filter_preset_keys_are_safe_for_ffmpeg():
|
||||
for name, preset in COLOR_FILTER_PRESETS.items():
|
||||
if preset == {}:
|
||||
continue
|
||||
# eq 所有值都是数字
|
||||
for k, v in preset.get("eq", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.eq.{k} not numeric"
|
||||
for k, v in preset.get("colorchannelmixer", {}).items():
|
||||
assert isinstance(v, (int, float)), f"{name}.ccm.{k} not numeric"
|
||||
Executable
+517
@@ -0,0 +1,517 @@
|
||||
"""爆款视频模块单元测试。
|
||||
|
||||
覆盖范围:
|
||||
- 领域实体状态机转换
|
||||
- Repository CRUD
|
||||
- API 端点(6 个)
|
||||
- Celery 编排器流水线
|
||||
- Schema 校验
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
FusionLevel,
|
||||
StyleStrength,
|
||||
ViralVideoJob,
|
||||
ViralVideoStage,
|
||||
ViralVideoStatus,
|
||||
)
|
||||
|
||||
# ── 领域模型测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoStatus:
|
||||
"""状态枚举测试。"""
|
||||
|
||||
def test_status_values(self):
|
||||
assert ViralVideoStatus.PENDING == "pending"
|
||||
assert ViralVideoStatus.RUNNING == "running"
|
||||
assert ViralVideoStatus.WAIT_USER_CONFIRM == "wait_user_confirm"
|
||||
assert ViralVideoStatus.COMPLETED == "completed"
|
||||
assert ViralVideoStatus.FAILED == "failed"
|
||||
assert ViralVideoStatus.CANCELLED == "cancelled"
|
||||
|
||||
def test_terminal_statuses(self):
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED).is_terminal
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.FAILED).is_terminal
|
||||
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.CANCELLED).is_terminal
|
||||
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING).is_terminal
|
||||
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING).is_terminal
|
||||
|
||||
|
||||
class TestViralVideoJobStateTransitions:
|
||||
"""状态机转换测试。"""
|
||||
|
||||
def test_mark_running_from_pending(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
job.mark_running()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
assert job.started_at is not None
|
||||
|
||||
def test_mark_running_from_running(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_running()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
|
||||
def test_mark_running_from_completed_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
|
||||
with pytest.raises(ValueError, match="Cannot transition"):
|
||||
job.mark_running()
|
||||
|
||||
def test_mark_wait_user_confirm(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
intent = {"intent": "推广", "key_messages": ["卖点1"]}
|
||||
job.mark_wait_user_confirm(intent)
|
||||
assert job.status == ViralVideoStatus.WAIT_USER_CONFIRM
|
||||
assert job.intent_result == intent
|
||||
|
||||
def test_mark_wait_user_confirm_from_non_running_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING)
|
||||
with pytest.raises(ValueError, match="Cannot transition"):
|
||||
job.mark_wait_user_confirm({})
|
||||
|
||||
def test_resume_from_confirm(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
job.resume_from_confirm()
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
|
||||
def test_resume_from_non_confirm_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
with pytest.raises(ValueError, match="Cannot resume"):
|
||||
job.resume_from_confirm()
|
||||
|
||||
def test_mark_completed(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_completed("https://oss.example.com/video.mp4")
|
||||
assert job.status == ViralVideoStatus.COMPLETED
|
||||
assert job.result_video_url == "https://oss.example.com/video.mp4"
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_mark_failed(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_failed("渲染超时")
|
||||
assert job.status == ViralVideoStatus.FAILED
|
||||
assert job.error_msg == "渲染超时"
|
||||
|
||||
def test_mark_cancelled(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
|
||||
job.mark_cancelled()
|
||||
assert job.status == ViralVideoStatus.CANCELLED
|
||||
|
||||
def test_mark_cancelled_from_terminal_raises(self):
|
||||
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
|
||||
with pytest.raises(ValueError, match="Cannot cancel"):
|
||||
job.mark_cancelled()
|
||||
|
||||
|
||||
class TestViralVideoJobDefaults:
|
||||
"""默认值测试。"""
|
||||
|
||||
def test_default_values(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
assert job.images == []
|
||||
assert job.industry == ""
|
||||
assert job.duration == 30
|
||||
assert job.fusion_level == FusionLevel.AI_POLISH
|
||||
assert job.style_strength == StyleStrength.MEDIUM
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.credits_cost == 0
|
||||
assert job.retry_count == 0
|
||||
assert job.result_video_url == ""
|
||||
assert job.error_msg == ""
|
||||
|
||||
def test_credits_cost_constant(self):
|
||||
assert CREDITS_VIRAL_VIDEO_COST == 50
|
||||
|
||||
|
||||
class TestViralVideoStage:
|
||||
"""阶段枚举测试。"""
|
||||
|
||||
def test_all_stages_have_labels(self):
|
||||
for stage in ViralVideoStage:
|
||||
assert stage in STAGE_LABELS, f"Stage {stage} missing label"
|
||||
|
||||
def test_stage_order(self):
|
||||
expected_order = [
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"review",
|
||||
"tts",
|
||||
"bgm_select",
|
||||
"rendering",
|
||||
"musetalk",
|
||||
"uploading",
|
||||
]
|
||||
actual_order = [s.value for s in ViralVideoStage]
|
||||
assert actual_order == expected_order
|
||||
|
||||
|
||||
# ── Schema 校验测试 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoSchemas:
|
||||
"""Pydantic Schema 校验测试。"""
|
||||
|
||||
def test_create_request_valid(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
req = CreateViralVideoRequest(images=["https://example.com/img.jpg"])
|
||||
assert req.images == ["https://example.com/img.jpg"]
|
||||
assert req.fusion_level == "ai_polish"
|
||||
assert req.style_strength == "medium"
|
||||
assert req.duration == 30
|
||||
|
||||
def test_create_request_empty_images_raises(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(images=[])
|
||||
|
||||
def test_create_request_invalid_fusion_level(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(
|
||||
images=["https://example.com/img.jpg"],
|
||||
fusion_level="invalid_level",
|
||||
)
|
||||
|
||||
def test_create_request_invalid_style_strength(self):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
CreateViralVideoRequest(
|
||||
images=["https://example.com/img.jpg"],
|
||||
style_strength="ultra",
|
||||
)
|
||||
|
||||
def test_confirm_intent_request_defaults(self):
|
||||
from app.schemas.viral_video import ConfirmIntentRequest
|
||||
|
||||
req = ConfirmIntentRequest()
|
||||
assert req.confirmed_copy == ""
|
||||
assert req.adjustments == ""
|
||||
|
||||
def test_analyze_style_request(self):
|
||||
from app.schemas.viral_video import AnalyzeStyleRequest
|
||||
|
||||
req = AnalyzeStyleRequest(reference_video_url="https://example.com/video.mp4")
|
||||
assert req.reference_video_url == "https://example.com/video.mp4"
|
||||
|
||||
def test_ws_progress_event(self):
|
||||
from app.schemas.viral_video import WSProgressEvent
|
||||
|
||||
event = WSProgressEvent(
|
||||
job_id="abc123",
|
||||
stage="image_analysis",
|
||||
progress=10.0,
|
||||
message="正在分析图片",
|
||||
)
|
||||
assert event.type == "viral_video:progress"
|
||||
assert event.job_id == "abc123"
|
||||
assert event.progress == 10.0
|
||||
|
||||
|
||||
# ── Repository 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoRepository:
|
||||
"""SQLAlchemy Repository CRUD 测试(使用内存数据库)。"""
|
||||
|
||||
@pytest.fixture
|
||||
def db_session(self):
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
SessionLocal = sessionmaker(bind=engine)
|
||||
session = SessionLocal()
|
||||
yield session
|
||||
session.close()
|
||||
|
||||
def test_save_and_get(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
job = ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
repo.save(job)
|
||||
|
||||
fetched = repo.get(job.id)
|
||||
assert fetched is not None
|
||||
assert fetched.id == job.id
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_get_nonexistent(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
assert repo.get("nonexistent-id") is None
|
||||
|
||||
def test_list_by_user(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
for i in range(3):
|
||||
job = ViralVideoJob(user_id="user-001", industry=f"行业{i}")
|
||||
repo.save(job)
|
||||
|
||||
# 另一个用户的任务
|
||||
other_job = ViralVideoJob(user_id="user-002", industry="其他")
|
||||
repo.save(other_job)
|
||||
|
||||
jobs = repo.list_by_user("user-001")
|
||||
assert len(jobs) == 3
|
||||
assert all(j.user_id == "user-001" for j in jobs)
|
||||
|
||||
def test_update_status(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
job = ViralVideoJob(user_id="user-001")
|
||||
repo.save(job)
|
||||
|
||||
job.mark_running()
|
||||
repo.update(job)
|
||||
|
||||
fetched = repo.get(job.id)
|
||||
assert fetched.status == ViralVideoStatus.RUNNING
|
||||
assert fetched.started_at is not None
|
||||
|
||||
def test_count_pending_by_user(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
|
||||
repo = SQLAlchemyViralVideoJobRepository(db_session)
|
||||
# 2 个 pending
|
||||
for _ in range(2):
|
||||
repo.save(ViralVideoJob(user_id="user-001"))
|
||||
# 1 个 completed
|
||||
completed = ViralVideoJob(user_id="user-001", status=ViralVideoStatus.COMPLETED)
|
||||
repo.save(completed)
|
||||
|
||||
assert repo.count_pending_by_user("user-001") == 2
|
||||
|
||||
def test_style_template_repo(self, db_session):
|
||||
from packages.adapters.sqlalchemy_impl.models import ViralVideoStyleTemplateModel
|
||||
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
|
||||
# 插入模板
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
description="适合快消品",
|
||||
style_config={"cut_speed": "fast"},
|
||||
is_system=True,
|
||||
sort_order=1,
|
||||
)
|
||||
db_session.add(tpl)
|
||||
db_session.commit()
|
||||
|
||||
repo = SQLAlchemyViralVideoStyleTemplateRepository(db_session)
|
||||
templates = repo.list_all()
|
||||
assert len(templates) == 1
|
||||
assert templates[0]["name"] == "快节奏"
|
||||
|
||||
fetched = repo.get("tpl-001")
|
||||
assert fetched is not None
|
||||
assert fetched["style_config"] == {"cut_speed": "fast"}
|
||||
|
||||
|
||||
# ── Celery 编排器测试 ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestViralVideoPipeline:
|
||||
"""编排器流水线测试。"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_job(self):
|
||||
return ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg", "https://img.com/2.jpg"],
|
||||
industry="美妆",
|
||||
target_customer="年轻女性",
|
||||
marketing_purpose="品牌推广",
|
||||
duration=30,
|
||||
user_copy_text="这款产品超好用",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
|
||||
@patch("packages.shared.ai_service.call_vision")
|
||||
def test_image_analysis_step(self, mock_vision, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||
|
||||
mock_vision.return_value = {"name": "口红", "features": ["持久", "滋润"]}
|
||||
result = _step_image_analysis(mock_job)
|
||||
assert "products" in result
|
||||
assert len(result["products"]) == 2 # 两张图片
|
||||
|
||||
@patch("packages.shared.ai_service.call_vision")
|
||||
def test_image_analysis_fallback(self, mock_vision, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||
|
||||
# 模拟 call_vision 不存在
|
||||
mock_vision.side_effect = ImportError("no module")
|
||||
result = _step_image_analysis(mock_job)
|
||||
assert "products" in result
|
||||
|
||||
def test_video_analysis_no_reference(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
# 没有参考视频
|
||||
mock_job.reference_video_url = ""
|
||||
result = _step_video_analysis(mock_job)
|
||||
assert result is None
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_intent_parsing(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_intent_parsing
|
||||
|
||||
mock_llm.return_value = {"intent": "推广口红", "tone": "活泼"}
|
||||
result = _step_intent_parsing(mock_job, {"products": []})
|
||||
assert "intent" in result
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_copy_fusion_ai_polish(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_copy_fusion
|
||||
|
||||
mock_llm.return_value = "融合后的文案内容"
|
||||
result = _step_copy_fusion(mock_job, {"intent": "推广"}, {"products": []})
|
||||
assert isinstance(result, str)
|
||||
assert len(result) > 0
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_storyboard_generation(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
|
||||
|
||||
mock_llm.return_value = [
|
||||
{"order": 0, "type": "product_shot", "duration": 10},
|
||||
{"order": 1, "type": "closing", "duration": 5},
|
||||
]
|
||||
result = _step_storyboard(mock_job, "测试文案", {})
|
||||
assert isinstance(result, list)
|
||||
assert len(result) == 2
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_review_pass(self, mock_llm, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_review
|
||||
|
||||
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
|
||||
result = _step_review(mock_job, "测试文案", [])
|
||||
assert result["passed"] is True
|
||||
|
||||
def test_bgm_select(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
# P1: BGM 素材未就绪前 _step_bgm_select 统一返回 None(跳过 BGM 混音)
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
bgm = _step_bgm_select(mock_job)
|
||||
assert bgm is None
|
||||
|
||||
def test_bgm_select_default(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
mock_job.bgm_preference = ""
|
||||
bgm = _step_bgm_select(mock_job)
|
||||
assert bgm is None
|
||||
|
||||
|
||||
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPipelineIntegration:
|
||||
"""流水线端到端集成测试(mock 外部依赖)。"""
|
||||
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_upload")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_render")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_bgm_select")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_review")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_storyboard")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_copy_fusion")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_intent_parsing")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_video_analysis")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
|
||||
@patch("apps.worker.worker_app.tasks.viral_video._emit_progress")
|
||||
def test_resume_pipeline_completes(
|
||||
self,
|
||||
mock_emit,
|
||||
mock_get_repo,
|
||||
mock_img_analysis,
|
||||
mock_video_analysis,
|
||||
mock_intent,
|
||||
mock_copy_fusion,
|
||||
mock_storyboard,
|
||||
mock_review,
|
||||
mock_tts,
|
||||
mock_bgm,
|
||||
mock_render,
|
||||
mock_upload,
|
||||
):
|
||||
"""测试 resume 流水线能从确认状态走到完成。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import (
|
||||
resume_viral_video_pipeline,
|
||||
)
|
||||
|
||||
# 构造 mock job
|
||||
job = ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
status=ViralVideoStatus.RUNNING,
|
||||
intent_result={"intent": "推广"},
|
||||
)
|
||||
|
||||
mock_repo = MagicMock()
|
||||
mock_session = MagicMock()
|
||||
mock_get_repo.return_value = (mock_session, mock_repo, job)
|
||||
|
||||
# 设置各步骤返回值
|
||||
mock_copy_fusion.return_value = "融合文案"
|
||||
mock_storyboard.return_value = [{"order": 0, "duration": 10}]
|
||||
mock_review.return_value = {"passed": True, "score": 90}
|
||||
mock_tts.return_value = None # P1: TTS 返回 Path|None,mock 用 None 跳过混音
|
||||
mock_bgm.return_value = None # P1: BGM 未就绪前返回 None
|
||||
mock_render.return_value = "/tmp/video.mp4"
|
||||
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
||||
|
||||
result = resume_viral_video_pipeline.run("job-001")
|
||||
|
||||
assert result["ok"] is True
|
||||
assert result["video_url"] == "https://oss.example.com/final.mp4"
|
||||
assert job.status == ViralVideoStatus.COMPLETED
|
||||
assert job.credits_cost == CREDITS_VIRAL_VIDEO_COST
|
||||
@@ -0,0 +1,99 @@
|
||||
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
|
||||
|
||||
验证:
|
||||
- 3 张新表可在内存 SQLite 上创建
|
||||
- 默认值与基本 CRUD 正常
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
Base,
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
|
||||
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
|
||||
|
||||
class TestViralVideoJobModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(
|
||||
id="job-001",
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_default_values(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(id="job-002", user_id="user-002")
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoJobModel, "job-002")
|
||||
assert fetched.images == []
|
||||
assert fetched.fusion_level == "ai_polish"
|
||||
assert fetched.style_strength == "medium"
|
||||
assert fetched.status == "pending"
|
||||
assert fetched.credits_cost == 0
|
||||
assert fetched.retry_count == 0
|
||||
assert fetched.style_guide is None
|
||||
assert fetched.intent_result is None
|
||||
|
||||
|
||||
class TestViralVideoStyleTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
style_config={"cut_speed": "fast"},
|
||||
sort_order=1,
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
|
||||
assert fetched.name == "快节奏"
|
||||
assert fetched.style_config == {"cut_speed": "fast"}
|
||||
assert fetched.sort_order == 1
|
||||
|
||||
|
||||
class TestViralVideoPromptTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoPromptTemplateModel(
|
||||
id="pt-001",
|
||||
prompt_type="image_analysis",
|
||||
name="图片分析模板",
|
||||
content="请分析图片:{image_url}",
|
||||
variables=["image_url"],
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
|
||||
assert fetched.prompt_type == "image_analysis"
|
||||
assert fetched.content == "请分析图片:{image_url}"
|
||||
assert fetched.variables == ["image_url"]
|
||||
assert fetched.version == 1
|
||||
assert fetched.is_active is True
|
||||
@@ -0,0 +1,230 @@
|
||||
"""#2106 P0 修复单测:Seedance 对接、image_analysis 持久化、TTS Path 统一、BGM/MuseTalk 跳过。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# worker 容器 PYTHONPATH 包含 apps/worker(worker 侧代码使用顶层包名 services/、viral_video/)
|
||||
_WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
|
||||
if str(_WORKER_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_WORKER_ROOT))
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_job():
|
||||
return ViralVideoJob(
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=15,
|
||||
user_copy_text="测试文案",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
|
||||
|
||||
# ── P0-2: _step_video_analysis import 路径 ──────────────────────────
|
||||
|
||||
|
||||
class TestVideoAnalysisImport:
|
||||
def test_no_reference_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = ""
|
||||
assert _step_video_analysis(mock_job) is None
|
||||
|
||||
def test_with_reference_returns_dict_or_none(self, mock_job):
|
||||
"""有参考视频 URL 时,不管分析成功/失败/占位,返回 dict(不抛异常)。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||
|
||||
mock_job.reference_video_url = "https://example.com/ref.mp4"
|
||||
result = _step_video_analysis(mock_job)
|
||||
# 允许占位/失败/真实返回,但绝不能抛异常
|
||||
assert result is None or isinstance(result, dict)
|
||||
|
||||
|
||||
# ── P0-3: image_analysis 字段 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestImageAnalysisField:
|
||||
def test_default_none(self):
|
||||
job = ViralVideoJob(user_id="u1")
|
||||
assert job.image_analysis is None
|
||||
|
||||
def test_persist_and_read(self, mock_job):
|
||||
mock_job.image_analysis = {"products": [{"name": "口红"}]}
|
||||
assert mock_job.image_analysis["products"][0]["name"] == "口红"
|
||||
|
||||
|
||||
# ── P0-1: storyboard 规范化 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStoryboardNormalize:
|
||||
def test_normalize_fills_defaults(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
raw = [{"order": 0, "description": "镜头一"}]
|
||||
out = _normalize_storyboard(raw, total_duration=10, n_segments=1, copy_text="文案")
|
||||
assert len(out) == 1
|
||||
assert out[0]["duration"] >= 3
|
||||
assert out[0]["ken_burns"] in {"zoom_in", "zoom_out", "pan_left", "pan_right", "static"}
|
||||
assert out[0]["type"] == "product_shot"
|
||||
assert out[0]["text"] == ""
|
||||
|
||||
def test_normalize_scales_to_total_duration(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _normalize_storyboard
|
||||
|
||||
raw = [
|
||||
{"order": 0, "duration": 10, "description": "a"},
|
||||
{"order": 1, "duration": 10, "description": "b"},
|
||||
]
|
||||
out = _normalize_storyboard(raw, total_duration=10, n_segments=2, copy_text="x")
|
||||
total = sum(s["duration"] for s in out)
|
||||
assert total == 10
|
||||
|
||||
def test_fallback_storyboard(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _fallback_storyboard
|
||||
|
||||
out = _fallback_storyboard("文案", total_duration=15, n_segments=3)
|
||||
assert len(out) == 3
|
||||
assert sum(s["duration"] for s in out) == 15
|
||||
assert all(s["duration"] >= 3 for s in out)
|
||||
|
||||
def test_storyboard_llm_list(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
|
||||
|
||||
with patch("packages.shared.ai_service.call_llm") as mock_llm:
|
||||
mock_llm.return_value = [
|
||||
{"order": 0, "description": "产品特写", "duration": 5, "text": "t1"},
|
||||
{"order": 1, "description": "使用场景", "duration": 5, "text": "t2"},
|
||||
{"order": 2, "description": "CTA", "duration": 5, "text": "t3"},
|
||||
]
|
||||
result = _step_storyboard(mock_job, "文案", {"products": []})
|
||||
assert len(result) == 3
|
||||
assert all("description" in s for s in result)
|
||||
|
||||
|
||||
# ── P1: TTS 返回 Path|None ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTTSPath:
|
||||
def test_tts_returns_none_on_import_error(self, mock_job):
|
||||
"""get_tts_service 抛 ImportError 时 _step_tts 返回 None。"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
with patch("services.tts_service_factory.get_tts_service", side_effect=ImportError("no tts")):
|
||||
assert vv._step_tts(mock_job, "文案") is None
|
||||
|
||||
def test_tts_returns_none_when_path_not_exists(self, mock_job, tmp_path):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
fake_service = MagicMock()
|
||||
fake_service.synthesize.return_value = str(tmp_path / "not_exist.mp3")
|
||||
with patch("services.tts_service_factory.get_tts_service", return_value=fake_service):
|
||||
assert vv._step_tts(mock_job, "文案") is None
|
||||
|
||||
def test_tts_returns_path_when_exists(self, mock_job, tmp_path):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
audio = tmp_path / "voice.mp3"
|
||||
audio.write_bytes(b"ID3fake")
|
||||
fake_service = MagicMock()
|
||||
fake_service.synthesize.return_value = audio
|
||||
with patch("services.tts_service_factory.get_tts_service", return_value=fake_service):
|
||||
result = vv._step_tts(mock_job, "文案")
|
||||
assert isinstance(result, Path)
|
||||
assert result.exists()
|
||||
|
||||
|
||||
# ── P1: BGM 跳过 / MuseTalk 无 persona 跳过 ───────────────────────
|
||||
|
||||
|
||||
class TestBGMSkip:
|
||||
def test_bgm_returns_none(self, mock_job):
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
|
||||
|
||||
mock_job.bgm_preference = "upbeat"
|
||||
assert _step_bgm_select(mock_job) is None
|
||||
|
||||
|
||||
# ── P0-1: call_video_generation 参数构造 ──────────────────────────
|
||||
|
||||
|
||||
class TestCallVideoGeneration:
|
||||
def test_returns_none_when_client_unavailable(self):
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = False
|
||||
mock_get.return_value = mock_client
|
||||
assert call_video_generation("prompt") is None
|
||||
|
||||
def test_delegates_to_client(self, tmp_path):
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
out = tmp_path / "v.mp4"
|
||||
out.write_bytes(b"fake")
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.return_value = str(out)
|
||||
mock_get.return_value = mock_client
|
||||
result = call_video_generation(prompt="测试", image_url="https://img/x.jpg", duration=5, ratio="9:16")
|
||||
assert result == str(out)
|
||||
mock_client.video_generation.assert_called_once()
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
assert kwargs["prompt"] == "测试"
|
||||
assert kwargs["image_url"] == "https://img/x.jpg"
|
||||
assert kwargs["duration"] == 5
|
||||
|
||||
|
||||
# ── P0-1: _step_render 占位片段生成 ──────────────────────────────
|
||||
|
||||
|
||||
class TestPlaceholderClip:
|
||||
def test_make_placeholder_clip(self, tmp_path):
|
||||
import shutil
|
||||
|
||||
from apps.worker.worker_app.tasks.viral_video import _make_placeholder_clip, _probe_ok
|
||||
|
||||
if not shutil.which("ffmpeg"):
|
||||
pytest.skip("ffmpeg not available")
|
||||
|
||||
out = _make_placeholder_clip(tmp_path, 0, 3)
|
||||
assert out.exists()
|
||||
assert _probe_ok(str(out))
|
||||
|
||||
|
||||
# ── P0-1: DoubaoClient.video_generation 在不可用时返回 None ───────
|
||||
|
||||
|
||||
class TestDoubaoClientVideoGen:
|
||||
def test_unavailable_returns_none(self):
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = "" # is_available -> False
|
||||
assert client.video_generation("prompt") is None
|
||||
|
||||
|
||||
# ── P0-3: resume 从 job 读 image_analysis ────────────────────────
|
||||
|
||||
|
||||
class TestResumeReadsImageAnalysis:
|
||||
def test_resume_uses_persisted_image_analysis(self):
|
||||
"""resume_pipeline 应从 job.image_analysis 读(P0-3 持久化)。"""
|
||||
import inspect
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "job.image_analysis" in src
|
||||
@@ -0,0 +1,176 @@
|
||||
"""viral_video.py HTTP 端点单元测试(celery send_task 分支覆盖)。
|
||||
|
||||
直接调用路由函数(不启动 TestClient),通过 patch 注入 repo/session/user,
|
||||
覆盖 4 个 celery_app.send_task(...) 调用点:
|
||||
|
||||
- create_viral_video (generate) -> worker.run_viral_video_pipeline
|
||||
- retry_viral_video_job (retry) -> worker.run_viral_video_pipeline
|
||||
- confirm_intent -> worker.resume_viral_video_pipeline
|
||||
- analyze_style -> worker.run_video_style_analysis
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _auth_user(uid: str = "u1"):
|
||||
return SimpleNamespace(user=SimpleNamespace(id=uid))
|
||||
|
||||
|
||||
def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending", **kwargs):
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
job = MagicMock()
|
||||
job.id = job_id
|
||||
job.user_id = user_id
|
||||
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||
job.images = kwargs.pop("images", ["img-1"])
|
||||
job.industry = kwargs.pop("industry", "电商")
|
||||
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
||||
for k, v in {
|
||||
"persona_id": "",
|
||||
"viral_structure": "",
|
||||
"marketing_purpose": "",
|
||||
"bgm_preference": "",
|
||||
"duration": 30,
|
||||
"user_copy_text": "",
|
||||
"fusion_level": "ai_polish",
|
||||
"reference_audio_path": "",
|
||||
"reference_video_url": "",
|
||||
"style_strength": "medium",
|
||||
"style_template_id": "",
|
||||
"retry_count": 0,
|
||||
"error_msg": "",
|
||||
"result_video_url": "",
|
||||
"style_guide": None,
|
||||
"created_at": None,
|
||||
"started_at": None,
|
||||
"completed_at": None,
|
||||
"stage": "",
|
||||
"progress": 0.0,
|
||||
"intent_result": None,
|
||||
"updated_at": None,
|
||||
}.items():
|
||||
setattr(job, k, kwargs.pop(k, v))
|
||||
return job
|
||||
|
||||
|
||||
# ── generate ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateViralVideo:
|
||||
def _req(self, **kw):
|
||||
from app.schemas.viral_video import CreateViralVideoRequest
|
||||
|
||||
d = {"images": ["https://x.com/a.jpg"], "industry": "电商", "target_customer": "年轻人"}
|
||||
d.update(kw)
|
||||
return CreateViralVideoRequest(**d)
|
||||
|
||||
def test_generate_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
req = self._req()
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
saved_job = _make_job(job_id="job-new", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
|
||||
def fake_save(job):
|
||||
job.id = saved_job.id
|
||||
|
||||
repo.save.side_effect = fake_save
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.create_viral_video(req, authenticated_user=user, session=session)
|
||||
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=[saved_job.id])
|
||||
assert resp.id == saved_job.id
|
||||
|
||||
|
||||
# ── retry ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRetryViralVideo:
|
||||
def test_retry_dispatches_celery_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-retry", user_id="u1", status=ViralVideoStatus.FAILED, retry_count=1)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.retry_viral_video_job("job-retry", authenticated_user=user, session=session)
|
||||
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.retry_count == 2
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=["job-retry"])
|
||||
assert resp.id == "job-retry"
|
||||
|
||||
|
||||
# ── confirm-intent ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConfirmIntent:
|
||||
def test_confirm_intent_dispatches_resume_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmIntentRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-cfm", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmIntentRequest(confirmed_copy="确认后的文案")
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.confirm_intent("job-cfm", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.user_copy_text == "确认后的文案"
|
||||
job.resume_from_confirm.assert_called_once()
|
||||
mock_send.assert_called_once_with("worker.resume_viral_video_pipeline", args=["job-cfm"])
|
||||
assert resp.id == "job-cfm"
|
||||
|
||||
|
||||
# ── analyze-style ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAnalyzeStyle:
|
||||
def test_analyze_style_dispatches_analysis_task(self):
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import AnalyzeStyleRequest
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-sty", user_id="u1", status="pending")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = AnalyzeStyleRequest(reference_video_url="https://x.com/ref.mp4", style_template_id="tpl-1")
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.analyze_style("job-sty", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.reference_video_url == "https://x.com/ref.mp4"
|
||||
assert job.style_template_id == "tpl-1"
|
||||
mock_send.assert_called_once_with("worker.run_video_style_analysis", args=["job-sty"])
|
||||
assert resp.job_id == "job-sty"
|
||||
assert resp.status == "analyzing"
|
||||
@@ -0,0 +1,432 @@
|
||||
"""Unit tests for the viral_video WebSocket progress endpoint and worker event format.
|
||||
|
||||
These tests exercise:
|
||||
* the worker _emit_progress helper (JSON serialization + event_type kwarg)
|
||||
* the pure helper functions on the API route module
|
||||
* WebSocket authentication / ownership / 404 behaviour
|
||||
* Initial-snapshot / terminal-job fast-close behaviour of the WS endpoint
|
||||
|
||||
The CI unit-test environment sets ``USE_IN_MEMORY_DB=true`` and relies on
|
||||
``settings.effective_database_url`` returning a SQLite URL. ``app/db.py`` and
|
||||
``app/dependencies.py`` have been fixed to honour ``effective_database_url``
|
||||
(matching the worker), so these tests never need a real Postgres or Redis.
|
||||
|
||||
Imports go through the ``apps.worker.*`` namespace (not bare ``worker_app.*``)
|
||||
to stay consistent with the existing integration tests and avoid creating a
|
||||
second module object that would make cross-file patches invisible.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.websockets import WebSocketDisconnect
|
||||
|
||||
# Ensure CI-friendly env is set BEFORE any app import so SQLite is used.
|
||||
os.environ.setdefault("USE_IN_MEMORY_DB", "true")
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "test-secret")
|
||||
os.environ.setdefault("DATABASE_URL", "postgresql+psycopg://no:such@127.0.0.1:1/none")
|
||||
|
||||
import app.db as _app_db # noqa: E402
|
||||
from app.api.routes import viral_video as vv_module # noqa: E402
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as worker_vv # noqa: E402
|
||||
|
||||
|
||||
def _make_job(**kwargs):
|
||||
defaults = dict(
|
||||
id="job-1",
|
||||
user_id="user-1",
|
||||
status="running",
|
||||
current_stage="analyzing",
|
||||
progress_percent=30,
|
||||
status_message="looking good",
|
||||
error_msg=None,
|
||||
is_terminal=False,
|
||||
result_video_url=None,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Worker event serialization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWorkerEmitProgress:
|
||||
def test_emit_progress_serialises_with_json_dumps(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress("job-1", "analyzing", 12, message="hi")
|
||||
fake_r.publish.assert_called_once()
|
||||
channel, payload = fake_r.publish.call_args.args
|
||||
assert channel == "viral_video:job-1"
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["stage"] == "analyzing"
|
||||
assert parsed["type"] == "viral_video:progress"
|
||||
assert parsed["progress"] == 12
|
||||
assert parsed["job_id"] == "job-1"
|
||||
assert "'stage'" not in payload # JSON uses double quotes, not Python repr
|
||||
|
||||
def test_emit_progress_respects_event_type(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-2",
|
||||
"done",
|
||||
100,
|
||||
message="ok",
|
||||
event_type="viral_video:completed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:completed"
|
||||
assert parsed["progress"] == 100
|
||||
|
||||
def test_emit_progress_failure_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-3",
|
||||
"failed",
|
||||
0,
|
||||
message="err",
|
||||
data={"error": "oom"},
|
||||
event_type="viral_video:failed",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:failed"
|
||||
assert parsed["data"]["error"] == "oom"
|
||||
|
||||
def test_emit_progress_wait_user_event(self):
|
||||
fake_r = MagicMock()
|
||||
with patch("redis.from_url", return_value=fake_r):
|
||||
worker_vv._emit_progress(
|
||||
"job-4",
|
||||
"intent_parsing",
|
||||
35,
|
||||
message="waiting for you",
|
||||
event_type="viral_video:wait_user",
|
||||
)
|
||||
_, payload = fake_r.publish.call_args.args
|
||||
parsed = json.loads(payload)
|
||||
assert parsed["type"] == "viral_video:wait_user"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure helpers on the route module
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSHelpers:
|
||||
def test_estimate_progress_maps_status(self):
|
||||
assert vv_module._estimate_progress(_make_job(status="pending")) == 0.0
|
||||
assert vv_module._estimate_progress(_make_job(status="running")) == 5.0
|
||||
assert vv_module._estimate_progress(_make_job(status="wait_user_confirm")) == 35.0
|
||||
assert vv_module._estimate_progress(_make_job(status="completed")) == 100.0
|
||||
assert vv_module._estimate_progress(_make_job(status="failed")) == 0.0
|
||||
|
||||
def test_initial_message_readable(self):
|
||||
job = _make_job(status="running")
|
||||
msg = vv_module._initial_message(job)
|
||||
assert isinstance(msg, str) and msg
|
||||
job_failed = _make_job(status="failed", error_msg="boom")
|
||||
assert "boom" in vv_module._initial_message(job_failed)
|
||||
job_wait = _make_job(status="wait_user_confirm")
|
||||
assert "等待" in vv_module._initial_message(job_wait)
|
||||
|
||||
def test_stage_from_status_falls_back(self):
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="pending")), str)
|
||||
assert isinstance(vv_module._stage_from_status(_make_job(status="weird_unknown")), str)
|
||||
|
||||
def test_job_status_handles_enum_and_string(self):
|
||||
job = _make_job(status="running")
|
||||
assert vv_module._job_status(job) == "running"
|
||||
job_enum = _make_job(status=SimpleNamespace(value="completed"))
|
||||
assert vv_module._job_status(job_enum) == "completed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket authentication / ownership / 404
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSRejectsUnauthenticated:
|
||||
def test_no_token_closes_with_4401(self):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
with patch.object(vv_module, "_ws_authenticate_user", return_value=None):
|
||||
client = TestClient(app)
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-1"):
|
||||
pass
|
||||
assert exc.value.code == 4401
|
||||
|
||||
|
||||
def _build_client(*, auth_user, repo_get_return, redis_instance=None):
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = repo_get_return
|
||||
sess = MagicMock()
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=auth_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch("redis.from_url", return_value=redis_instance or MagicMock()),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
return TestClient(app), fake_repo, sess, patches
|
||||
|
||||
|
||||
class TestWSOwnershipAnd404:
|
||||
def test_other_users_job_closes_with_4403(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
other_job = _make_job(user_id="user-b")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=other_job)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-x?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4403
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
def test_missing_job_closes_with_4404(self):
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
client, _repo, _sess, patches = _build_client(auth_user=fake_user, repo_get_return=None)
|
||||
try:
|
||||
with pytest.raises(WebSocketDisconnect) as exc:
|
||||
with client.websocket_connect("/ws/job-missing?token=valid-token"):
|
||||
pass
|
||||
assert exc.value.code == 4404
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _ws_authenticate_user direct unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWSAuthenticateUser:
|
||||
def test_empty_token_returns_none(self):
|
||||
assert vv_module._ws_authenticate_user("") is None
|
||||
|
||||
def test_decode_exception_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", side_effect=Exception("bad token")):
|
||||
assert vv_module._ws_authenticate_user("not-a-jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_missing_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_non_string_sub_returns_none(self):
|
||||
sess_factory = MagicMock()
|
||||
with patch.object(_app_db, "SessionLocal", sess_factory):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": 123}):
|
||||
assert vv_module._ws_authenticate_user("jwt") is None
|
||||
sess_factory.assert_not_called()
|
||||
|
||||
def test_success_returns_user(self):
|
||||
sess = MagicMock()
|
||||
fake_user = SimpleNamespace(id="u1")
|
||||
fake_user_repo = MagicMock()
|
||||
fake_user_repo.find_by_id.return_value = fake_user
|
||||
with patch.object(_app_db, "SessionLocal", return_value=sess):
|
||||
with patch("app.auth._decode_user_token", return_value={"sub": "u1"}):
|
||||
with patch(
|
||||
"app.dependencies.get_user_repository",
|
||||
return_value=fake_user_repo,
|
||||
):
|
||||
result = vv_module._ws_authenticate_user("valid.jwt")
|
||||
assert result is fake_user
|
||||
fake_user_repo.find_by_id.assert_called_once_with("u1")
|
||||
sess.close.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket initial-snapshot / terminal-job fast-close tests.
|
||||
#
|
||||
# The Redis pubsub reader thread is factored into ``_run_pubsub_forwarder`` and
|
||||
# marked ``# pragma: no cover`` (integration-tested with a live Redis). These
|
||||
# tests patch it out so we can deterministically verify the pre-subscribe
|
||||
# handshake without needing a real Redis or real thread scheduling.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _run_ws_handshake(*, job):
|
||||
"""Drive a WS handshake; collect JSON messages before connection closes."""
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
|
||||
fake_user = SimpleNamespace(id=getattr(job, "user_id", "user-a"))
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = job
|
||||
sess = MagicMock()
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
return received, fake_repo, sess
|
||||
|
||||
|
||||
class TestWSInitialSnapshot:
|
||||
def test_running_job_sends_initial_snapshot(self):
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
received, repo, sess = _run_ws_handshake(job=job)
|
||||
assert received[0]["type"] == "viral_video:progress"
|
||||
assert received[0]["job_id"] == "job-1"
|
||||
assert received[0]["data"]["status"] == "running"
|
||||
# Session was used for both ownership check and initial snapshot.
|
||||
assert sess.close.call_count >= 2
|
||||
|
||||
def test_running_job_with_enum_status(self):
|
||||
job = _make_job(
|
||||
status=SimpleNamespace(value="wait_user_confirm"),
|
||||
user_id="user-a",
|
||||
is_terminal=False,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
assert received[0]["data"]["status"] == "wait_user_confirm"
|
||||
assert received[0]["progress"] == 35.0
|
||||
assert "等待" in received[0]["message"]
|
||||
|
||||
def test_already_completed_job_sends_completion_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url="https://example.com/v.mp4",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
types = [m["type"] for m in received]
|
||||
assert "viral_video:progress" in types
|
||||
assert "viral_video:completed" in types
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == "https://example.com/v.mp4"
|
||||
assert completed["progress"] == 100
|
||||
|
||||
def test_already_failed_job_sends_failed_event_and_closes(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg="out of memory",
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == "out of memory"
|
||||
assert failed["progress"] == 0
|
||||
|
||||
def test_completed_job_without_result_url_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="completed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
result_video_url=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
completed = next(m for m in received if m["type"] == "viral_video:completed")
|
||||
assert completed["data"]["video_url"] == ""
|
||||
|
||||
def test_failed_job_without_error_msg_sends_empty_string(self):
|
||||
job = _make_job(
|
||||
status="failed",
|
||||
user_id="user-a",
|
||||
is_terminal=True,
|
||||
error_msg=None,
|
||||
)
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
failed = next(m for m in received if m["type"] == "viral_video:failed")
|
||||
assert failed["data"]["error"] == ""
|
||||
|
||||
def test_initial_snapshot_exception_is_swallowed(self):
|
||||
"""If sending the initial snapshot raises, the endpoint should log and
|
||||
still proceed to the Redis forwarder (doesn't crash)."""
|
||||
job = _make_job(status="running", user_id="user-a", is_terminal=False)
|
||||
|
||||
async def _fake_forwarder(websocket, redis_lib, settings, job_id):
|
||||
await websocket.send_json({"type": "forwarder_reached"})
|
||||
await websocket.close()
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(vv_module.router)
|
||||
fake_user = SimpleNamespace(id="user-a")
|
||||
fake_repo = MagicMock()
|
||||
calls = {"n": 0}
|
||||
|
||||
def _get(job_id):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 2:
|
||||
raise RuntimeError("boom in snapshot")
|
||||
return job
|
||||
|
||||
fake_repo.get.side_effect = _get
|
||||
sess = MagicMock()
|
||||
patches = [
|
||||
patch.object(vv_module, "_ws_authenticate_user", return_value=fake_user),
|
||||
patch.object(vv_module, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo),
|
||||
patch.object(_app_db, "SessionLocal", return_value=sess),
|
||||
patch.object(vv_module, "_run_pubsub_forwarder", new=_fake_forwarder),
|
||||
]
|
||||
for p in patches:
|
||||
p.start()
|
||||
received = []
|
||||
try:
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/ws/job-1?token=valid") as ws:
|
||||
for _ in range(5):
|
||||
try:
|
||||
msg = ws.receive_json()
|
||||
received.append(msg)
|
||||
except Exception:
|
||||
break
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
assert any(m["type"] == "forwarder_reached" for m in received)
|
||||
@@ -65,10 +65,24 @@ def test_generate_video_preserves_original_function():
|
||||
assert original.__name__ == "generate_video", f"expected __name__='generate_video', got '{original.__name__}'"
|
||||
|
||||
|
||||
def test_sync_task_config_to_plan_is_plain_function():
|
||||
"""Helper must NOT be registered as a Celery task."""
|
||||
from worker_app.tasks.generation import _sync_task_config_to_plan
|
||||
def test_build_task_config_override_is_plain_function():
|
||||
"""#2098: 原 _sync_task_config_to_plan 已拆分为 _build_task_config_override + _download_voice_for_task,
|
||||
均为普通函数,不应被注册为 Celery task。"""
|
||||
from worker_app.tasks.generation import _build_task_config_override, _download_voice_for_task
|
||||
|
||||
assert not hasattr(
|
||||
_sync_task_config_to_plan, "run"
|
||||
), "_sync_task_config_to_plan must be a plain function, not a Celery task"
|
||||
for fn in (_build_task_config_override, _download_voice_for_task):
|
||||
assert not hasattr(fn, "run"), f"{fn.__name__} must be a plain function, not a Celery task"
|
||||
|
||||
# Bug A: override 对 title_config 做 key 归一化 (font_size→size, font_color→color)
|
||||
override = _build_task_config_override(
|
||||
{
|
||||
"title_config": {"font_size": 48, "font_color": "#ff0000", "text": "hi"},
|
||||
"bgm_config": {"url": "http://x/bgm.mp3"},
|
||||
"output_width": 1080,
|
||||
"output_height": 1920,
|
||||
}
|
||||
)
|
||||
assert override["title"]["size"] == 48
|
||||
assert override["title"]["color"] == "#ff0000"
|
||||
assert override["bgm"]["url"] == "http://x/bgm.mp3"
|
||||
assert override["export"]["resolution"] == "1080x1920"
|
||||
|
||||
Reference in New Issue
Block a user