Compare commits
151 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 6b2b30a6e9 | |||
| d0582d1600 | |||
| bd617128ce | |||
| 7800ff4c3d | |||
| 5d6a4675fb | |||
| 774845bf91 | |||
| 7a63905a1c | |||
| 4a4b8f4a05 | |||
| d8e995afff | |||
| 2ce31a4438 | |||
| 9a4206e65c | |||
| 70aeb5e642 | |||
| d707b64876 | |||
| fe4464df2e | |||
| 0c0fe4619d | |||
| ad686dcd8b | |||
| 347cc82ffa | |||
| d3ee11d27a | |||
| 16ef616907 | |||
| ea51b372af | |||
| 9c807444da | |||
| b31205e965 | |||
| b4b53c9e5d | |||
| cf0a08f503 | |||
| 3ede4dca1f | |||
| 133a6c5914 | |||
| 1e8ba91bba | |||
| 2a739dee17 | |||
| 6b75eb5f67 | |||
| dbb4e57810 | |||
| ecb049a57e | |||
| 58d6852a71 | |||
| 66033c520f | |||
| 07069144c7 | |||
| 5bbc34d4c3 | |||
| 83be8a7d35 | |||
| 6ee9ca6a33 | |||
| 9af967c4f0 | |||
| 75ca55b5e5 | |||
| a0d20bd55f | |||
| d2bd6cbc01 | |||
| 383367718c | |||
| 0ad5647d36 | |||
| e922b0b472 | |||
| 61142c0936 | |||
| 522668006f | |||
| f417611829 | |||
| e1ecea7a6e | |||
| 0200d499ef | |||
| eda3a3a540 | |||
| 649420bd35 | |||
| 069544da38 | |||
| 10007507a5 | |||
| 19024da223 | |||
| 033c4a2eab | |||
| 4ed906e5fa | |||
| d390d7c310 | |||
| 7fd9c0cf43 | |||
| 4b04c6401c | |||
| 0effc450a9 | |||
| b2fd6fe46b | |||
| e935d1d72a | |||
| 1f8c8d033e | |||
| a9cbe7d4c9 | |||
| 3a8ef857ac | |||
| fc6ebbecb6 | |||
| 3d8882c479 | |||
| a74be7c717 | |||
| 09b8b2990f | |||
| cdce1b2e10 | |||
| 5cefbc9c05 | |||
| 41fe2a96a6 | |||
| f0862934f8 | |||
| 774c4fc0df | |||
| c7f8db383f | |||
| 17a95eb8f0 | |||
| 3fbc1bbfe6 | |||
| fa9545f79b | |||
| 87eb480f3c | |||
| 8bdc39a1ab | |||
| 5e61dbe4f9 | |||
| 22e04d65a7 | |||
| 6ff57b2feb | |||
| 2981d20d5b | |||
| 6cddd72910 | |||
| 6d5c44d6be | |||
| 665a3063b6 | |||
| 24724dca9f | |||
| d08835ec9f | |||
| 77ce4a1a0d | |||
| 7ad722e6c6 | |||
| b54dda6526 | |||
| 69da326ed6 | |||
| f7f600d091 | |||
| bf9249da19 | |||
| ca834b23cb | |||
| 37f7aa3329 | |||
| 794f5f374b | |||
| 34305974ad | |||
| e83a7cad2e | |||
| c45a2ce9b1 | |||
| 9814fcdc22 | |||
| 6636dc45f7 | |||
| e11e4f0e99 | |||
| a7d6ba473b | |||
| 966da04c9c | |||
| fdeb792bab | |||
| ff1d878c62 | |||
| f19be5fd09 | |||
| eeb8a05b69 | |||
| efb7fa5729 | |||
| 5e1520230f | |||
| 96bcec5fdd | |||
| dfc5e5a5b6 | |||
| 0a004db1bd | |||
| 3be06c5763 | |||
| 02199d80ee | |||
| ca6803e1a5 | |||
| d449496f90 | |||
| a04e363d1b | |||
| da59a6c9a6 | |||
| 11c554e43a | |||
| f9f3e6bfb9 | |||
| 3904a8f3b5 | |||
| ba3e97c986 | |||
| 3106496c12 | |||
| f79f75b863 | |||
| b7a439d319 | |||
| 4d7c80ae07 | |||
| 26c0140d79 | |||
| c507f76c14 | |||
| 8a5cfe831e | |||
| d7d3f3184b | |||
| ec9240b52a | |||
| d0e5ef1753 | |||
| 6e8199581d | |||
| 1e23a3f094 | |||
| b04a803655 | |||
| e496f127a3 | |||
| 423be1446f | |||
| 6d9d2e8179 | |||
| b233529eee | |||
| 6de6971e7b | |||
| 548be6aced | |||
| 588a4b7320 | |||
| 1a93a9c00e | |||
| 01991f14d7 | |||
| 8abdeb9551 | |||
| 97ad0ae2e5 | |||
| 59c05148ab | |||
| a00031e100 |
+26
-2
@@ -79,14 +79,33 @@ CELERY_BROKER_URL=redis://localhost:6379/0
|
||||
CELERY_RESULT_BACKEND=redis://localhost:6379/1
|
||||
|
||||
|
||||
# ==================== Worker 配置 ====================
|
||||
# ==================== Worker 配置(#2073 队列分流) ====================
|
||||
#
|
||||
# 容器内跑三个独立进程:beat(只发定时任务)+ generation worker(实时高优)
|
||||
# + transcode worker(后台批量/清理)。三个进程的并发与开关独立配置。
|
||||
|
||||
# Worker 进程名称
|
||||
WORKER_NAME=xiaoxia-saas-worker
|
||||
|
||||
# Worker 并发数(同时执行的任务数)
|
||||
# 总并发参考(兼容旧变量):
|
||||
# - 若 GENERATION_CONCURRENCY 与 TRANSCODE_CONCURRENCY 都未显式设置,
|
||||
# entrypoint 会按此总数对半分配(gen=ceil(total/2), trans=剩余,各至少 1);
|
||||
# - 任一个 *_CONCURRENCY 显式设置后,按显式值生效,忽略此变量对应部分。
|
||||
WORKER_CONCURRENCY=4
|
||||
|
||||
# Generation worker 并发数(用户实时任务:视频生成/TTS/音色克隆/lipsync/数字人)
|
||||
# 实时链路对延迟敏感,建议 2C 以上机器设为 2;高负载场景可加到 4。
|
||||
GENERATION_CONCURRENCY=2
|
||||
|
||||
# Transcode worker 并发数(后台批量:素材入库转码/AI 分类打标/质量评分/查重/批量下载)
|
||||
# 后台任务可排队,独立伸缩;素材入库量大时可加到 4。
|
||||
TRANSCODE_CONCURRENCY=2
|
||||
|
||||
# 是否在本容器启动 celery beat 进程(默认 1)。
|
||||
# 默认 beat 与 worker 同容器部署;若要独立 beat 容器部署,worker 容器设为 0、
|
||||
# beat 容器单独跑 `celery -A worker_app.celery_app beat` 并设 BEAT_ENABLED=1。
|
||||
BEAT_ENABLED=1
|
||||
|
||||
# 每个子进程最多处理多少任务后重启(防止内存泄漏)
|
||||
WORKER_MAX_TASKS_PER_CHILD=1000
|
||||
|
||||
@@ -193,9 +212,14 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
|
||||
|
||||
DOUBAO_API_KEY=your-doubao-api-key
|
||||
DOUBAO_MODEL=doubao-seed-1-6-250615
|
||||
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
|
||||
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
|
||||
DOUBAO_TIMEOUT=30
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
# 视觉模型:pro 精度高,lite 速度快(viral-video 商品识别默认用 lite 提速)
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -1302,8 +1302,16 @@ jobs:
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/configs/douyin_cookies.txt"
|
||||
echo "✅ Douyin cookies uploaded"
|
||||
|
||||
# 上传 infra/docker 配置到服务器(compose 单一事实来源)
|
||||
echo "Uploading infra/docker configs to staging server..."
|
||||
ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" \
|
||||
"mkdir -p /var/lib/xiaoxia-saas-staging/infra/docker"
|
||||
scp -P "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no infra/docker/compose.yml \
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/infra/docker/compose.yml"
|
||||
echo "✅ infra/docker/compose.yml uploaded"
|
||||
|
||||
# 通过环境变量传递凭证,避免命令行引号转义问题
|
||||
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh"
|
||||
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} COMPOSE_SYNC=0 sh"
|
||||
|
||||
# 清理 CI runner 上的渲染文件
|
||||
rm -f .env.rendered
|
||||
|
||||
+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")
|
||||
@@ -0,0 +1,51 @@
|
||||
"""viral video add copy_result + voice/video columns
|
||||
|
||||
Revision ID: 088_viral_video_copy_result
|
||||
Revises: 087_viral_video_image_analysis
|
||||
Create Date: 2026-10-01
|
||||
|
||||
v1.6 爆款视频字段补齐:
|
||||
- copy_result JSON: 编导分镜脚本完整结构(overview/scene_and_lighting/shots/hard_constraints/negative_prompts/voiceover_script)
|
||||
- voice_id/voice_source: TTS 音色参数
|
||||
- video_ratio/video_model: Seedance 视频比例/模型
|
||||
注意:线上启动也有幂等 ADD COLUMN 补列逻辑 (_ensure_viral_video_columns),本 migration 提供标准 Alembic 路径,
|
||||
两套机制互不冲突(IF NOT EXISTS 等价行为)。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "088_viral_video_copy_result"
|
||||
down_revision = "087_viral_video_image_analysis"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 幂等添加列(通过单独执行 + 异常忽略兼容已由 backfill 补上的环境)
|
||||
cols = [
|
||||
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
|
||||
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
|
||||
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
|
||||
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
|
||||
("copy_result", "JSON"),
|
||||
]
|
||||
conn = op.get_bind()
|
||||
for name, ddl in cols:
|
||||
try:
|
||||
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
|
||||
except Exception:
|
||||
# 不支持 IF NOT EXISTS 的库(如老版本 SQLite)直接尝试 ADD COLUMN,失败则忽略
|
||||
try:
|
||||
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
for name in ("copy_result", "video_model", "video_ratio", "voice_source", "voice_id"):
|
||||
try:
|
||||
op.drop_column("viral_video_jobs", name)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,62 @@
|
||||
"""viral video add storyboard + generated_copy_text (complement 088)
|
||||
|
||||
Revision ID: 089_viral_video_cols
|
||||
Revises: 088_viral_video_copy_result
|
||||
Create Date: 2026-10-01
|
||||
|
||||
#2129 兜底迁移:补齐 _VIRAL_VIDEO_BACKFILL_COLS 中所有列,覆盖
|
||||
# watchtower 自动部署未跑历史 migration、且 AUTO_CREATE_SCHEMA=false 时
|
||||
# _ensure_viral_video_columns 未执行的场景。
|
||||
# 幂等 ADD COLUMN IF NOT EXISTS,已存在则跳过。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "089_viral_video_cols"
|
||||
down_revision = "088_viral_video_copy_result"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 扩展 alembic_version.version_num 字段长度(原来 VARCHAR(32) 装不下长 revision id)
|
||||
conn = op.get_bind()
|
||||
try:
|
||||
conn.execute(sa.text("ALTER TABLE alembic_version ALTER COLUMN version_num TYPE VARCHAR(256)"))
|
||||
except Exception:
|
||||
pass
|
||||
cols = [
|
||||
("storyboard", "JSON"),
|
||||
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
|
||||
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
|
||||
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
|
||||
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
|
||||
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
|
||||
("copy_result", "JSON"),
|
||||
]
|
||||
for name, ddl in cols:
|
||||
try:
|
||||
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN IF NOT EXISTS {name} {ddl}"))
|
||||
except Exception:
|
||||
try:
|
||||
conn.execute(sa.text(f"ALTER TABLE viral_video_jobs ADD COLUMN {name} {ddl}"))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
for name in (
|
||||
"copy_result",
|
||||
"video_model",
|
||||
"video_ratio",
|
||||
"voice_source",
|
||||
"voice_id",
|
||||
"generated_copy_text",
|
||||
"storyboard",
|
||||
):
|
||||
try:
|
||||
op.drop_column("viral_video_jobs", name)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,35 @@
|
||||
"""viral video add phase_message column (#2134)
|
||||
|
||||
Revision ID: 090_viral_video_phase_msg
|
||||
Revises: 089_viral_video_cols
|
||||
Create Date: 2026-10-02
|
||||
|
||||
#2134 阶段细粒度提示:viral_video 表新增 phase_message 列(中文阶段提示文案)。
|
||||
current_stage 列已在之前版本存在,本迁移只补 phase_message。
|
||||
幂等 ADD COLUMN IF NOT EXISTS。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "090_viral_video_phase_msg"
|
||||
down_revision = "089_viral_video_cols"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# SQLite/PostgreSQL 兼容的幂等添加列
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "phase_message" not in cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("phase_message", sa.String(length=500), nullable=False, server_default=""),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("viral_video_jobs", "phase_message")
|
||||
@@ -0,0 +1,49 @@
|
||||
"""viral video add current_stage column (#2137 follow-up)
|
||||
|
||||
Revision ID: 091_viral_video_stage
|
||||
Revises: 090_viral_video_phase_msg
|
||||
Create Date: 2026-10-02
|
||||
|
||||
#2137 follow-up fix: 090 migration missed current_stage column on viral_video_jobs,
|
||||
causing UndefinedColumn errors and 500s on all authenticated viral-video endpoints.
|
||||
Idempotently add current_stage and double-check phase_message.
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "091_viral_video_stage"
|
||||
down_revision = "090_viral_video_phase_msg"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "current_stage" not in cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column(
|
||||
"current_stage",
|
||||
sa.String(length=200),
|
||||
nullable=False,
|
||||
server_default="",
|
||||
),
|
||||
)
|
||||
if "phase_message" not in cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column(
|
||||
"phase_message",
|
||||
sa.String(length=500),
|
||||
nullable=False,
|
||||
server_default="",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("viral_video_jobs", "current_stage")
|
||||
@@ -0,0 +1,42 @@
|
||||
"""viral_video_jobs 增加 heartbeat_at 列(worker 心跳,用于僵尸任务超时回收)
|
||||
|
||||
Revision ID: 092_viral_video_heartbeat
|
||||
Revises: 091_viral_video_stage
|
||||
Create Date: 2026-10-02
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "092_viral_video_heartbeat"
|
||||
down_revision = "091_viral_video_stage"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "heartbeat_at" not in cols:
|
||||
op.add_column("viral_video_jobs", sa.Column("heartbeat_at", sa.DateTime(), nullable=True))
|
||||
op.execute(
|
||||
"UPDATE viral_video_jobs SET heartbeat_at = updated_at " "WHERE status = 'running' AND heartbeat_at IS NULL"
|
||||
)
|
||||
try:
|
||||
op.create_index("ix_viral_video_jobs_heartbeat_at", "viral_video_jobs", ["heartbeat_at"])
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "heartbeat_at" in cols:
|
||||
try:
|
||||
op.drop_index("ix_viral_video_jobs_heartbeat_at", table_name="viral_video_jobs")
|
||||
except Exception:
|
||||
pass
|
||||
op.drop_column("viral_video_jobs", "heartbeat_at")
|
||||
@@ -0,0 +1,87 @@
|
||||
"""viral_video 动态积分定价 + 积分字段从 Integer 改为 Float (#2151)
|
||||
|
||||
Revision ID: 093
|
||||
Revises: 092_viral_video_heartbeat
|
||||
Create Date: 2026-10-02
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "093"
|
||||
down_revision = "092_viral_video_heartbeat"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
|
||||
# 1) points_accounts 三列 Integer -> Float
|
||||
pa_cols = {c["name"]: c for c in inspector.get_columns("points_accounts")}
|
||||
for col in ("balance", "total_earned", "total_spent"):
|
||||
if col in pa_cols:
|
||||
op.alter_column(
|
||||
"points_accounts",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 2) points_transactions amount/balance_after Integer -> Float
|
||||
pt_cols = {c["name"]: c for c in inspector.get_columns("points_transactions")}
|
||||
for col in ("amount", "balance_after"):
|
||||
if col in pt_cols:
|
||||
op.alter_column(
|
||||
"points_transactions",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 3) users.points_balance Integer -> Float
|
||||
user_cols = {c["name"]: c for c in inspector.get_columns("users")}
|
||||
if "points_balance" in user_cols:
|
||||
op.alter_column(
|
||||
"users",
|
||||
"points_balance",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 4) viral_video_jobs.credits_cost Integer -> Float
|
||||
vv_cols = {c["name"]: c for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "credits_cost" in vv_cols:
|
||||
op.alter_column(
|
||||
"viral_video_jobs",
|
||||
"credits_cost",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 5) viral_video_jobs 新增列
|
||||
if "video_resolution" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("video_resolution", sa.String(20), nullable=False, server_default="720p"),
|
||||
)
|
||||
if "credits_prepaid" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_prepaid", sa.Float(), nullable=False, server_default="0"),
|
||||
)
|
||||
if "credits_transaction_id" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_transaction_id", sa.String(36), nullable=False, server_default=""),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -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=["爆款视频"])
|
||||
|
||||
@@ -29,8 +29,6 @@ from app.services.ai_avatar_render_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -44,7 +42,6 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService
|
||||
|
||||
|
||||
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
|
||||
@points_gate("ai_digital_human", per_unit=15)
|
||||
def create_render_job(
|
||||
body: CreateAiAvatarRenderRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -27,7 +27,6 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
)
|
||||
from packages.application import ListGeneratedVideosByTaskUseCase
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
|
||||
@@ -346,7 +345,6 @@ def _is_trusted_media_url(url: str) -> bool:
|
||||
|
||||
|
||||
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
||||
@points_gate("ai_cover")
|
||||
def generate_cover(
|
||||
body: GenerateCoverRequest,
|
||||
template_id: str = Query(..., description="模板 ID"),
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -43,7 +41,6 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -272,7 +269,6 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
|
||||
|
||||
|
||||
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
|
||||
@points_gate("ai_video", quantity_field="preview_count")
|
||||
def create_preview_generation_task(
|
||||
request: CreatePreviewGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -302,26 +298,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 +576,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 +587,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,7 +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 +278,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)
|
||||
@@ -376,7 +464,6 @@ def _resolve_project_and_library(
|
||||
|
||||
|
||||
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
|
||||
@points_gate("ai_video", quantity_field="count")
|
||||
def create_generation_task(
|
||||
request: CreateGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -594,21 +681,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 +1002,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 +1156,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 +1336,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 +1380,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,
|
||||
|
||||
@@ -9,6 +9,11 @@ from fastapi.responses import JSONResponse
|
||||
router = APIRouter(tags=["Health"])
|
||||
|
||||
|
||||
|
||||
def _pg_url(url: str) -> str:
|
||||
"""Convert SQLAlchemy URL (postgresql+psycopg://...) to libpq connection string."""
|
||||
return url.replace("postgresql+psycopg://", "postgresql://", 1).replace("postgresql+psycopg2://", "postgresql://", 1)
|
||||
|
||||
@router.get("/health", status_code=status.HTTP_200_OK)
|
||||
async def health_check():
|
||||
return {
|
||||
@@ -49,7 +54,7 @@ async def _check_database() -> dict:
|
||||
"message": "Using in-memory database",
|
||||
}
|
||||
try:
|
||||
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
|
||||
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT 1")
|
||||
cur.fetchone()
|
||||
@@ -124,7 +129,7 @@ async def _check_migrations() -> dict:
|
||||
"message": "Using in-memory database, no migrations needed",
|
||||
}
|
||||
try:
|
||||
conn = psycopg.connect(settings.DATABASE_URL, connect_timeout=3)
|
||||
conn = psycopg.connect(_pg_url(settings.DATABASE_URL), connect_timeout=3)
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("""
|
||||
SELECT COUNT(*) FROM information_schema.tables
|
||||
@@ -137,3 +142,5 @@ async def _check_migrations() -> dict:
|
||||
return {"status": "unhealthy", "message": f"Missing tables, found {count}/5"}
|
||||
except Exception as error:
|
||||
return {"status": "unhealthy", "message": f"Migration check failed: {error}"}
|
||||
|
||||
|
||||
|
||||
@@ -12,11 +12,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from datetime import UTC
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.dependencies import (
|
||||
get_db_session,
|
||||
get_voice_clone_profile_repository,
|
||||
@@ -32,9 +30,6 @@ from app.services.mediakit_client import MediaKitError
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -61,37 +56,6 @@ def create_lipsync_job(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟)
|
||||
if body.audio_url and body.audio_duration and body.audio_duration > 0:
|
||||
est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0))
|
||||
elif body.script_text:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text) / 240))
|
||||
else:
|
||||
est_minutes = 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""提交对口型任务.
|
||||
|
||||
三种模式:
|
||||
@@ -101,6 +65,8 @@ def create_lipsync_job(
|
||||
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings},
|
||||
后端同步ffprobe+写入timings+直接提交MediaKit(~2-3s)。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
job = svc.create_job(
|
||||
user_id=user_id,
|
||||
@@ -118,18 +84,8 @@ def create_lipsync_job(
|
||||
project_id=body.project_id,
|
||||
)
|
||||
except ValueError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}")
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 502
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -145,24 +101,11 @@ def create_lipsync_job(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"创建对口型任务失败: {exc}",
|
||||
) from exc
|
||||
|
||||
# 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等)
|
||||
# 若任务已创建且状态为 failed,退费
|
||||
if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed":
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return job
|
||||
|
||||
|
||||
@@ -176,37 +119,14 @@ def preview_tts(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""步骤1「生成配音」同步 TTS 预合成.
|
||||
|
||||
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
|
||||
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL(~24h 有效)。
|
||||
耗时约 2-3 秒。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
result = svc.preview_tts(
|
||||
user_id=user_id,
|
||||
@@ -218,11 +138,6 @@ def preview_tts(
|
||||
emotion=body.emotion,
|
||||
)
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 400
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -237,11 +152,6 @@ def preview_tts(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"TTS 合成失败: {exc}",
|
||||
|
||||
@@ -169,17 +169,7 @@ def check_points(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""消费前检查余额是否足够。未知 scene_key 返回 400(而非 500)。"""
|
||||
if body.scene_key not in POINTS_SCENES:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"code": "UNKNOWN_SCENE",
|
||||
"message": f"未知场景: {body.scene_key}",
|
||||
"valid_scenes": sorted(POINTS_SCENES.keys()),
|
||||
},
|
||||
)
|
||||
|
||||
"""消费前检查余额是否足够。已下线/未知场景返回 cost=0(免费)。"""
|
||||
# 积分系统暂停(ENABLE_CREDIT_SYSTEM=false):所有场景直接放行,需 0 积分
|
||||
if not _credits_enabled():
|
||||
svc = _get_service()
|
||||
@@ -195,13 +185,6 @@ def check_points(
|
||||
is_mem = _is_member(current_user)
|
||||
mt = _member_type(current_user)
|
||||
|
||||
# 混剪场景先检查免费额度
|
||||
is_free_quota = False
|
||||
if body.scene_key == "ai_video" and not is_mem:
|
||||
svc = _get_service()
|
||||
if svc.check_daily_free_clip(current_user.user.id, db):
|
||||
is_free_quota = True
|
||||
|
||||
required = calculate_points_cost(
|
||||
body.scene_key,
|
||||
is_mem,
|
||||
@@ -215,11 +198,11 @@ def check_points(
|
||||
balance = account["balance"]
|
||||
|
||||
return PointsCheckResponse(
|
||||
allowed=is_free_quota or balance >= required,
|
||||
allowed=balance >= required,
|
||||
required_points=required,
|
||||
current_balance=balance,
|
||||
remaining_after=balance - required,
|
||||
is_free_quota=is_free_quota,
|
||||
is_free_quota=False,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -44,7 +44,6 @@ from app.services.script_asr_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -373,7 +372,6 @@ def douyin_diag():
|
||||
|
||||
|
||||
@router.post("/extract-from-douyin", response_model=ExtractFromDouyinResponse)
|
||||
@points_gate("douyin_extract")
|
||||
def extract_from_douyin(
|
||||
request: ExtractFromDouyinRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -497,7 +495,6 @@ def extract_from_douyin(
|
||||
|
||||
|
||||
@router.post("/ai-rewrite", response_model=AiRewriteResponse)
|
||||
@points_gate("ai_rewrite")
|
||||
def ai_rewrite(
|
||||
request: AiRewriteRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -537,7 +534,6 @@ def ai_rewrite(
|
||||
|
||||
|
||||
@router.post("/ai-generate-titles", response_model=AiGenerateTitlesResponse)
|
||||
@points_gate("ai_title")
|
||||
def ai_generate_titles(
|
||||
request: AiGenerateTitlesRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -4,14 +4,12 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.core.celery_app import celery_app
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
@@ -53,8 +51,6 @@ from packages.application.tts_job.use_cases import (
|
||||
)
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.domain.voice_presets import list_voices
|
||||
from packages.ports.asset_library_repository import AssetLibraryRepository
|
||||
from packages.ports.asset_repository import AssetRepository
|
||||
@@ -144,31 +140,6 @@ def synthesize(
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
|
||||
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
|
||||
actual_voice_id = request.voice_id
|
||||
@@ -231,7 +202,6 @@ def synthesize(
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
)
|
||||
|
||||
synthesis_error: Exception | None = None
|
||||
try:
|
||||
job = workflow.start_synthesis(job.id)
|
||||
except Exception as e:
|
||||
@@ -239,18 +209,10 @@ def synthesize(
|
||||
# 但 DB 异常、网络异常等意外错误可能逃逸。
|
||||
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
|
||||
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
|
||||
synthesis_error = e
|
||||
try:
|
||||
job = workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception as inner_e:
|
||||
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 合成失败且已扣积分 → 退费
|
||||
if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
|
||||
if job.status.value == "processing":
|
||||
# 分段合成任务 vs 普通单段任务
|
||||
@@ -269,13 +231,6 @@ def synthesize(
|
||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 调度失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return TTSSynthesizeResponse(
|
||||
job_id=job.id,
|
||||
status=job.status,
|
||||
@@ -610,31 +565,6 @@ def preview_tts(
|
||||
用于前端预览配音效果,限制文本长度 200 字以内。
|
||||
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
|
||||
actual_voice_id = request.voice_id
|
||||
profile = voice_clone_repo.get(request.voice_id)
|
||||
@@ -664,12 +594,6 @@ def preview_tts(
|
||||
language=getattr(request, "language", "zh-CN"),
|
||||
)
|
||||
except (CosyVoiceError, ValueError) as e:
|
||||
# 合成失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
|
||||
if isinstance(e, CosyVoiceError):
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||
|
||||
@@ -191,6 +191,23 @@ def _find_duplicate_asset(
|
||||
return None
|
||||
|
||||
|
||||
|
||||
def _get_existing_asset_url(existing: Any, storage_service: Any) -> str:
|
||||
"""安全获取已存在素材的公网 URL,兼容 domain Asset(无 file_url 字段)和 ORM model。"""
|
||||
# Domain Asset 只有 storage_key 字段;ORM model 有 file_url 但存的也是 storage_key
|
||||
key = ""
|
||||
for attr in ("storage_key", "file_url"):
|
||||
v = getattr(existing, attr, None)
|
||||
if v:
|
||||
key = v
|
||||
break
|
||||
if not key:
|
||||
return ""
|
||||
try:
|
||||
return storage_service.get_url(key) or ""
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
def _create_pending_asset(
|
||||
asset_repository,
|
||||
project_id,
|
||||
@@ -290,6 +307,43 @@ def _submit_ingest_job(
|
||||
return job
|
||||
|
||||
|
||||
def _find_active_ingest_job(ingest_job_repository: Any, asset_id: str) -> Any | None:
|
||||
"""查询 asset 上是否存在"仍在跑或已成功"的 ingest job(FAILED 视为不存在,需重提)。"""
|
||||
if not asset_id:
|
||||
return None
|
||||
find = getattr(ingest_job_repository, "find_by_asset_id", None)
|
||||
if not callable(find):
|
||||
# 旧仓储未实现 find_by_asset_id,无法判断 → 保守返回 None(走正常流程,
|
||||
# _submit_ingest_job 自身有数据库唯一约束/幂等兜底,不会重复建 job)
|
||||
return None
|
||||
try:
|
||||
return find(asset_id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[upload] find_by_asset_id 查询失败,按无 job 处理: asset=%s", asset_id, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def _is_true_duplicate(existing_asset: Asset, ingest_job_repository: Any) -> tuple[bool, Any | None]:
|
||||
"""判断 `existing_asset` 是真重复(应短路返 duplicated)还是占位(应补提 ingest)。
|
||||
|
||||
返回 (is_duplicate, existing_job):
|
||||
- READY 素材:真重复,job 可能为 None(已就绪不需要 job_id)
|
||||
- PROCESSING/UPLOADING 且已有在跑/已完成 ingest job:幂等重试,真重复,job 返回给前端轮询
|
||||
- PROCESSING/UPLOADING 且无 job:prepare 建的占位 / 之前 ingest 创建失败 → 非重复,需补提 ingest
|
||||
- ERROR/DELETED:非重复(允许重新上传覆盖)
|
||||
"""
|
||||
status = getattr(existing_asset, "status", None)
|
||||
if status == AssetStatus.READY:
|
||||
return True, None
|
||||
if status in (AssetStatus.PROCESSING, AssetStatus.UPLOADING):
|
||||
job = _find_active_ingest_job(ingest_job_repository, existing_asset.id)
|
||||
if job is not None:
|
||||
return True, job
|
||||
return False, None
|
||||
# ERROR / DELETED / 其它:走正常流程重新 ingest
|
||||
return False, None
|
||||
|
||||
|
||||
@router.post("/direct/prepare", response_model=DirectUploadPrepareResponse)
|
||||
async def prepare_direct_upload(
|
||||
request: DirectUploadPrepareRequest,
|
||||
@@ -353,6 +407,7 @@ async def prepare_direct_upload(
|
||||
duplicated=True,
|
||||
skip_transfer=True,
|
||||
asset_id=existing.id,
|
||||
url=_get_existing_asset_url(existing, storage_service),
|
||||
)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
@@ -406,6 +461,7 @@ async def prepare_direct_upload(
|
||||
duplicated=False,
|
||||
skip_transfer=False,
|
||||
asset_id=pending_asset_id,
|
||||
url="",
|
||||
)
|
||||
|
||||
|
||||
@@ -444,12 +500,25 @@ async def complete_direct_upload(
|
||||
file_size=request.file_size,
|
||||
)
|
||||
if existing is not None:
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id="",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
url=storage_service.get_url(existing.storage_key),
|
||||
is_dup, existing_job = _is_true_duplicate(existing, ingest_job_repository)
|
||||
if is_dup:
|
||||
logger.info(
|
||||
"[upload] complete 幂等命中真重复: asset=%s status=%s job=%s",
|
||||
existing.id,
|
||||
getattr(existing, "status", None),
|
||||
getattr(existing_job, "id", None),
|
||||
)
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id=getattr(existing_job, "id", "") or "",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
url=storage_service.get_url(existing.storage_key),
|
||||
)
|
||||
logger.info(
|
||||
"[upload] complete 命中占位 asset(status=%s 无 ingest job),继续补提 ingest: asset=%s",
|
||||
getattr(existing, "status", None),
|
||||
existing.id,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -479,14 +548,24 @@ async def complete_direct_upload(
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
storage_key=normalized_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=request.file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
# 幂等保护:补提占位场景下可能已有 job(极端竞态),先查一次
|
||||
existing_job = _find_active_ingest_job(ingest_job_repository, pending_asset.id)
|
||||
if existing_job is not None:
|
||||
logger.info(
|
||||
"[upload] complete 补提时发现 job 已存在(竞态/并发重试),复用: asset=%s job=%s",
|
||||
pending_asset.id,
|
||||
existing_job.id,
|
||||
)
|
||||
job = existing_job
|
||||
else:
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
storage_key=normalized_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=request.file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=normalized_key,
|
||||
ingest_job_id=job.id,
|
||||
@@ -533,12 +612,27 @@ async def upload_asset(
|
||||
file_size=0,
|
||||
)
|
||||
if existing is not None:
|
||||
return UploadAssetResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id="",
|
||||
url="",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
is_dup, existing_job = _is_true_duplicate(existing, ingest_job_repository)
|
||||
if is_dup:
|
||||
logger.info(
|
||||
"[upload] multipart 幂等命中真重复: asset=%s status=%s job=%s",
|
||||
existing.id,
|
||||
getattr(existing, "status", None),
|
||||
getattr(existing_job, "id", None),
|
||||
)
|
||||
return UploadAssetResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id=getattr(existing_job, "id", "") or "",
|
||||
url=storage_service.get_url(existing.storage_key)
|
||||
if getattr(existing, "status", None) == AssetStatus.READY
|
||||
else "",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
)
|
||||
logger.info(
|
||||
"[upload] multipart 命中占位 asset(status=%s 无 ingest job),继续补提 ingest: asset=%s",
|
||||
getattr(existing, "status", None),
|
||||
existing.id,
|
||||
)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
@@ -574,14 +668,23 @@ async def upload_asset(
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
existing_job = _find_active_ingest_job(ingest_job_repository, pending_asset.id)
|
||||
if existing_job is not None:
|
||||
logger.info(
|
||||
"[upload] multipart 补提时发现 job 已存在(竞态/并发重试),复用: asset=%s job=%s",
|
||||
pending_asset.id,
|
||||
existing_job.id,
|
||||
)
|
||||
job = existing_job
|
||||
else:
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
|
||||
return UploadAssetResponse(
|
||||
storage_key=storage_key,
|
||||
|
||||
@@ -0,0 +1,957 @@
|
||||
"""爆款视频 API 路由。
|
||||
|
||||
v1.6 三步分步流水线端点(单次 Seedance 出片版):
|
||||
POST /api/v1/viral-video/analyze-images 阶段1:创建任务 + 仅做图片/视频分析,暂停在 image_analyzed
|
||||
POST /api/v1/viral-video/{job_id}/generate-copy 阶段2:用户填完参数后跑意图+文案+分镜+审核,暂停在 copy_generated
|
||||
POST /api/v1/viral-video/{job_id}/confirm-copy 阶段3:用户确认/编辑文案后跑渲染,直到完成
|
||||
|
||||
旧端点(兼容保留,旧前端/一键生成模式):
|
||||
POST /api/v1/viral-video/generate 一键入队,前半段跑到 wait_user_confirm
|
||||
POST /api/v1/viral-video/{job_id}/confirm-intent 旧的意图确认后继续渲染
|
||||
|
||||
通用:
|
||||
GET /api/v1/viral-video/{job_id} 查询任务状态(含 image_analysis/copy_result 编导脚本)
|
||||
GET /api/v1/viral-video/history 历史记录
|
||||
POST /api/v1/viral-video/{job_id}/retry 重试失败任务
|
||||
POST /api/v1/viral-video/{job_id}/analyze-style 触发风格分析
|
||||
GET /api/v1/viral-video/style-templates 风格模板列表
|
||||
WS /api/v1/viral-video/ws/{job_id}?token= WebSocket 进度推送
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.viral_video import (
|
||||
AnalyzeImagesRequest,
|
||||
AnalyzeStyleRequest,
|
||||
AnalyzeStyleResponse,
|
||||
ConfirmCopyRequest,
|
||||
ConfirmIntentRequest,
|
||||
CreateViralVideoRequest,
|
||||
CreditsFormulaBreakdown,
|
||||
EstimateCreditsRequest,
|
||||
EstimateCreditsResponse,
|
||||
GenerateCopyRequest,
|
||||
RetryViralVideoRequest,
|
||||
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 _build_copy_result(job) -> dict | None:
|
||||
"""v1.6: 返回编导分镜脚本 CopyResult 结构(给前端/Seedance 使用)。
|
||||
|
||||
- 若 job.copy_result 已持久化(v1.6 worker 生成),直接返回(补 final_copy 兜底)。
|
||||
- 否则从老字段(generated_copy_text=口播, storyboard=分镜列表, intent_result)拼装兼容结构。
|
||||
"""
|
||||
cr = getattr(job, "copy_result", None)
|
||||
if isinstance(cr, dict) and cr:
|
||||
out = dict(cr)
|
||||
# 向后兼容字段
|
||||
voiceover = out.get("voiceover_script", "") or ""
|
||||
out.setdefault("final_copy", voiceover)
|
||||
out.setdefault("suggested_copy", voiceover)
|
||||
out.setdefault("title", "")
|
||||
return out
|
||||
# 兼容 v1.5 老数据:storyboard 是老格式 [{order,type,description,text,duration,...}]
|
||||
copy_text = getattr(job, "generated_copy_text", "") or ""
|
||||
sb = getattr(job, "storyboard", None) or []
|
||||
intent = getattr(job, "intent_result", None) or {}
|
||||
if not copy_text and not sb:
|
||||
return None
|
||||
title = ""
|
||||
if isinstance(intent, dict):
|
||||
title = intent.get("suggested_title") or intent.get("intent", "") or ""
|
||||
shots = []
|
||||
for seg in sb:
|
||||
if isinstance(seg, dict):
|
||||
shots.append(
|
||||
{
|
||||
"time_range": "",
|
||||
"shot_type_angle_movement": seg.get("ken_burns", ""),
|
||||
"scene_and_dialogue": (seg.get("text") or "")
|
||||
+ (" " + seg.get("description", "") if seg.get("description") else ""),
|
||||
"action_details": "",
|
||||
"audio_bgm": "",
|
||||
"transition": seg.get("transition", "硬切"),
|
||||
"reference_image_index": None,
|
||||
}
|
||||
)
|
||||
ratio = getattr(job, "video_ratio", None) or "9:16"
|
||||
return {
|
||||
"overview": {"theme": title, "total_duration": getattr(job, "duration", 15), "aspect_ratio": ratio},
|
||||
"scene_and_lighting": "",
|
||||
"shots": shots,
|
||||
"hard_constraints": ["无字幕", "无水印", "人物一致性"],
|
||||
"negative_prompts": ["字幕", "水印", "错误文字", "五官崩坏"],
|
||||
"voiceover_script": copy_text,
|
||||
"final_copy": copy_text,
|
||||
"suggested_copy": copy_text,
|
||||
"title": title,
|
||||
}
|
||||
|
||||
|
||||
def _to_response(job) -> ViralVideoJobResponse:
|
||||
return ViralVideoJobResponse(
|
||||
id=job.id,
|
||||
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 or 15,
|
||||
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,
|
||||
current_stage=getattr(job, "current_stage", "") or "",
|
||||
phase_message=getattr(job, "phase_message", "") or "",
|
||||
image_analysis=getattr(job, "image_analysis", None),
|
||||
storyboard=getattr(job, "storyboard", None),
|
||||
generated_copy_text=getattr(job, "generated_copy_text", "") or "",
|
||||
copy_result=_build_copy_result(job),
|
||||
voice_id=getattr(job, "voice_id", "") or "",
|
||||
voice_source=getattr(job, "voice_source", "") or "",
|
||||
video_ratio=getattr(job, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(job, "video_model", "") or "",
|
||||
intent_result=job.intent_result,
|
||||
result_video_url=job.result_video_url,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
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 or 15,
|
||||
user_copy_text=request.user_copy_text,
|
||||
fusion_level=request.fusion_level,
|
||||
reference_audio_path=request.reference_audio_path,
|
||||
reference_video_url=request.reference_video_url,
|
||||
style_strength=request.style_strength,
|
||||
style_template_id=request.style_template_id,
|
||||
voice_id=getattr(request, "voice_id", "") or "",
|
||||
voice_source=getattr(request, "voice_source", "") or "",
|
||||
video_ratio=getattr(request, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(request, "video_model", "") or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
copy_result=None,
|
||||
)
|
||||
|
||||
# 持久化
|
||||
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.post("/analyze-images", response_model=ViralVideoJobResponse)
|
||||
def analyze_images(
|
||||
request: AnalyzeImagesRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""v1.5 阶段1:创建任务并仅做图片/视频 VLM 分析,跑完后状态=image_analyzed。
|
||||
|
||||
前端拿到 image_analysis(商品名/品牌/特征/颜色/材质等结构化结果)展示给用户;
|
||||
用户填完营销参数后再调 /{id}/generate-copy 进入阶段2。
|
||||
"""
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
repo = _get_job_repo(session)
|
||||
job = ViralVideoJob(
|
||||
user_id=authenticated_user.user.id,
|
||||
images=list(request.images),
|
||||
reference_video_url=request.reference_video_url or "",
|
||||
style_template_id=request.style_template_id or "",
|
||||
style_strength=request.style_strength or "medium",
|
||||
voice_id=request.voice_id or "",
|
||||
voice_source=request.voice_source or "",
|
||||
video_ratio=request.video_ratio or "9:16",
|
||||
video_model=request.video_model or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
duration=request.duration or 15,
|
||||
)
|
||||
repo.save(job)
|
||||
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_analyze", args=[job.id])
|
||||
logger.info("[爆款视频][阶段1] analyze-images 入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段1] analyze-images 入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"任务入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/generate-copy", response_model=ViralVideoJobResponse)
|
||||
def generate_copy(
|
||||
job_id: str,
|
||||
request: GenerateCopyRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""v1.6 阶段2:用户填完营销参数后,跑 意图解析 → 编导分镜脚本生成 → 合规审核。
|
||||
|
||||
跑完后状态=copy_generated,响应 copy_result(含 overview/scene_and_lighting/shots/
|
||||
hard_constraints/negative_prompts/voiceover_script),前端展示脚本与口播供用户编辑;
|
||||
确认/编辑后调 /{id}/confirm-copy 进入阶段3(TTS + 单次 Seedance 出片)。
|
||||
"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED):
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案")
|
||||
|
||||
# 允许失败任务重试:重置
|
||||
if job.status == ViralVideoStatus.FAILED:
|
||||
job.retry_count += 1
|
||||
job.error_msg = ""
|
||||
|
||||
# 把用户填的营销参数写到 job 上
|
||||
job.industry = request.industry or job.industry
|
||||
job.target_customer = request.target_customer or job.target_customer
|
||||
job.persona_id = request.persona_id or job.persona_id
|
||||
job.viral_structure = request.viral_structure or job.viral_structure
|
||||
job.marketing_purpose = request.marketing_purpose or job.marketing_purpose
|
||||
job.bgm_preference = request.bgm_preference or job.bgm_preference
|
||||
if request.duration:
|
||||
job.duration = max(5, min(30, int(request.duration)))
|
||||
job.user_copy_text = request.user_copy_text if request.user_copy_text else job.user_copy_text
|
||||
job.fusion_level = request.fusion_level or job.fusion_level
|
||||
job.reference_audio_path = request.reference_audio_path or job.reference_audio_path
|
||||
job.reference_video_url = request.reference_video_url or job.reference_video_url
|
||||
job.style_strength = request.style_strength or job.style_strength
|
||||
job.style_template_id = request.style_template_id or job.style_template_id
|
||||
if request.style_guide is not None:
|
||||
job.style_guide = request.style_guide
|
||||
job.voice_id = request.voice_id or job.voice_id
|
||||
job.voice_source = request.voice_source or job.voice_source
|
||||
job.video_ratio = request.video_ratio or job.video_ratio or "9:16"
|
||||
job.video_model = request.video_model or job.video_model or ""
|
||||
job.video_resolution = getattr(request, "video_resolution", "") or job.video_resolution or "720p"
|
||||
|
||||
job.resume_from_image_analyzed()
|
||||
repo.update(job)
|
||||
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_generate_copy", args=[job.id])
|
||||
logger.info("[爆款视频][阶段2] generate-copy 入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段2] generate-copy 入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"任务入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/confirm-copy", response_model=ViralVideoJobResponse)
|
||||
def confirm_copy(
|
||||
job_id: str,
|
||||
request: ConfirmCopyRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""v1.6 阶段3:用户确认/编辑口播后开始 TTS + 单次 Seedance 生成 + 上传。"""
|
||||
repo = _get_job_repo(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status != ViralVideoStatus.COPY_GENERATED:
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
|
||||
|
||||
# 积分预扣(已扣过/重试任务跳过)
|
||||
from app.config import settings as _settings
|
||||
|
||||
if _settings.points_enabled:
|
||||
already_paid = (float(getattr(job, "credits_prepaid", 0) or 0) > 0) or (
|
||||
float(getattr(job, "credits_cost", 0) or 0) > 0
|
||||
)
|
||||
if not already_paid:
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
w, h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
job.video_ratio or "9:16",
|
||||
)
|
||||
est_credits = calculate_viral_video_credits(
|
||||
int(job.duration or 15), w, h, job.video_model or "seedance-2.5"
|
||||
)
|
||||
svc = PointsService()
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, est_credits, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {est_credits} 积分,当前余额 {balance}",
|
||||
"required": est_credits,
|
||||
"balance": balance,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = est_credits
|
||||
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||
repo.update(job)
|
||||
|
||||
job.resume_from_copy_generated(edited_copy=request.edited_copy or None)
|
||||
repo.update(job)
|
||||
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_render", args=[job.id])
|
||||
logger.info("[爆款视频][阶段3] confirm-copy 入队: job_id=%s", job.id)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段3] confirm-copy 入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"任务入队失败: {e}")
|
||||
repo.update(job)
|
||||
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/estimate-credits", response_model=EstimateCreditsResponse)
|
||||
def estimate_credits(
|
||||
request: EstimateCreditsRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EstimateCreditsResponse:
|
||||
"""爆款视频积分预估(纯计算,不扣费、不创建任务)。
|
||||
|
||||
返回 estimated_credits 与 formula_breakdown(tokens / video_cost / fixed_cost /
|
||||
profit_multiplier / model_price / width / height / fps),便于前端展示计费明细。
|
||||
同时兼容前端传 model 或 video_model、resolution 或 video_resolution、ratio 或 video_ratio。
|
||||
"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
model = (request.model or "").strip() or "seedance-2.5"
|
||||
resolution = (request.resolution or "").strip() or "720p"
|
||||
ratio = (request.ratio or "").strip() or "9:16"
|
||||
duration = int(request.duration or 15)
|
||||
|
||||
w, h = resolve_video_dimensions(resolution, ratio)
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(
|
||||
duration, w, h, model,
|
||||
)
|
||||
breakdown = CreditsFormulaBreakdown(**bd)
|
||||
return EstimateCreditsResponse(estimated_credits=credits, formula_breakdown=breakdown)
|
||||
|
||||
|
||||
@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,
|
||||
request: RetryViralVideoRequest | None = None,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。
|
||||
|
||||
可选 body (RetryViralVideoRequest):若传入新的 duration/video_resolution/video_ratio/
|
||||
video_model,会重新预估积分并与原 credits_prepaid 做差额多退少补(不足抛 402 阻止重试);
|
||||
不传 body 或参数无变化时,保持原参数、原预扣金额不变,仅重置状态并入队。
|
||||
credits_prepaid 为 0 的老任务首次重试会走预扣流程(与 confirm-copy 一致)。
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
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="无权操作此任务")
|
||||
|
||||
# 判定是否为僵尸 running 任务:running 超过 10 分钟且心跳停止超过 2 分钟
|
||||
now = datetime.now(timezone.utc)
|
||||
is_stale_running = False
|
||||
if job.status == ViralVideoStatus.RUNNING and job.started_at is not None:
|
||||
hb = getattr(job, "heartbeat_at", None) or job.updated_at
|
||||
if (now - job.started_at).total_seconds() > 10 * 60 and hb is not None and (now - hb).total_seconds() > 2 * 60:
|
||||
is_stale_running = True
|
||||
|
||||
if job.status != ViralVideoStatus.FAILED and not is_stale_running:
|
||||
raise HTTPException(status_code=409, detail="只有失败或超时的任务可以重试")
|
||||
|
||||
# ── 参数变更检测 + 积分多退少补 ──────────────────────────────────────
|
||||
req = request or RetryViralVideoRequest()
|
||||
new_duration = req.duration
|
||||
new_resolution = (req.video_resolution or "").strip() or None
|
||||
new_ratio = (req.video_ratio or "").strip() or None
|
||||
new_model = (req.video_model or "").strip() or None
|
||||
|
||||
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||
old_resolution = (getattr(job, "video_resolution", "720p") or "720p").strip() or "720p"
|
||||
old_ratio = (getattr(job, "video_ratio", "9:16") or "9:16").strip() or "9:16"
|
||||
old_model = (getattr(job, "video_model", "") or "").strip()
|
||||
|
||||
# 仅当有任意字段传入且值不同才算"参数变更"
|
||||
param_changed = bool(
|
||||
(new_duration is not None and int(new_duration) != old_duration)
|
||||
or (new_resolution is not None and new_resolution != old_resolution)
|
||||
or (new_ratio is not None and new_ratio != old_ratio)
|
||||
or (new_model is not None and new_model != old_model)
|
||||
)
|
||||
|
||||
from app.config import settings as _settings
|
||||
|
||||
need_points_settle = False
|
||||
new_est = 0.0
|
||||
if _settings.points_enabled and param_changed:
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
eff_dur = int(new_duration if new_duration is not None else old_duration)
|
||||
eff_res = new_resolution if new_resolution is not None else old_resolution
|
||||
eff_ratio = new_ratio if new_ratio is not None else old_ratio
|
||||
eff_model = new_model if new_model is not None else (old_model or "seedance-2.5")
|
||||
w, h = resolve_video_dimensions(eff_res, eff_ratio)
|
||||
new_est, _ = calculate_viral_video_credits_with_breakdown(eff_dur, w, h, eff_model or "seedance-2.5")
|
||||
need_points_settle = True
|
||||
|
||||
# 写入新参数(即使不开 points 也要允许用户重试时改参数)
|
||||
if new_duration is not None:
|
||||
job.duration = max(5, min(30, int(new_duration)))
|
||||
if new_resolution is not None:
|
||||
job.video_resolution = new_resolution
|
||||
if new_ratio is not None:
|
||||
job.video_ratio = new_ratio
|
||||
if new_model is not None:
|
||||
job.video_model = new_model
|
||||
|
||||
if need_points_settle:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
old_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
svc = PointsService()
|
||||
diff = round(new_est - old_prepaid, 2)
|
||||
if abs(diff) >= 0.01:
|
||||
if diff > 0:
|
||||
# 新预扣更多:补扣差额
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, diff, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"重试参数变更后需补扣 {diff} 积分,余额不足(当前 {balance},需 {new_est})",
|
||||
"required": new_est,
|
||||
"balance": balance,
|
||||
"delta": diff,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid + diff, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 补扣差额 job_id=%s diff=%.2f new_prepaid=%.2f",
|
||||
job.id, diff, job.credits_prepaid,
|
||||
)
|
||||
else:
|
||||
# 新预扣更少:退还差额
|
||||
refund = round(-diff, 2)
|
||||
txn_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
svc.refund_points(
|
||||
user_id=authenticated_user.user.id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=session,
|
||||
ref_id=txn_id or job.id,
|
||||
description="爆款视频重试参数变更退费",
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid - refund, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 退还差额 job_id=%s refund=%.2f new_prepaid=%.2f",
|
||||
job.id, refund, job.credits_prepaid,
|
||||
)
|
||||
# 差额为 0 则不调整
|
||||
|
||||
# 重置状态
|
||||
job.retry_count += 1
|
||||
job.status = ViralVideoStatus.PENDING
|
||||
job.error_msg = "" if not is_stale_running else "任务执行超时,已重置重试"
|
||||
job.started_at = None
|
||||
job.completed_at = None
|
||||
job.current_stage = ""
|
||||
job.phase_message = ""
|
||||
job.heartbeat_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 stale=%s params_changed=%s",
|
||||
job.id, job.retry_count, is_stale_running, param_changed,
|
||||
)
|
||||
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": "",
|
||||
"image_analyzed": "image_analysis",
|
||||
"copy_generated": "review",
|
||||
"wait_user_confirm": "intent_parsing",
|
||||
"completed": "uploading",
|
||||
"failed": "",
|
||||
"cancelled": "",
|
||||
}
|
||||
|
||||
_STATUS_PROGRESS = {
|
||||
"pending": 0.0,
|
||||
"running": 5.0,
|
||||
"image_analyzed": 15.0,
|
||||
"copy_generated": 70.0,
|
||||
"wait_user_confirm": 35.0,
|
||||
"completed": 100.0,
|
||||
"failed": 0.0,
|
||||
"cancelled": 0.0,
|
||||
}
|
||||
|
||||
_STATUS_MESSAGE = {
|
||||
"pending": "任务已创建,等待执行",
|
||||
"running": "任务执行中",
|
||||
"image_analyzed": "图片分析完成,等待填写营销参数",
|
||||
"copy_generated": "文案与分镜已生成,等待确认文案",
|
||||
"wait_user_confirm": "等待用户确认意图文案",
|
||||
"completed": "视频生成完成",
|
||||
"failed": "任务失败",
|
||||
"cancelled": "任务已取消",
|
||||
}
|
||||
|
||||
|
||||
def _stage_from_status(job) -> str:
|
||||
return _STATUS_STAGE.get(_job_status(job), "")
|
||||
|
||||
|
||||
def _estimate_progress(job) -> float:
|
||||
"""根据 status 粗略估算百分比(0-100),用于连接初始快照;
|
||||
连接建立后由 Redis 推送的真实事件持续更新。
|
||||
"""
|
||||
return _STATUS_PROGRESS.get(_job_status(job), 5.0)
|
||||
|
||||
|
||||
def _initial_message(job) -> str:
|
||||
"""给新连接的前端一个可读的初始状态文案。"""
|
||||
status_val = _job_status(job)
|
||||
if status_val == "failed" and job.error_msg:
|
||||
return f"任务失败: {job.error_msg}"
|
||||
return _STATUS_MESSAGE.get(status_val, "任务准备中")
|
||||
@@ -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]:
|
||||
|
||||
@@ -13,9 +13,9 @@ from pydantic import BaseModel, Field
|
||||
class PointsBalanceResponse(BaseModel):
|
||||
"""积分余额 + 会员状态"""
|
||||
|
||||
balance: int = Field(..., description="当前积分余额")
|
||||
total_earned: int = Field(..., description="累计获得积分")
|
||||
total_spent: int = Field(..., description="累计消耗积分")
|
||||
balance: float = Field(..., description="当前积分余额")
|
||||
total_earned: float = Field(..., description="累计获得积分")
|
||||
total_spent: float = Field(..., description="累计消耗积分")
|
||||
is_member: bool = Field(default=False, description="是否付费会员")
|
||||
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
|
||||
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
|
||||
@@ -30,8 +30,8 @@ class PointsTransactionItem(BaseModel):
|
||||
id: str
|
||||
type: str = Field(..., description="类型: add/deduct")
|
||||
source: str = Field(..., description="来源场景")
|
||||
amount: int
|
||||
balance_after: int
|
||||
amount: float
|
||||
balance_after: float
|
||||
description: str = ""
|
||||
ref_id: str = ""
|
||||
created_at: Optional[str] = None
|
||||
@@ -99,9 +99,9 @@ class PointsCheckResponse(BaseModel):
|
||||
"""消费前余额检查响应"""
|
||||
|
||||
allowed: bool
|
||||
required_points: int
|
||||
current_balance: int
|
||||
remaining_after: int
|
||||
required_points: float
|
||||
current_balance: float
|
||||
remaining_after: float
|
||||
is_free_quota: bool = False
|
||||
|
||||
|
||||
@@ -112,7 +112,7 @@ class PointsDeductRequest(BaseModel):
|
||||
"""积分扣减请求"""
|
||||
|
||||
scene_key: str
|
||||
amount: int
|
||||
amount: float
|
||||
description: Optional[str] = ""
|
||||
ref_id: Optional[str] = ""
|
||||
|
||||
@@ -170,7 +170,7 @@ class MembershipStatusResponse(BaseModel):
|
||||
is_member: bool
|
||||
member_type: Optional[str] = None
|
||||
member_expires_at: Optional[datetime] = None
|
||||
points_balance: int
|
||||
points_balance: float
|
||||
max_resolution: str = Field(
|
||||
default="1080p",
|
||||
description="可用最高分辨率: 720p(free) / 1080p(paid)",
|
||||
|
||||
@@ -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
+319
@@ -0,0 +1,319 @@
|
||||
"""爆款视频 API schemas (v1.6 单次 Seedance 出片版)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
# -- 枚举常量 --
|
||||
|
||||
VALID_FUSION_LEVELS = ("ai_full", "full_ai", "ai_polish", "user_primary")
|
||||
VALID_STYLE_STRENGTHS = ("light", "medium", "strict")
|
||||
VALID_STAGES = (
|
||||
"image_analysis",
|
||||
"video_analysis",
|
||||
"intent_parsing",
|
||||
"script_generation",
|
||||
"review",
|
||||
"tts",
|
||||
"rendering",
|
||||
"uploading",
|
||||
)
|
||||
VALID_VIDEO_RATIOS = ("9:16", "16:9", "1:1", "4:3", "3:4", "21:9")
|
||||
VALID_DURATIONS = (5, 10, 15, 20, 25, 30)
|
||||
VALID_VIDEO_RESOLUTIONS = ("480p", "720p", "1080p", "普清", "高清", "超清")
|
||||
|
||||
|
||||
# -- 编导脚本结构(v1.6) --
|
||||
|
||||
|
||||
class ShotScript(BaseModel):
|
||||
"""逐镜头分镜。"""
|
||||
|
||||
time_range: str = Field(default="", description="时间区间,如 0-3秒")
|
||||
shot_type_angle_movement: str = Field(default="", description="景别/角度/运镜,如『近景俯拍45度,缓慢推镜』")
|
||||
scene_and_dialogue: str = Field(default="", description="场景描述+口播台词")
|
||||
action_details: str = Field(default="", description="人物动作、表情、物品操作细节")
|
||||
audio_bgm: str = Field(default="", description="环境音+BGM提示")
|
||||
transition: str = Field(default="硬切", description="转场方式:硬切/淡入淡出/叠化")
|
||||
reference_image_index: int | None = Field(
|
||||
default=None, description="参考图片索引(0-based,对应上传的第几张产品图)"
|
||||
)
|
||||
|
||||
|
||||
class CopyResultOverview(BaseModel):
|
||||
theme: str = ""
|
||||
total_duration: int = 15
|
||||
aspect_ratio: str = "9:16"
|
||||
|
||||
|
||||
class CopyResult(BaseModel):
|
||||
"""v1.6 编导分镜脚本结构(给前端 + Seedance 用)。"""
|
||||
|
||||
overview: CopyResultOverview = Field(default_factory=CopyResultOverview)
|
||||
scene_and_lighting: str = ""
|
||||
shots: list[ShotScript] = Field(default_factory=list)
|
||||
hard_constraints: list[str] = Field(default_factory=list)
|
||||
negative_prompts: list[str] = Field(default_factory=list)
|
||||
voiceover_script: str = Field(
|
||||
default="", description="纯口播对白,从各镜 scene_and_dialogue 的对白部分拼接,供 TTS 使用"
|
||||
)
|
||||
# 向后兼容:final_copy = voiceover_script
|
||||
final_copy: str = ""
|
||||
suggested_copy: str = ""
|
||||
title: str = ""
|
||||
|
||||
|
||||
# -- Request Schemas --
|
||||
|
||||
|
||||
class CreateViralVideoRequest(BaseModel):
|
||||
"""旧接口:一键创建(保留兼容)。"""
|
||||
|
||||
images: list[str] = Field(..., min_length=1, max_length=20)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
duration: int = Field(default=15, ge=5, le=30, description="视频时长(秒),5-30")
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = "ai_polish"
|
||||
reference_audio_path: str = ""
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = "medium"
|
||||
style_template_id: str = ""
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
def _v_fl(cls, v: str) -> str:
|
||||
if v == "full_ai":
|
||||
return "ai_full"
|
||||
if v not in VALID_FUSION_LEVELS:
|
||||
raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}")
|
||||
return v
|
||||
|
||||
@field_validator("style_strength")
|
||||
@classmethod
|
||||
def _v_ss(cls, v: str) -> str:
|
||||
if v not in VALID_STYLE_STRENGTHS:
|
||||
raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}")
|
||||
return v
|
||||
|
||||
|
||||
class AnalyzeImagesRequest(BaseModel):
|
||||
"""v1.5+ 阶段1:创建任务 + 图片/视频分析。"""
|
||||
|
||||
images: list[str] = Field(..., min_length=1, max_length=30)
|
||||
reference_video_url: str = ""
|
||||
style_template_id: str = ""
|
||||
style_strength: str = "medium"
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
|
||||
class GenerateCopyRequest(BaseModel):
|
||||
"""v1.5+ 阶段2:填完营销参数,生成编导脚本。"""
|
||||
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
viral_structure: str = ""
|
||||
marketing_purpose: str = ""
|
||||
bgm_preference: str = ""
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = "ai_polish"
|
||||
reference_audio_path: str = ""
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = "medium"
|
||||
style_template_id: str = ""
|
||||
style_guide: dict | None = None
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
def _v_fl(cls, v: str) -> str:
|
||||
if v == "full_ai":
|
||||
return "ai_full"
|
||||
if v not in VALID_FUSION_LEVELS:
|
||||
raise ValueError(f"fusion_level must be one of {VALID_FUSION_LEVELS}")
|
||||
return v
|
||||
|
||||
@field_validator("style_strength")
|
||||
@classmethod
|
||||
def _v_ss(cls, v: str) -> str:
|
||||
if v not in VALID_STYLE_STRENGTHS:
|
||||
raise ValueError(f"style_strength must be one of {VALID_STYLE_STRENGTHS}")
|
||||
return v
|
||||
|
||||
|
||||
class ConfirmCopyRequest(BaseModel):
|
||||
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
|
||||
|
||||
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
|
||||
|
||||
|
||||
class ConfirmIntentRequest(BaseModel):
|
||||
"""旧 confirm-intent(兼容)。"""
|
||||
|
||||
confirmed_copy: str = ""
|
||||
adjustments: str = ""
|
||||
|
||||
|
||||
class AnalyzeStyleRequest(BaseModel):
|
||||
reference_video_url: str = Field(..., description="参考视频 URL")
|
||||
style_template_id: str = ""
|
||||
|
||||
|
||||
# -- Response Schemas --
|
||||
|
||||
|
||||
class ViralVideoJobResponse(BaseModel):
|
||||
"""爆款视频任务响应(v1.6 包含 copy_result 编导脚本结构)。"""
|
||||
|
||||
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 = 15
|
||||
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
|
||||
current_stage: str = (
|
||||
"" # 细粒度阶段 snake_case(analyzing_images/parsing_intent/generating_script/reviewing/tts_synthesizing/rendering_video/uploading)
|
||||
)
|
||||
phase_message: str = "" # 中文阶段提示文案(前端轮询/SSE 直接展示)
|
||||
image_analysis: dict | None = None
|
||||
# v1.6 编导脚本(推荐前端使用)
|
||||
copy_result: dict | None = None
|
||||
# v1.5 兼容字段
|
||||
storyboard: list | None = None
|
||||
generated_copy_text: str = ""
|
||||
# 音色/视频参数
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_cost: float = 0.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
|
||||
|
||||
|
||||
# -- 积分预估 --
|
||||
|
||||
|
||||
class EstimateCreditsRequest(BaseModel):
|
||||
"""爆款视频积分预估请求。
|
||||
|
||||
前端可传 model 或 video_model(兼容老字段);resolution/ratio/duration 为预估所需参数。
|
||||
"""
|
||||
|
||||
model: str = Field(default="", alias="video_model")
|
||||
resolution: str = Field(default="720p", alias="video_resolution")
|
||||
ratio: str = Field(default="9:16", alias="video_ratio")
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
model_config = {"populate_by_name": True}
|
||||
|
||||
|
||||
class CreditsFormulaBreakdown(BaseModel):
|
||||
"""爆款视频积分计费公式明细(前端展示用)。"""
|
||||
|
||||
tokens: float = Field(..., description="估算视频 tokens 数 (duration*width*height*fps/1024)")
|
||||
video_cost: float = Field(..., description="视频生成成本(元)= tokens/1e6 * model_price")
|
||||
fixed_cost: float = Field(..., description="固定成本(元),含 VLM/LLM/TTS/OSS/服务器")
|
||||
profit_multiplier: float = Field(..., description="利润系数(默认 1.3)")
|
||||
model_price: float = Field(..., description="模型单价(元/百万 tokens)")
|
||||
width: int = Field(..., description="视频宽度像素")
|
||||
height: int = Field(..., description="视频高度像素")
|
||||
fps: int = Field(..., description="视频帧率")
|
||||
|
||||
|
||||
class EstimateCreditsResponse(BaseModel):
|
||||
"""爆款视频积分预估响应。"""
|
||||
|
||||
estimated_credits: float
|
||||
formula_breakdown: CreditsFormulaBreakdown = Field(..., description="计费公式明细")
|
||||
|
||||
|
||||
class RetryViralVideoRequest(BaseModel):
|
||||
"""重试爆款视频任务的请求体(可选,允许改参数重新预估积分多退少补)。
|
||||
|
||||
不传 body 或字段全缺省:保持原参数、不重新扣点,走默认重置+入队逻辑。
|
||||
传入新的 duration/video_resolution/video_ratio/video_model:重新预估积分,
|
||||
与原 credits_prepaid 比较后多退少补(差额补扣不足抛 402)。
|
||||
"""
|
||||
|
||||
duration: int | None = Field(default=None, ge=5, le=30, description="重试时新的视频时长(秒)")
|
||||
video_resolution: str | None = Field(default=None, description="重试时新的分辨率,如 720p/1080p")
|
||||
video_ratio: str | None = Field(default=None, description="重试时新的画幅比,如 9:16/16:9")
|
||||
video_model: str | None = Field(default=None, description="重试时新的视频模型,如 seedance-2.5")
|
||||
|
||||
|
||||
# -- WebSocket 事件 Schema --
|
||||
|
||||
|
||||
class WSProgressEvent(BaseModel):
|
||||
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)
|
||||
@@ -11,14 +11,12 @@
|
||||
存储路径与元信息约定),返回 asset_id —— 下游仍以 voice_library_id(实为
|
||||
audio asset id)消费,渲染链路零改动。
|
||||
|
||||
积分扣点与 /tts 合成端点保持一致(ai_voice 场景),失败退费。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
@@ -32,13 +30,10 @@ from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.use_cases import CreateTTSJobUseCase
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.shared.storage import SharedStorageService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_POINTS_SCENE = "ai_voice"
|
||||
_SYNTH_TIMEOUT = 180.0 # 叙事配音在 HTTP 请求内同步等待,长文案分段合成时留出余量
|
||||
_CONTENT_TYPE_MAP = {"mp3": "audio/mpeg", "wav": "audio/wav", "pcm": "audio/pcm", "opus": "audio/opus"}
|
||||
|
||||
@@ -273,24 +268,6 @@ def prepare_narrative_voice(
|
||||
voice_clone_repository=voice_clone_repository,
|
||||
)
|
||||
|
||||
# 积分扣点(与 /tts 合成端点同口径),失败时在合成失败分支退费
|
||||
points_svc = PointsService() if points_enabled else None
|
||||
points_deducted = 0
|
||||
if points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(content) / 240))
|
||||
points_deducted = calculate_points_cost(
|
||||
_POINTS_SCENE,
|
||||
is_member=is_member,
|
||||
duration_minutes=est_minutes,
|
||||
member_type=member_type,
|
||||
)
|
||||
deduct_res = points_svc.deduct_points(user_id, points_deducted, _POINTS_SCENE, db)
|
||||
if not deduct_res["success"]:
|
||||
raise NarrativeError(
|
||||
f"积分不足,需要 {points_deducted} 积分,当前余额 {deduct_res['balance']}",
|
||||
status_code=402,
|
||||
)
|
||||
|
||||
use_case = CreateTTSJobUseCase(tts_repository)
|
||||
job = use_case.execute(
|
||||
user_id=user_id,
|
||||
@@ -311,19 +288,9 @@ def prepare_narrative_voice(
|
||||
workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("标记叙事 TTS job 失败出错: job_id=%s", job.id, exc_info=True)
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 失败退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError(f"配音合成失败:{e}", status_code=502) from e
|
||||
|
||||
if not job.is_completed:
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 未完成退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError("配音合成未完成,请稍后重试", status_code=504)
|
||||
|
||||
asset = _save_tts_job_as_voice_asset(
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -3,7 +3,8 @@
|
||||
*/
|
||||
|
||||
/** 任务状态 */
|
||||
export type TaskStatus = "pending" | "waiting" | "running" | "completed" | "failed" | "cancelled"
|
||||
export type TaskStatus =
|
||||
"pending" | "waiting" | "running" | "awaiting_cover" | "completed" | "failed" | "cancelled"
|
||||
|
||||
/** 任务类型 */
|
||||
export type TaskType = "ingest" | "generation" | string
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
import apiClient from "@/api/client"
|
||||
import type {
|
||||
GenerateViralVideoRequest,
|
||||
HistoryResponse,
|
||||
StyleTemplate,
|
||||
ViralVideoJob,
|
||||
ImageAnalysisResult,
|
||||
CopyResult,
|
||||
AnalyzeImagesRequest,
|
||||
GenerateCopyRequest,
|
||||
ConfirmCopyRequest,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
export function generateViralVideo(payload: GenerateViralVideoRequest) {
|
||||
return apiClient.post<ViralVideoJob>("/viral-video/generate", payload).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 查询单个任务 */
|
||||
export function getViralVideoJob(id: string) {
|
||||
return apiClient.get<ViralVideoJob>(`/viral-video/${id}`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 用户确认/修改 AI 理解的意图后继续 */
|
||||
export function confirmViralVideoIntent(
|
||||
id: string,
|
||||
payload: { confirmed_copy?: string; edits?: Record<string, unknown> },
|
||||
) {
|
||||
return apiClient
|
||||
.post<ViralVideoJob>(`/viral-video/${id}/confirm-intent`, payload)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 重试失败任务 */
|
||||
export function retryViralVideo(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/retry`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 历史记录(分页) */
|
||||
export function getViralVideoHistory(params?: { page?: number; page_size?: number }) {
|
||||
return apiClient.get<HistoryResponse>("/viral-video/history", { params }).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 预设风格模板 */
|
||||
export function getViralStyleTemplates() {
|
||||
return apiClient.get<StyleTemplate[]>("/viral-video/style-templates").then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 上传参考视频后触发风格分析 */
|
||||
export function analyzeViralStyle(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 动态预估积分消耗(STEP3 参数变化时调用) */
|
||||
export function estimateViralVideoCredits(params: {
|
||||
video_model: string
|
||||
resolution: string
|
||||
video_ratio: string
|
||||
duration: number
|
||||
}) {
|
||||
return apiClient
|
||||
.post<{ estimated_credits: number }>("/viral-video/estimate-credits", params)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
/** ── 三步拆分:前端 mock 辅助函数(后端新接口上线后可替换) ── */
|
||||
|
||||
/**
|
||||
* 客户端图片分析 mock(后端未提供 analyze-only 端点前的占位方案):
|
||||
* 基于已上传图片生成一份示例识别汇览,让 STEP1→STEP2 交互可走通。
|
||||
* 后端上线后改为调用真实接口。
|
||||
*/
|
||||
export function mockImageAnalysis(images: { name: string }[]): Promise<ImageAnalysisResult> {
|
||||
return new Promise((resolve) => {
|
||||
setTimeout(() => {
|
||||
const products = images.slice(0, 3).map((img, i) => {
|
||||
const n = img.name.replace(/\.[^.]+$/, "")
|
||||
return {
|
||||
name: n || `商品 ${i + 1}`,
|
||||
spec: i === 0 ? "500ml/瓶" : i === 1 ? "300g/盒" : undefined,
|
||||
brand: i === 0 ? "示例品牌" : undefined,
|
||||
features:
|
||||
i === 0
|
||||
? "瓶身透明、蓝色标签、白色瓶盖;标签上印有品牌Logo和产品名称;光线均匀,主体居中"
|
||||
: i === 1
|
||||
? "盒装包装、主色调为米白+暖黄;正面有产品实物图;文字清晰可辨"
|
||||
: "产品主体清晰、背景干净、色彩鲜艳,突出核心卖点",
|
||||
label_text: i === 0 ? "包装正面印有产品名称、净含量、品牌Logo" : undefined,
|
||||
image_index: i,
|
||||
}
|
||||
})
|
||||
resolve({ products })
|
||||
}, 1800)
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
* 客户端文案生成 mock(后端未提供 generate-copy 端点前的占位方案):
|
||||
* 后端上线后改为调用真实接口。
|
||||
*/
|
||||
export function mockGenerateCopy(params: {
|
||||
product: string
|
||||
sellingPoints?: string[]
|
||||
tone?: string
|
||||
duration?: number
|
||||
marketingPurpose?: string
|
||||
industry?: string
|
||||
targetCustomer?: string
|
||||
}): Promise<CopyResult> {
|
||||
return new Promise((resolve) => {
|
||||
setTimeout(() => {
|
||||
const product = params.product || "这款产品"
|
||||
const tone = params.tone || "亲切务实"
|
||||
const purpose = params.marketingPurpose || "品牌种草"
|
||||
resolve({
|
||||
title: `【${purpose}】${product},用过的人都说好!`,
|
||||
final_copy: `你有没有发现,选对一款${params.industry || "好物"}真的能让生活省心很多?\n\n今天给大家推荐这款${product}。${tone.includes("亲切") ? "说实话," : ""}我自己用了一段时间,最直观的感受就是——好用、省心、值得回购。\n\n✅ 亮点一:品质到位,用料扎实,细节处见用心\n✅ 亮点二:使用体验舒服,日常高频场景都能打\n✅ 亮点三:性价比很能打,这个价位真的没什么可挑的\n\n如果你也在找一款靠谱的${params.industry || "日常好物"},真的建议试试${product},不会让你失望。点击左下角,直接入手!`,
|
||||
suggested_copy: `你有没有发现,选对一款${params.industry || "好物"}真的能让生活省心很多?\n\n今天给大家推荐这款${product}。${tone.includes("亲切") ? "说实话," : ""}我自己用了一段时间,最直观的感受就是——好用、省心、值得回购。\n\n✅ 亮点一:品质到位,用料扎实,细节处见用心\n✅ 亮点二:使用体验舒服,日常高频场景都能打\n✅ 亮点三:性价比很能打,这个价位真的没什么可挑的\n\n如果你也在找一款靠谱的${params.industry || "日常好物"},真的建议试试${product},不会让你失望。点击左下角,直接入手!`,
|
||||
})
|
||||
}, 2200)
|
||||
})
|
||||
}
|
||||
|
||||
/** ── 三步拆分 v1.5 真实后端 API(PR #2117 合入后启用,前端可替换 mock 调用) ── */
|
||||
|
||||
/** 阶段1:上传图片后仅做 VLM 图片分析 + 可选参考视频风格分析,完成后状态=image_analyzed */
|
||||
export function analyzeViralImages(payload: AnalyzeImagesRequest) {
|
||||
return apiClient.post<ViralVideoJob>("/viral-video/analyze-images", payload).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 阶段2:用户填完营销参数后生成文案+分镜+合规审核,完成后状态=copy_generated,返回 copy_result */
|
||||
export function generateViralCopy(id: string, payload: GenerateCopyRequest) {
|
||||
return apiClient
|
||||
.post<ViralVideoJob>(`/viral-video/${id}/generate-copy`, payload)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 阶段3:用户确认/编辑文案后开始 TTS→渲染→上传,完成后状态=completed */
|
||||
export function confirmViralCopy(id: string, payload: ConfirmCopyRequest = {}) {
|
||||
return apiClient
|
||||
.post<ViralVideoJob>(`/viral-video/${id}/confirm-copy`, payload)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
export type FusionLevel = "ai_full" | "ai_polish" | "user_primary"
|
||||
export const FUSION_LEVELS: { value: FusionLevel; label: string; desc: string }[] = [
|
||||
{ value: "ai_full", label: "AI 全写", desc: "给我方向,全由AI创作" },
|
||||
{ value: "ai_polish", label: "AI润色", desc: "我写草稿,AI帮我润色" },
|
||||
{ value: "user_primary", label: "按我写的来", desc: "几乎不改我的文案" },
|
||||
]
|
||||
|
||||
export type StyleStrength = "light" | "medium" | "strict"
|
||||
export const STYLE_STRENGTHS: { value: StyleStrength; label: string }[] = [
|
||||
{ value: "light", label: "轻度借鉴" },
|
||||
{ value: "medium", label: "中度参考" },
|
||||
{ value: "strict", label: "像素级复刻" },
|
||||
]
|
||||
|
||||
/** v1.6 前端时长下拉选项(5/10/15/20/25/30秒) */
|
||||
export const VALID_DURATIONS = [5, 10, 15, 20, 25, 30] as const
|
||||
export type VideoDuration = (typeof VALID_DURATIONS)[number]
|
||||
|
||||
/** v1.6 支持的画幅比例 */
|
||||
export const VALID_RATIOS = ["9:16", "16:9", "1:1"] as const
|
||||
export type VideoRatio = (typeof VALID_RATIOS)[number]
|
||||
|
||||
export type ViralVideoStatus =
|
||||
| "pending"
|
||||
| "running"
|
||||
| "wait_user_confirm"
|
||||
| "image_analyzed"
|
||||
| "copy_generated"
|
||||
| "completed"
|
||||
| "failed"
|
||||
| "cancelled"
|
||||
|
||||
/**
|
||||
* v1.6 后端流水线阶段。单次 Seedance 出片版:
|
||||
* image_analysis → video_analysis(可选) → intent_parsing → script_generation → review → tts → rendering → uploading
|
||||
*/
|
||||
export type ViralVideoStage =
|
||||
| "image_analysis"
|
||||
| "video_analysis"
|
||||
| "intent_parsing"
|
||||
| "script_generation"
|
||||
| "review"
|
||||
| "tts"
|
||||
| "rendering"
|
||||
| "uploading"
|
||||
|
||||
/** 图片+视频分析阶段:属于「分析图片」按钮的范围 */
|
||||
const IMAGE_ANALYSIS_STAGES = new Set<ViralVideoStage>(["image_analysis", "video_analysis"])
|
||||
/** 编导脚本阶段:属于「生成文案」按钮的范围 */
|
||||
const COPY_STAGES = new Set<ViralVideoStage>(["intent_parsing", "script_generation", "review"])
|
||||
/** 视频生成阶段:属于「开始生成视频」按钮的范围(v1.6: TTS+单次Seedance+上传) */
|
||||
const VIDEO_STAGES = new Set<ViralVideoStage>(["tts", "rendering", "uploading"])
|
||||
|
||||
export function isImageAnalysisStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return !!stage && IMAGE_ANALYSIS_STAGES.has(stage)
|
||||
}
|
||||
export function isCopyStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return !!stage && COPY_STAGES.has(stage)
|
||||
}
|
||||
export function isVideoStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return !!stage && VIDEO_STAGES.has(stage)
|
||||
}
|
||||
/** 兼容旧调用:分析图片+生成文案 的所有前置阶段 */
|
||||
export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return isImageAnalysisStage(stage) || isCopyStage(stage)
|
||||
}
|
||||
|
||||
/** 单张图片 VLM 识别出的商品信息 */
|
||||
export interface ImageProductAnalysis {
|
||||
name?: string
|
||||
category?: string
|
||||
brand?: string
|
||||
colors?: string[]
|
||||
material_or_texture?: string
|
||||
key_features?: string[]
|
||||
visual_style?: string
|
||||
scene?: string
|
||||
target_audience_hint?: string
|
||||
text_on_image?: string
|
||||
/** 旧字段兼容 */
|
||||
spec?: string
|
||||
features?: string[] | string
|
||||
label_text?: string
|
||||
selling_points?: string
|
||||
image_index?: number
|
||||
}
|
||||
|
||||
export interface ImageAnalysisResult {
|
||||
products?: ImageProductAnalysis[]
|
||||
}
|
||||
|
||||
/** v1.6 编导分镜脚本 - 单镜头 */
|
||||
export interface ShotScript {
|
||||
/** 时间区间,如 "0-3秒" */
|
||||
time_range?: string
|
||||
/** 景别/角度/运镜,如 "近景俯拍45度,缓慢推镜" */
|
||||
shot_type_angle_movement?: string
|
||||
/** 场景描述+对白 */
|
||||
scene_and_dialogue?: string
|
||||
/** 人物动作/表情/物品操作细节 */
|
||||
action_details?: string
|
||||
/** 环境音+BGM提示 */
|
||||
audio_bgm?: string
|
||||
/** 转场方式(硬切/淡入淡出/叠化/结束) */
|
||||
transition?: string
|
||||
/** 参考图片索引(0-based,对应上传产品图数组) */
|
||||
reference_image_index?: number | null
|
||||
}
|
||||
|
||||
/** v1.6 编导分镜脚本 - 总览 */
|
||||
export interface CopyResultOverview {
|
||||
theme?: string
|
||||
total_duration?: number
|
||||
aspect_ratio?: string
|
||||
}
|
||||
|
||||
/** v1.6 编导分镜脚本(核心输出结构,给 Seedance 做 prompt,给 TTS 取 voiceover_script) */
|
||||
export interface CopyResult {
|
||||
overview?: CopyResultOverview
|
||||
/** 整体场景+光线描述 */
|
||||
scene_and_lighting?: string
|
||||
/** 逐镜头时间轴 */
|
||||
shots?: ShotScript[]
|
||||
/** 硬性约束(禁止字幕/水印/变形等) */
|
||||
hard_constraints?: string[]
|
||||
/** 负面提示词 */
|
||||
negative_prompts?: string[]
|
||||
/** 完整口播稿(纯文本,用于 TTS 合成) */
|
||||
voiceover_script?: string
|
||||
/** 向后兼容:= voiceover_script */
|
||||
final_copy?: string
|
||||
/** 向后兼容:= voiceover_script */
|
||||
suggested_copy?: string
|
||||
title?: string
|
||||
/** v1.5 旧字段兼容(老数据降级时可能出现) */
|
||||
scenes?: Array<{ shot: string; narration: string; duration?: number }>
|
||||
}
|
||||
|
||||
export interface StyleTemplate {
|
||||
id: string
|
||||
name: string
|
||||
description?: string
|
||||
thumbnail_url?: string
|
||||
style_config?: Record<string, unknown>
|
||||
tags?: string[]
|
||||
}
|
||||
|
||||
export interface IntentResult {
|
||||
intent?: string
|
||||
key_messages?: string[]
|
||||
tone?: string
|
||||
target_emotion?: string
|
||||
call_to_action?: string
|
||||
suggested_title?: string
|
||||
/** v1.5 旧字段兼容 */
|
||||
product?: string
|
||||
selling_points?: string[]
|
||||
target_audience?: string
|
||||
structure?: string
|
||||
duration?: number
|
||||
suggested_copy?: string
|
||||
}
|
||||
|
||||
export interface ViralVideoJob {
|
||||
id: string
|
||||
status: ViralVideoStatus
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
style_guide?: string | Record<string, unknown>
|
||||
user_copy_text?: string
|
||||
/** v1.6: = copy_result.voiceover_script(从 copy_result 派生,向后兼容) */
|
||||
final_copy_text?: string
|
||||
generated_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_mode?: "global" | "per_video"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
bgm_preference?: string
|
||||
intent_result?: IntentResult
|
||||
intent_text?: string
|
||||
/** v1.6 编导分镜脚本(核心产物) */
|
||||
copy_result?: CopyResult
|
||||
/** 向后兼容:= copy_result.shots */
|
||||
storyboard?: ShotScript[]
|
||||
image_analysis?: ImageAnalysisResult
|
||||
/** 视频比例:9:16 / 16:9 / 1:1,默认 9:16 */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空=后端默认) */
|
||||
video_model?: string
|
||||
/** 视频时长(秒,5-30,默认15) */
|
||||
duration?: number
|
||||
progress_stage?: ViralVideoStage
|
||||
progress_percent?: number
|
||||
progress_message?: string
|
||||
output_url?: string
|
||||
result_video_url?: string
|
||||
error_message?: string
|
||||
error_msg?: string
|
||||
credits_cost?: number
|
||||
created_at?: string
|
||||
updated_at?: string
|
||||
}
|
||||
|
||||
export interface GenerateViralVideoRequest {
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
douyin_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
bgm_preference?: string
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
language?: string
|
||||
persona_id?: string
|
||||
viral_structure?: string
|
||||
marketing_purpose?: string
|
||||
/** 视频时长(5-30秒,默认15) */
|
||||
duration?: number
|
||||
video_model?: string
|
||||
video_ratio?: string
|
||||
/** 三步拆分:step 控制后端执行到哪一步暂停 */
|
||||
step?: "analyze" | "generate_copy" | "generate_video"
|
||||
}
|
||||
|
||||
export interface HistoryResponse {
|
||||
items: ViralVideoJob[]
|
||||
total: number
|
||||
page: number
|
||||
page_size: number
|
||||
}
|
||||
|
||||
/** v1.6 阶段1请求:图片/视频分析(POST /viral-video/analyze-images) */
|
||||
export interface AnalyzeImagesRequest {
|
||||
images: string[]
|
||||
reference_video_url?: string
|
||||
style_template_id?: string
|
||||
style_strength?: StyleStrength
|
||||
/** TTS 音色 ID(STEP1 已选音色时传) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
/** Seedance 视频比例:9:16 | 16:9 | 1:1 */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
video_model?: string
|
||||
/** 视频时长(秒,5-30,默认15) */
|
||||
duration?: number
|
||||
}
|
||||
|
||||
/** v1.6 阶段2请求:填完营销参数后生成编导分镜脚本(POST /viral-video/{id}/generate-copy) */
|
||||
export interface GenerateCopyRequest {
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
persona_id?: string
|
||||
viral_structure?: string
|
||||
marketing_purpose?: string
|
||||
bgm_preference?: string
|
||||
/** 视频时长(秒,5-30,默认15) */
|
||||
duration?: number
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
reference_audio_path?: string
|
||||
reference_video_url?: string
|
||||
style_strength?: StyleStrength
|
||||
style_template_id?: string
|
||||
style_guide?: string | Record<string, unknown>
|
||||
/** TTS 音色 ID(优先级高于 persona_id) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
/** Seedance 视频比例(9:16/16:9/1:1 等) */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
|
||||
export interface ConfirmCopyRequest {
|
||||
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
|
||||
edited_copy?: string
|
||||
}
|
||||
|
||||
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
|
||||
export interface StoryboardSegment {
|
||||
order: number
|
||||
type: string
|
||||
description: string
|
||||
text: string
|
||||
duration: number
|
||||
ken_burns?: string
|
||||
transition?: string
|
||||
}
|
||||
@@ -2,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),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
||||
@@ -580,20 +580,6 @@ const AiAvatarPage: React.FC = () => {
|
||||
|
||||
return (
|
||||
<div className="aa-page">
|
||||
<div className="aa-page-header">
|
||||
<h1>AI数字人</h1>
|
||||
</div>
|
||||
|
||||
{/* 步骤切换导航条 */}
|
||||
<div className="aa-step-nav">
|
||||
<span className={`aa-step-nav__item${currentStep === 1 ? " active" : ""}`}>
|
||||
1. 视频 / 配音 / 文案
|
||||
</span>
|
||||
<span className={`aa-step-nav__item${currentStep === 2 ? " active" : ""}`}>
|
||||
2. 对口型 / 标题 / 封面 / 生成
|
||||
</span>
|
||||
</div>
|
||||
|
||||
<div className="aa-page-body">
|
||||
{/* ════ 步骤 1:出镜视频 / 配音库 / 文案 ════ */}
|
||||
{currentStep === 1 && (
|
||||
|
||||
@@ -13,8 +13,6 @@ import CloneModal from "@/components/voice/CloneModal"
|
||||
import VoiceSelectModal from "./components/VoiceSelectModal"
|
||||
import ScriptSelectModal from "./components/ScriptSelectModal"
|
||||
import TtsVoiceModal from "./components/TtsVoiceModal"
|
||||
import GenerateHeader from "./components/GenerateHeader"
|
||||
import GenerateStepsBar from "./components/GenerateStepsBar"
|
||||
import GenerateStepContent from "./components/GenerateStepContent"
|
||||
import GenerateStepActions from "./components/GenerateStepActions"
|
||||
import { useGenerateFormState } from "./hooks/useGenerateFormState"
|
||||
@@ -88,7 +86,6 @@ const GeneratePage: React.FC = () => {
|
||||
style,
|
||||
autoSubtitles,
|
||||
bgm,
|
||||
editPlanId,
|
||||
sourceEditPlanId,
|
||||
previewTaskId,
|
||||
setPreviewTaskId,
|
||||
@@ -239,9 +236,14 @@ const GeneratePage: React.FC = () => {
|
||||
voiceModePerVideo,
|
||||
variantCoverUrls: previewCovers,
|
||||
selectedVariantIndexes: isBatch ? selectedVariantIds : undefined,
|
||||
onGenerationSuccess: () => {
|
||||
onGenerationSuccess: (status?: "completed" | "awaiting_cover") => {
|
||||
setPreviewTaskId(null)
|
||||
setStoredSourceEditPlanId(null)
|
||||
// #2088:渲染完成后自动跳到封面选择页(step 5),不再等用户手动点「下一步」
|
||||
// awaiting_cover 和 completed 都走封面页(completed 是旧 worker 或 finalize 后状态,仍支持选封面)
|
||||
if (status === "awaiting_cover" || status === "completed" || !status) {
|
||||
setCurrentStep(5)
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
@@ -518,10 +520,6 @@ const GeneratePage: React.FC = () => {
|
||||
|
||||
return (
|
||||
<div className="xx-generate-page">
|
||||
<GenerateHeader fromEditPlan={!!editPlanId} />
|
||||
|
||||
<GenerateStepsBar currentStep={currentStep} onStepClick={setCurrentStep} />
|
||||
|
||||
<div className={layoutClassName}>
|
||||
{/* ════ 步骤1~2 表单 / 步骤3 标题设置 / 步骤4 确认生成进度 / 步骤5 封面 ════ */}
|
||||
<div className="xx-generate-form">
|
||||
|
||||
@@ -11,7 +11,12 @@
|
||||
* 防止长标题在窄列里溢出导致与相邻卡片进度条视觉重叠。
|
||||
*/
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import {
|
||||
LoadingOutlined,
|
||||
CheckCircleFilled,
|
||||
CloseCircleOutlined,
|
||||
ClockCircleOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
@@ -61,6 +66,11 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#ef4444" }}
|
||||
/>
|
||||
) : task.status === "queued" ? (
|
||||
<ClockCircleOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
style={{ color: "#faad14" }}
|
||||
/>
|
||||
) : (
|
||||
<LoadingOutlined
|
||||
className="xx-batch-gen-card-icon"
|
||||
@@ -85,6 +95,21 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
<div className="xx-batch-gen-card-pct">{Math.round(task.progress)}%</div>
|
||||
</>
|
||||
)}
|
||||
{task.status === "queued" && (
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 8,
|
||||
color: "var(--text-secondary, #faad14)",
|
||||
fontSize: 13,
|
||||
padding: "8px 0",
|
||||
}}
|
||||
>
|
||||
<ClockCircleOutlined />
|
||||
<span>排队等待中,前面任务完成后自动开始渲染</span>
|
||||
</div>
|
||||
)}
|
||||
{(task.status === "completed" || task.status === "awaiting_cover") && video && (
|
||||
// 竖屏自适应容器(#1750):成片固定 1080×1920(9:16),
|
||||
// 视频按真实宽高比 contain 显示,黑底居中,杜绝横屏播放器左右大黑边
|
||||
|
||||
@@ -12,6 +12,7 @@ import Step2MaterialSelect from "../components/Step2MaterialSelect"
|
||||
import Step4TitleSettings from "../components/Step4TitleSettings"
|
||||
import Step6CoverSettings from "../components/Step6CoverSettings"
|
||||
import BatchGenerationGrid from "./BatchGenerationGrid"
|
||||
import Step3VoiceWithMode from "./Step3VoiceWithMode"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
import type { TitleTemplate } from "@/components/title/template-types"
|
||||
@@ -151,6 +152,12 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedCoverTemplate,
|
||||
onSelectedCoverTemplateChange,
|
||||
onConfirmGenerate,
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
voiceModePerVideo,
|
||||
onVoiceModePerVideoChange,
|
||||
voiceLibraryIds,
|
||||
onVoiceLibraryIdsChange,
|
||||
} = props
|
||||
|
||||
switch (currentStep) {
|
||||
@@ -184,32 +191,46 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
)
|
||||
case 3:
|
||||
return (
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1 ? props.selectedVariantIds?.length || 1 : 1
|
||||
}
|
||||
/>
|
||||
<>
|
||||
<Step4TitleSettings
|
||||
titleSettings={titleSettings}
|
||||
onTitleSettingsChange={onTitleSettingsChange}
|
||||
onUpdatePosition={onUpdatePosition}
|
||||
onUpdateFont={onUpdateFont}
|
||||
onUpdateSize={onUpdateSize}
|
||||
onToggleBold={onToggleBold}
|
||||
onToggleItalic={onToggleItalic}
|
||||
onToggleStroke={onToggleStroke}
|
||||
onToggleShadow={onToggleShadow}
|
||||
onApplyPreset={onApplyPreset}
|
||||
onUpdateStyle={onUpdateStyle}
|
||||
activePreset={activePreset}
|
||||
titlePresets={titlePresets}
|
||||
enableTemplates={enableTemplates}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={onApplyTemplate}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={onPreviewTitlesChange}
|
||||
onConfirmGenerate={onConfirmGenerate}
|
||||
generating={props.generating}
|
||||
selectedCount={
|
||||
props.previewCount && props.previewCount > 1
|
||||
? props.selectedVariantIds?.length || 1
|
||||
: 1
|
||||
}
|
||||
/>
|
||||
{/* 批量配音选择:共用/独立切换(#2096) */}
|
||||
<Step3VoiceWithMode
|
||||
previewCount={previewCount}
|
||||
selectedVoice={selectedVoice}
|
||||
onSelectedVoiceChange={onSelectedVoiceChange}
|
||||
voiceModePerVideo={voiceModePerVideo}
|
||||
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
|
||||
voiceLibraryIds={voiceLibraryIds}
|
||||
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
|
||||
/>
|
||||
</>
|
||||
)
|
||||
case 4:
|
||||
return (
|
||||
|
||||
@@ -157,7 +157,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
/>
|
||||
</div>
|
||||
) : (
|
||||
/* ── 批量:N 个独立标题输入框 ── */
|
||||
/* ── 批量:N 个独立标题输入框(两列布局 #2096) ── */
|
||||
<div className="xx-batch-titles">
|
||||
<div
|
||||
style={{
|
||||
@@ -170,17 +170,25 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
为每个视频输入独立标题。标题样式(字体/颜色/位置)全局统一。
|
||||
</div>
|
||||
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: 640 }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(2, minmax(0, 1fr))",
|
||||
gap: 16,
|
||||
}}
|
||||
>
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i} style={{ maxWidth: "100%" }}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<TitleLibraryAutoComplete
|
||||
placeholder={`输入或选择视频 ${i + 1} 的标题`}
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(val) => updateVariantTitle(i, val)}
|
||||
options={titleOptions}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import CoverSettingsModal from "./cover-settings/CoverSettingsModal"
|
||||
import CoverEditorModal from "./cover-settings/CoverEditorModal"
|
||||
import { useSharedCover } from "@/components/cover/useSharedCover"
|
||||
import { generateCover as apiGenerateCover } from "@/api/generation"
|
||||
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
|
||||
interface Step6CoverSettingsProps {
|
||||
coverSettings: CoverConfig
|
||||
@@ -176,15 +176,8 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
|
||||
thumbnail_url: previewUrl,
|
||||
mode: "upload",
|
||||
})
|
||||
// 查找图片素材库(复用批量封面的逻辑)
|
||||
const libs = await getAssetLibraries()
|
||||
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
|
||||
if (!imageLib) {
|
||||
hide()
|
||||
message.error("未找到素材库,请先创建图片素材库")
|
||||
return previewUrl
|
||||
}
|
||||
const result = await uploadAssetDirect({ file, library_id: imageLib.id })
|
||||
// 后端自动在默认项目下确保图片素材库存在(P0 404 修复)
|
||||
const result = await uploadAssetDirect({ file, kind: "image" })
|
||||
const realUrl = result?.url || ""
|
||||
if (!realUrl) {
|
||||
hide()
|
||||
|
||||
@@ -41,8 +41,8 @@ export interface UseGenerateVideoProps {
|
||||
enabled: boolean
|
||||
music_id?: string
|
||||
}
|
||||
/** 生成成功后的回调(用于清除持久化的 previewTaskId 等状态) */
|
||||
onGenerationSuccess?: () => void
|
||||
/** 生成成功后的回调(用于清除持久化的 previewTaskId 等状态);status=awaiting_cover 表示需进封面选择 */
|
||||
onGenerationSuccess?: (status?: "completed" | "awaiting_cover") => void
|
||||
/* ── 批量生成(#1677)── */
|
||||
/** 生成数量(1=单条旧逻辑,>1=批量) */
|
||||
previewCount?: number
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { useRef, useCallback, useState } from "react"
|
||||
import { useRef, useCallback, useState, useEffect } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { getGenerationTask, retryTask as retryGenerationTaskApi } from "@/api/tasks/tasks"
|
||||
@@ -10,7 +10,7 @@ export interface BatchTaskState {
|
||||
taskId: string
|
||||
/** 变体序号(0-based,与标题/封面数组对齐) */
|
||||
variantIndex: number
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed"
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed" | "queued"
|
||||
progress: number
|
||||
error: string | null
|
||||
/** 完成后的成片视频 */
|
||||
@@ -19,7 +19,7 @@ export interface BatchTaskState {
|
||||
|
||||
interface UseGenerationPollingOptions {
|
||||
onProgress: (progress: number) => void
|
||||
onComplete: (videos: unknown[]) => void
|
||||
onComplete: (videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => void
|
||||
onFailed: (errorMsg: string) => void
|
||||
/** 批量:单任务状态变化(第5步逐卡片展示) */
|
||||
onBatchTaskUpdate?: (taskId: string, patch: Partial<BatchTaskState>) => void
|
||||
@@ -31,13 +31,19 @@ const MAX_RETRYABLE_ERRORS = 10
|
||||
const MAX_RESULTS_RETRIES = 3
|
||||
|
||||
/**
|
||||
* 生成状态轮询 Hook(v4 — 批量任务独立状态 + 单任务重试)
|
||||
* 生成状态轮询 Hook(v5 — awaiting_cover 状态识别 + visibilitychange 恢复 + 状态透传)
|
||||
*
|
||||
* startPolling(taskId) 轮询单个任务;
|
||||
* startPollingBatch(tasks) 并行轮询 N 个任务:
|
||||
* - 每个任务独立进度/状态/失败,通过 onBatchTaskUpdate 实时回传
|
||||
* - 全部成功才 onComplete(聚合视频按变体顺序);任一失败不影响其他任务继续
|
||||
* - retryTask(taskId) 单独重试失败任务(重新轮询,后端任务仍在跑则直接接续)
|
||||
*
|
||||
* v5 修复(#2088):
|
||||
* 1. 单任务路径透传 taskStatus(completed / awaiting_cover)到 onComplete,外层据此区分跳转
|
||||
* 2. 监听 visibilitychange,页面从后台切回可见时立即补拉一次,解决切后台 setInterval 被浏览器
|
||||
* 降频/冻结导致进度卡在 56% 的问题
|
||||
* 3. 非 4xx/5xx 网络错误按 3s 退避重试(已有 MAX_RETRYABLE_ERRORS=10 兜底)
|
||||
*/
|
||||
export function useGenerationPolling({
|
||||
onProgress,
|
||||
@@ -49,12 +55,15 @@ export function useGenerationPolling({
|
||||
const cancelledRef = useRef(false)
|
||||
/** 批量任务上下文:taskId → 变体序号 */
|
||||
const batchContextRef = useRef<Map<string, number>>(new Map())
|
||||
/** 当前活跃的「立刻补拉一次」函数(visibilitychange 回调使用) */
|
||||
const immediateTickRef = useRef<(() => void) | null>(null)
|
||||
const [, forceTick] = useState(0)
|
||||
|
||||
const clearTimer = useCallback(() => {
|
||||
cancelledRef.current = true
|
||||
progressTimer.current.forEach((t) => clearTimeout(t))
|
||||
progressTimer.current = []
|
||||
immediateTickRef.current = null
|
||||
}, [])
|
||||
|
||||
/** 任务完成后拉取结果列表,带重试 */
|
||||
@@ -113,6 +122,7 @@ export function useGenerationPolling({
|
||||
|
||||
if (task.status === "completed" || task.status === "awaiting_cover") {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const videos = await fetchResultsWithRetry(taskId)
|
||||
if (cancelledRef.current) return
|
||||
if (videos === null) {
|
||||
@@ -128,6 +138,7 @@ export function useGenerationPolling({
|
||||
|
||||
if (task.status === "failed" || task.status === "cancelled") {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const rawMsg =
|
||||
task.error_info?.error_message ||
|
||||
task.error_message ||
|
||||
@@ -138,6 +149,7 @@ export function useGenerationPolling({
|
||||
return
|
||||
}
|
||||
|
||||
// running / pending / waiting:更新进度并安排下一次轮询
|
||||
const pct = Math.max(0, Math.min(99, Math.round(Number(task.progress) || 0)))
|
||||
callbacks?.onTaskProgress?.(pct)
|
||||
if (!callbacks && runId === 0) {
|
||||
@@ -149,16 +161,20 @@ export function useGenerationPolling({
|
||||
if (cancelledRef.current || done) return
|
||||
console.error("[轮询出错] taskId:", taskId, pollErr)
|
||||
const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined
|
||||
// 4xx 视为不可重试(任务不存在/权限问题等),直接失败
|
||||
if (status && status >= 400 && status < 500) {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const msg = extractErrorMessage(pollErr, status)
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
// 网络错误 / 5xx:3s 退避重试,最多 MAX_RETRYABLE_ERRORS 次
|
||||
consecutiveErrors += 1
|
||||
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const msg = "任务状态查询连续失败,请稍后在任务列表查看结果"
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
@@ -169,6 +185,18 @@ export function useGenerationPolling({
|
||||
}
|
||||
}
|
||||
|
||||
// 注册「立刻补拉一次」回调,供 visibilitychange 恢复时调用
|
||||
// 注意:必须在 done 后清理,避免切换页面时误触发已结束任务的补拉
|
||||
immediateTickRef.current = () => {
|
||||
if (!done && !cancelledRef.current) {
|
||||
// 清除未触发的 setTimeout,立即拉一次
|
||||
progressTimer.current.forEach((t) => clearTimeout(t))
|
||||
progressTimer.current = []
|
||||
consecutiveErrors = 0
|
||||
void poll()
|
||||
}
|
||||
}
|
||||
|
||||
const timer = setTimeout(poll, 1500)
|
||||
progressTimer.current.push(timer)
|
||||
})
|
||||
@@ -181,12 +209,22 @@ export function useGenerationPolling({
|
||||
(taskId: string) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.clear()
|
||||
pollSingleTask(taskId, 0)
|
||||
.then((videos) => {
|
||||
if (cancelledRef.current) return
|
||||
let resolvedStatus: "completed" | "awaiting_cover" = "completed"
|
||||
pollSingleTask(taskId, 0, {
|
||||
onTaskProgress: (pct) => onProgress(pct),
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
resolvedStatus = taskStatus ?? "completed"
|
||||
onProgress(100)
|
||||
onComplete(videos)
|
||||
message.success("视频生成完成!")
|
||||
onComplete(videos, resolvedStatus)
|
||||
},
|
||||
onTaskFailed: (msg) => onFailed(msg),
|
||||
})
|
||||
.then(() => {
|
||||
if (cancelledRef.current) return
|
||||
// awaiting_cover 是中间态(进封面选择页),不弹"完成"toast;completed 才弹
|
||||
if (resolvedStatus === "completed") {
|
||||
message.success("视频生成完成!")
|
||||
}
|
||||
})
|
||||
.catch((err: Error) => {
|
||||
if (cancelledRef.current) return
|
||||
@@ -201,7 +239,7 @@ export function useGenerationPolling({
|
||||
/**
|
||||
* 批量多任务轮询:
|
||||
* - 每个任务独立进度/状态回传 onBatchTaskUpdate
|
||||
* * 全部完成后按变体顺序聚合视频 onComplete
|
||||
* - 全部完成后按变体顺序聚合视频 onComplete
|
||||
* - 部分失败:整体不 onFailed(第5步逐卡片展示失败+重试按钮);全部失败才 onFailed
|
||||
*/
|
||||
const startPollingBatch = useCallback(
|
||||
@@ -211,6 +249,7 @@ export function useGenerationPolling({
|
||||
const progressMap = new Map<string, number>()
|
||||
const resultMap = new Map<string, unknown[]>()
|
||||
const failureMap = new Map<string, string>()
|
||||
const statusMap = new Map<string, "completed" | "awaiting_cover">()
|
||||
batchContextRef.current = new Map(tasks.map((t) => [t.taskId, t.variantIndex]))
|
||||
|
||||
const reportAggregateProgress = () => {
|
||||
@@ -225,7 +264,9 @@ export function useGenerationPolling({
|
||||
if (resultMap.size === tasks.length) {
|
||||
onProgress(100)
|
||||
const ordered = tasks.map((t) => resultMap.get(t.taskId) || []).flat()
|
||||
onComplete(ordered)
|
||||
// 批量:任一任务为 awaiting_cover,则整体透传 awaiting_cover(进封面页)
|
||||
const anyAwaiting = Array.from(statusMap.values()).some((s) => s === "awaiting_cover")
|
||||
onComplete(ordered, anyAwaiting ? "awaiting_cover" : "completed")
|
||||
message.success(`全部 ${tasks.length} 个视频生成完成!`)
|
||||
} else if (resultMap.size > 0) {
|
||||
// 部分失败:成功的视频聚合进成片列表(可进封面),失败卡片带重试按钮
|
||||
@@ -234,7 +275,8 @@ export function useGenerationPolling({
|
||||
.filter((t) => resultMap.has(t.taskId))
|
||||
.map((t) => resultMap.get(t.taskId) || [])
|
||||
.flat()
|
||||
onComplete(ordered)
|
||||
const anyAwaiting = Array.from(statusMap.values()).some((s) => s === "awaiting_cover")
|
||||
onComplete(ordered, anyAwaiting ? "awaiting_cover" : "completed")
|
||||
message.warning(
|
||||
`${failureMap.size} 个视频生成失败,可点击卡片上的「重试此视频」,成功的视频可先进入下一步`,
|
||||
)
|
||||
@@ -263,6 +305,7 @@ export function useGenerationPolling({
|
||||
progressMap.set(taskId, 100)
|
||||
resultMap.set(taskId, videos)
|
||||
const _finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
statusMap.set(taskId, _finalStatus)
|
||||
onBatchTaskUpdate?.(taskId, { status: _finalStatus, progress: 100, videos })
|
||||
reportAggregateProgress()
|
||||
checkAllSettled()
|
||||
@@ -309,5 +352,64 @@ export function useGenerationPolling({
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
/**
|
||||
* visibilitychange 恢复:页面从后台切回前台时,立刻触发一次补拉。
|
||||
* 解决浏览器后台标签页对 setTimeout 的 1Hz 节流/冻结导致的"进度卡 56%"问题。
|
||||
*/
|
||||
useEffect(() => {
|
||||
const handleVisibilityChange = () => {
|
||||
if (document.visibilityState === "visible" && immediateTickRef.current) {
|
||||
immediateTickRef.current()
|
||||
}
|
||||
}
|
||||
document.addEventListener("visibilitychange", handleVisibilityChange)
|
||||
// 页面聚焦也兜底一次(部分浏览器 visibilitychange 触发时机不一致)
|
||||
const handleFocus = () => {
|
||||
if (immediateTickRef.current) immediateTickRef.current()
|
||||
}
|
||||
window.addEventListener("focus", handleFocus)
|
||||
return () => {
|
||||
document.removeEventListener("visibilitychange", handleVisibilityChange)
|
||||
window.removeEventListener("focus", handleFocus)
|
||||
}
|
||||
}, [])
|
||||
|
||||
/**
|
||||
* 批量队列模式:逐任务追加到轮询队列(支持串行提交、429 排队重试场景)。
|
||||
* 与 startPollingBatch 不同的是:
|
||||
* - 不会 reset batchContextRef;多次调用会累积
|
||||
* - 不触发整体 onComplete / onFailed(完成判定交给外层 useEffect 按状态聚合)
|
||||
* - 仍通过 onBatchTaskUpdate 回传单任务状态
|
||||
*/
|
||||
const pollBatchTaskQueued = useCallback(
|
||||
(taskId: string, variantIndex: number) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.set(taskId, variantIndex)
|
||||
onBatchTaskUpdate?.(taskId, {
|
||||
taskId,
|
||||
variantIndex,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
})
|
||||
pollSingleTask(taskId, Date.now(), {
|
||||
onTaskProgress: (pct) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: pct })
|
||||
},
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
const finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
onBatchTaskUpdate?.(taskId, { status: finalStatus, progress: 100, videos })
|
||||
},
|
||||
onTaskFailed: (msg) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "failed", error: msg })
|
||||
},
|
||||
}).catch(() => {
|
||||
/* onTaskFailed 已处理 */
|
||||
})
|
||||
},
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, pollBatchTaskQueued, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
import { useCallback, useState } from "react"
|
||||
import { message } from "antd"
|
||||
import { generateCover } from "@/api/generation"
|
||||
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
/** onCoversChange 支持直接传值或函数式 updater(函数式用于串行回写避免闭包覆盖) */
|
||||
@@ -182,15 +182,9 @@ export function useBatchCovers({
|
||||
async (index: number, file: File) => {
|
||||
addUploading(index)
|
||||
try {
|
||||
const libs = await getAssetLibraries()
|
||||
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
|
||||
if (!imageLib) {
|
||||
message.error("未找到素材库,请先创建")
|
||||
return
|
||||
}
|
||||
const result = await uploadAssetDirect({
|
||||
file,
|
||||
library_id: imageLib.id,
|
||||
kind: "image",
|
||||
})
|
||||
const url = result?.url || ""
|
||||
if (url) {
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
* 视频生成 Hook
|
||||
* 封装视频生成的核心逻辑、状态管理、轮询等
|
||||
*/
|
||||
import { useState, useCallback, useEffect } from "react"
|
||||
import { useState, useCallback, useEffect, useRef } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor"
|
||||
import { createGenerationTask } from "@/api/tasks/tasks"
|
||||
import type { CreateGenerationTaskRequest } from "@/api/tasks/types"
|
||||
import type { UseGenerateVideoProps } from "./generate-video/types"
|
||||
import { getGenerationPhase } from "./generate-video/phase"
|
||||
import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling"
|
||||
@@ -13,6 +15,28 @@ import { validateGenerateInputs } from "./generate-video/buildPayload"
|
||||
import { calculateResolution } from "../utils/calculateResolution"
|
||||
import { extractBackendError, translateError } from "./generate-video/errorUtils"
|
||||
|
||||
export type GenerationCompleteStatus = "completed" | "awaiting_cover" | null
|
||||
|
||||
/** 判断是否是用户队列已满 429(需要排队重试而非直接报错) */
|
||||
function isUserQueueFullError(err: unknown): { waitMs: number } | null {
|
||||
if (!axios.isAxiosError(err)) return null
|
||||
if (err.response?.status !== 429 && err.response?.status !== 503) return null
|
||||
const detail = (err.response?.data as { detail?: unknown })?.detail
|
||||
const code =
|
||||
typeof detail === "object" && detail !== null ? (detail as { code?: string }).code : undefined
|
||||
if (code === "USER_QUEUE_FULL" || code === "SYSTEM_QUEUE_FULL") {
|
||||
const waitSec =
|
||||
typeof detail === "object" && detail !== null
|
||||
? Number((detail as { estimated_wait_seconds?: number }).estimated_wait_seconds) || 0
|
||||
: 0
|
||||
return { waitMs: Math.max(15_000, waitSec * 1000 || 30_000) }
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/** sleep */
|
||||
const sleep = (ms: number) => new Promise<void>((r) => setTimeout(r, ms))
|
||||
|
||||
export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { selectedTemplate, onGenerationSuccess } = props
|
||||
|
||||
@@ -22,11 +46,29 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const [generated, setGenerated] = useState(false)
|
||||
const [generateError, setGenerateError] = useState<string | null>(null)
|
||||
const [generatedVideos, setGeneratedVideos] = useState<GeneratedVideo[]>([])
|
||||
/** #2088:任务最终状态,区分 awaiting_cover(选封面)/ completed(已完成) */
|
||||
const [completionStatus, setCompletionStatus] = useState<GenerationCompleteStatus>(null)
|
||||
/** 单视频模式:当前任务 ID(封面 finalize 需要) */
|
||||
const [currentTaskId, setCurrentTaskId] = useState<string>("")
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
|
||||
|
||||
/** 排队中重试的定时器,unmount / 新提交时清理 */
|
||||
const queueTimersRef = useRef<number[]>([])
|
||||
const cancelledRef = useRef(false)
|
||||
|
||||
const clearQueueTimers = useCallback(() => {
|
||||
queueTimersRef.current.forEach((id) => clearTimeout(id))
|
||||
queueTimersRef.current = []
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
cancelledRef.current = true
|
||||
clearQueueTimers()
|
||||
}
|
||||
}, [clearQueueTimers])
|
||||
|
||||
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
|
||||
setBatchTasks((prev) => {
|
||||
const list = prev || []
|
||||
@@ -53,9 +95,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
(videos: unknown[]) => {
|
||||
setGenerating(false)
|
||||
(videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => {
|
||||
setGenerated(true)
|
||||
const finalStatus: GenerationCompleteStatus = taskStatus ?? "completed"
|
||||
setCompletionStatus(finalStatus)
|
||||
setGeneratedVideos(videos as GeneratedVideo[])
|
||||
// 批量:成功任务的 videos 已通过 onBatchTaskUpdate 写入,这里同步兜底
|
||||
setBatchTasks((prev) =>
|
||||
@@ -70,26 +113,35 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
: t,
|
||||
),
|
||||
)
|
||||
onGenerationSuccess?.()
|
||||
onGenerationSuccess?.(finalStatus)
|
||||
},
|
||||
[onGenerationSuccess],
|
||||
)
|
||||
const handleFailed = useCallback((errorMsg: string) => {
|
||||
setGenerating(false)
|
||||
setGenerateError(errorMsg)
|
||||
}, [])
|
||||
|
||||
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频。
|
||||
当全部任务都已结束(completed/awaiting_cover/failed)且无排队/渲染中任务时,关闭 generating。 */
|
||||
useEffect(() => {
|
||||
if (batchTasks.length === 0) return
|
||||
const byVariant = new Map<number, GeneratedVideo>()
|
||||
let hasQueued = false
|
||||
let hasRunning = false
|
||||
let hasSuccess = false
|
||||
let allDone = true
|
||||
batchTasks.forEach((t) => {
|
||||
if (
|
||||
t.status === "completed" ||
|
||||
(t.status === "awaiting_cover" && t.videos && t.videos.length > 0)
|
||||
) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
if (t.status === "queued") hasQueued = true
|
||||
else if (t.status === "running") hasRunning = true
|
||||
if (t.status === "completed" || t.status === "awaiting_cover") {
|
||||
hasSuccess = true
|
||||
if (t.videos && t.videos.length > 0) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
}
|
||||
}
|
||||
if (t.status !== "completed" && t.status !== "awaiting_cover" && t.status !== "failed") {
|
||||
allDone = false
|
||||
}
|
||||
})
|
||||
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
|
||||
@@ -99,17 +151,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
return ordered
|
||||
})
|
||||
if (allDone && !hasQueued && !hasRunning) {
|
||||
setGenerating(false)
|
||||
if (hasSuccess) {
|
||||
setGenerated(true)
|
||||
setCompletionStatus("awaiting_cover")
|
||||
}
|
||||
}
|
||||
}, [batchTasks])
|
||||
|
||||
const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({
|
||||
const { startPolling, pollBatchTaskQueued, retryTask, clearTimer } = useGenerationPolling({
|
||||
onProgress: handleProgress,
|
||||
onComplete: handleComplete,
|
||||
onFailed: handleFailed,
|
||||
onBatchTaskUpdate: handleBatchTaskUpdate,
|
||||
})
|
||||
|
||||
/** 根据 props 构造基础 payload(批量/单任务共用的字段) */
|
||||
const buildBasePayload = useCallback((): Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> => {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
const bgmConfig = {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
}
|
||||
|
||||
const titleConfig = props.titleSettings?.title
|
||||
? {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color: props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
}
|
||||
: undefined
|
||||
|
||||
const payload: Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> = {
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: bgmConfig as CreateGenerationTaskRequest["bgm_config"],
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(titleConfig ? ({ title_config: titleConfig } as Record<string, unknown>) : {}),
|
||||
}
|
||||
|
||||
return payload
|
||||
}, [props, selectedTemplate])
|
||||
|
||||
/* ── 生成视频 ──
|
||||
返回 true 表示任务创建成功并已开始轮询;false 表示校验未通过或创建失败 */
|
||||
返回 true 表示任务创建成功并已开始轮询(含排队中);false 表示校验未通过或创建失败 */
|
||||
const generate = useCallback(async (): Promise<boolean> => {
|
||||
const errorMsg = validateGenerateInputs(props)
|
||||
if (errorMsg) {
|
||||
@@ -117,34 +337,28 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
return false
|
||||
}
|
||||
|
||||
cancelledRef.current = false
|
||||
clearQueueTimers()
|
||||
setGenerating(true)
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
setGenerateError(null)
|
||||
setCompletionStatus(null)
|
||||
setBatchTasks([])
|
||||
setGeneratedVideos([])
|
||||
setCurrentTaskId("")
|
||||
clearTimer()
|
||||
|
||||
const basePayload = buildBasePayload()
|
||||
const assetIds = basePayload.asset_ids
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
|
||||
try {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
|
||||
// from-assets 已由 useStep2Materials 在用户选素材时(debounce 800ms)调用,
|
||||
// 后端已改为异步秒级返回,这里做一次轻量兜底:
|
||||
// 单次查 clips,已有则直接放行;没有则再调一次 from-assets。
|
||||
// from-assets 兜底:片段不存在则补一次
|
||||
if (assetIds.length > 0 && selectedTemplate) {
|
||||
try {
|
||||
const clipList = await getEditPlanClips(selectedTemplate, { limit: 500 })
|
||||
if (clipList.items.length === 0) {
|
||||
// 片段不存在(极端情况:useStep2Materials 的 debounce 还没触发)
|
||||
// 手动补一次 from-assets(后端秒级返回)
|
||||
await createClipsFromAssets(selectedTemplate, assetIds, "main")
|
||||
}
|
||||
} catch {
|
||||
@@ -152,221 +366,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
}
|
||||
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
const hide = message.loading(
|
||||
isBatch ? `正在生成 ${props.previewCount} 个视频...` : "正在生成预览视频...",
|
||||
0,
|
||||
)
|
||||
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
|
||||
const indexes =
|
||||
isBatch && props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = isBatch ? indexes.length : 1
|
||||
|
||||
// 标题文字数组:批量时按勾选顺序
|
||||
const titlesArr =
|
||||
isBatch && (props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: []
|
||||
// 配音数组:独立配音模式按勾选顺序;否则不传(回退共用 voice_library_id)
|
||||
const voiceArr =
|
||||
isBatch && props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map((i) => props.variantVoiceLibraryIds![i] || voiceLibraryId)
|
||||
: []
|
||||
// 封面数组:批量时按勾选顺序(未设置封面的变体传空串,后端回退智能封面)
|
||||
const coversArr =
|
||||
isBatch && props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: []
|
||||
// #1744 变体 plan 数组:预览阶段后端独立选片产出的 plan id,按勾选顺序回传,
|
||||
// 后端直接关联这些 plan 渲染(不再重新选片)→ 预览所见即成片。
|
||||
// 全部为空(降级本地模拟/后端端点未上线)时不传,后端走自身独立选片。
|
||||
const variantPlansArr =
|
||||
isBatch && props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: []
|
||||
const hasVariantPlans = variantPlansArr.some((id) => !!id)
|
||||
|
||||
try {
|
||||
const taskResp = await createGenerationTask({
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
},
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(isBatch ? { count: batchCount } : {}),
|
||||
...(titlesArr.length ? { titles: titlesArr } : {}),
|
||||
...(voiceArr.length ? { voice_library_ids: voiceArr } : {}),
|
||||
...(coversArr.length ? { cover_urls: coversArr } : {}),
|
||||
...(hasVariantPlans ? { variant_plan_ids: variantPlansArr } : {}),
|
||||
...(props.titleSettings?.title
|
||||
? {
|
||||
title_config: {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color:
|
||||
props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
})
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
if (taskIds.length > 1) {
|
||||
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
|
||||
setCurrentTaskId("")
|
||||
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
|
||||
} else {
|
||||
if (!isBatch) {
|
||||
/* ── 单视频:原逻辑(一次提交 count=1) ── */
|
||||
const hide = message.loading("正在生成预览视频...", 0)
|
||||
try {
|
||||
const taskResp = await createGenerationTask({ ...basePayload, count: 1 })
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
setCurrentTaskId(taskIds[0])
|
||||
startPolling(taskIds[0])
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
}
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
return true
|
||||
}
|
||||
|
||||
/* ── 批量:支持任意数量视频,按队列容量串行提交,429 自动排队重试 ── */
|
||||
const indexes = props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = indexes.length
|
||||
|
||||
const titlesAll =
|
||||
(props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: indexes.map(() => props.titleSettings?.title || "")
|
||||
const voiceArrAll =
|
||||
props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map(
|
||||
(i) => props.variantVoiceLibraryIds![i] || basePayload.voice_library_id || "",
|
||||
)
|
||||
: []
|
||||
const coversAll = props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: indexes.map(() => "")
|
||||
const plansAll = props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: indexes.map(() => "")
|
||||
|
||||
const hasAnyVoice = voiceArrAll.some((v) => !!v)
|
||||
const hasAnyCover = coversAll.some((u) => !!u)
|
||||
const hasAnyPlan = plansAll.some((id) => !!id)
|
||||
|
||||
// 先用占位 ID 把所有变体卡片置为 queued,UI 可见
|
||||
const placeholderIds = indexes.map((_, i) => `__queued_${Date.now()}_${i}`)
|
||||
const initialTasks: BatchTaskState[] = indexes.map((variantIndex, i) => ({
|
||||
taskId: placeholderIds[i],
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
}))
|
||||
setBatchTasks(initialTasks)
|
||||
|
||||
message.loading({
|
||||
content: `已提交 ${batchCount} 个视频任务,系统按队列容量依次渲染…`,
|
||||
key: "batch-gen",
|
||||
duration: 3,
|
||||
})
|
||||
|
||||
/** 将占位 taskId 更新为真实 taskId(卡片引用同一对象) */
|
||||
const replacePlaceholder = (placeholderId: string, realTaskId: string) => {
|
||||
setBatchTasks((prev) => {
|
||||
const idx = prev.findIndex((t) => t.taskId === placeholderId)
|
||||
if (idx === -1) return prev
|
||||
const next = [...prev]
|
||||
next[idx] = { ...next[idx], taskId: realTaskId }
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
/** 提交某一索引的单任务(count=1),成功后返回真实 taskId;429/503 则返回 waitMs */
|
||||
const submitOne = async (
|
||||
i: number,
|
||||
): Promise<{ queued: true; waitMs: number } | { queued: false; taskId: string }> => {
|
||||
const body: CreateGenerationTaskRequest = {
|
||||
...basePayload,
|
||||
count: 1,
|
||||
titles: [titlesAll[i] || ""],
|
||||
...(hasAnyVoice
|
||||
? { voice_library_ids: [voiceArrAll[i] || basePayload.voice_library_id || ""] }
|
||||
: {}),
|
||||
...(hasAnyCover ? { cover_urls: [coversAll[i] || ""] } : {}),
|
||||
...(hasAnyPlan && plansAll[i] ? { variant_plan_ids: [plansAll[i]] } : {}),
|
||||
}
|
||||
try {
|
||||
const resp = await createGenerationTask(body)
|
||||
const item = resp.items?.[0]
|
||||
const tid = item?.id
|
||||
if (!tid) throw new Error("创建任务成功但未返回任务 ID")
|
||||
return { queued: false, taskId: tid }
|
||||
} catch (err) {
|
||||
const q = isUserQueueFullError(err)
|
||||
if (q) return { queued: true, waitMs: q.waitMs }
|
||||
throw err
|
||||
}
|
||||
}
|
||||
|
||||
// 串行提交:每次提交一个;429/503 则等待后重试;其它错误立即标记该任务失败
|
||||
let fatalErr: unknown = null
|
||||
for (let i = 0; i < batchCount; i++) {
|
||||
if (cancelledRef.current) return false
|
||||
const variantIndex = indexes[i]
|
||||
const placeholderId = placeholderIds[i]
|
||||
let attempt = 0
|
||||
let submitted = false
|
||||
while (!submitted) {
|
||||
if (cancelledRef.current) return false
|
||||
attempt++
|
||||
try {
|
||||
const result = await submitOne(i)
|
||||
if (!result.queued) {
|
||||
replacePlaceholder(placeholderId, result.taskId)
|
||||
// 先更新到 running,再启动单任务增量轮询(不触发整体 onComplete)
|
||||
pollBatchTaskQueued(result.taskId, variantIndex)
|
||||
submitted = true
|
||||
} else {
|
||||
// 排队:保持 queued 状态,等待后重试
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
})
|
||||
if (attempt === 1) {
|
||||
message.info({
|
||||
content: `队列繁忙,${Math.round(result.waitMs / 1000)} 秒后自动继续提交后续视频…`,
|
||||
key: "batch-gen",
|
||||
duration: 4,
|
||||
})
|
||||
}
|
||||
await sleep(Math.min(result.waitMs, 60_000))
|
||||
}
|
||||
} catch (err) {
|
||||
// 非限流错误:该任务标记失败,继续后续任务(不阻断整个批量)
|
||||
console.error("[batch generate] 任务提交失败:", err)
|
||||
const msg = translateError(extractBackendError(err))
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "failed",
|
||||
error: msg,
|
||||
progress: 0,
|
||||
})
|
||||
submitted = true
|
||||
if (!fatalErr) fatalErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (fatalErr) {
|
||||
// 有任务失败但其余已成功,整体不 throw;由 UI 展示单个失败卡片
|
||||
}
|
||||
return true
|
||||
} catch (err: unknown) {
|
||||
console.error("[handleGenerate] 生成失败:", err)
|
||||
setGenerating(false)
|
||||
const backendMsg = extractBackendError(err)
|
||||
console.error("[handleGenerate] 错误信息:", backendMsg, "完整错误:", err)
|
||||
const finalMsg = translateError(backendMsg)
|
||||
setGenerateError(finalMsg)
|
||||
setGenerating(false)
|
||||
message.error(finalMsg)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
|
||||
}, [
|
||||
props,
|
||||
clearTimer,
|
||||
startPolling,
|
||||
selectedTemplate,
|
||||
buildBasePayload,
|
||||
handleBatchTaskUpdate,
|
||||
clearQueueTimers,
|
||||
pollBatchTaskQueued,
|
||||
])
|
||||
|
||||
const retry = useCallback(() => {
|
||||
setGenerateError(null)
|
||||
@@ -376,9 +554,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 第5步:单独重试某个失败任务 */
|
||||
const retryBatchTask = useCallback(
|
||||
(taskId: string) => {
|
||||
handleBatchTaskUpdate(taskId, { status: "running", progress: 0, error: null, videos: [] })
|
||||
retryTask(taskId)
|
||||
},
|
||||
[retryTask],
|
||||
[retryTask, handleBatchTaskUpdate],
|
||||
)
|
||||
|
||||
const dismissError = useCallback(() => {
|
||||
@@ -423,6 +602,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
generated,
|
||||
generateError,
|
||||
generatedVideos,
|
||||
completionStatus,
|
||||
currentTaskId,
|
||||
generate,
|
||||
retry,
|
||||
|
||||
@@ -45,6 +45,11 @@ export const STATUS_CONFIG: Record<
|
||||
color: "processing",
|
||||
icon: <SyncOutlined spin />,
|
||||
},
|
||||
awaiting_cover: {
|
||||
label: "待选封面",
|
||||
color: "warning",
|
||||
icon: <ClockCircleOutlined />,
|
||||
},
|
||||
completed: {
|
||||
label: "已完成",
|
||||
color: "success",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,264 @@
|
||||
/**
|
||||
* 爆款视频素材选择弹窗(通用版,支持 image/video/voice)
|
||||
* 基于 ai-avatar 的 ModalAssetPicker 改造:
|
||||
* - kind 可传 "image" | "video" | "voice"
|
||||
* - 多图场景 multiple=true 时底部"确认选择"
|
||||
* - 单选场景点击即回调关闭
|
||||
*/
|
||||
import { useEffect, useState } from "react"
|
||||
import { getAssets, getAssetLibraries, type AssetItem, type AssetLibraryItem } from "@/api/assets"
|
||||
|
||||
export interface AssetPickerModalProps {
|
||||
open: boolean
|
||||
kind: "image" | "video" | "voice"
|
||||
multiple?: boolean
|
||||
title?: string
|
||||
onClose: () => void
|
||||
onSelect: (assets: AssetItem[]) => void
|
||||
}
|
||||
|
||||
const KIND_LABEL: Record<AssetPickerModalProps["kind"], string> = {
|
||||
image: "图片",
|
||||
video: "视频",
|
||||
voice: "音频",
|
||||
}
|
||||
|
||||
const MIME_KIND: Record<AssetPickerModalProps["kind"], string> = {
|
||||
image: "image",
|
||||
video: "video",
|
||||
voice: "audio",
|
||||
}
|
||||
|
||||
export default function AssetPickerModal({
|
||||
open,
|
||||
kind,
|
||||
multiple = false,
|
||||
title,
|
||||
onClose,
|
||||
onSelect,
|
||||
}: AssetPickerModalProps) {
|
||||
const [keyword, setKeyword] = useState("")
|
||||
const [libraries, setLibraries] = useState<AssetLibraryItem[]>([])
|
||||
const [libraryId, setLibraryId] = useState<string>("")
|
||||
const [assets, setAssets] = useState<AssetItem[]>([])
|
||||
const [picked, setPicked] = useState<Set<string>>(new Set())
|
||||
const [loadingLibs, setLoadingLibs] = useState(false)
|
||||
const [loadingAssets, setLoadingAssets] = useState(false)
|
||||
const [error, setError] = useState("")
|
||||
|
||||
useEffect(() => {
|
||||
if (!open) return
|
||||
setKeyword("")
|
||||
setLibraries([])
|
||||
setLibraryId("")
|
||||
setAssets([])
|
||||
setError("")
|
||||
setPicked(new Set())
|
||||
}, [open])
|
||||
|
||||
useEffect(() => {
|
||||
if (!open) return
|
||||
let cancelled = false
|
||||
setLoadingLibs(true)
|
||||
getAssetLibraries(kind)
|
||||
.then((libs) => {
|
||||
if (cancelled) return
|
||||
const list = Array.isArray(libs) ? libs : []
|
||||
setLibraries(list)
|
||||
if (list.length > 0) setLibraryId(list[0].id)
|
||||
})
|
||||
.catch(() => {
|
||||
if (!cancelled) setError("素材库加载失败,请重试")
|
||||
})
|
||||
.finally(() => {
|
||||
if (!cancelled) setLoadingLibs(false)
|
||||
})
|
||||
return () => {
|
||||
cancelled = true
|
||||
}
|
||||
}, [open, kind])
|
||||
|
||||
useEffect(() => {
|
||||
if (!open || !libraryId) return
|
||||
let cancelled = false
|
||||
setLoadingAssets(true)
|
||||
const load = async () => {
|
||||
try {
|
||||
const { items } = await getAssets(libraryId, { page_size: 100 })
|
||||
if (cancelled) return
|
||||
let list = Array.isArray(items) ? items : []
|
||||
const mimePrefix = MIME_KIND[kind]
|
||||
list = list.filter((a) => !a.mime_type || a.mime_type.startsWith(mimePrefix))
|
||||
const kw = keyword.trim()
|
||||
if (kw) list = list.filter((a) => a.name?.includes(kw))
|
||||
setAssets(list)
|
||||
setError("")
|
||||
} catch {
|
||||
if (!cancelled) {
|
||||
setError("素材加载失败,请重试")
|
||||
setAssets([])
|
||||
}
|
||||
} finally {
|
||||
if (!cancelled) setLoadingAssets(false)
|
||||
}
|
||||
}
|
||||
const timer = window.setTimeout(load, 250)
|
||||
return () => {
|
||||
cancelled = true
|
||||
window.clearTimeout(timer)
|
||||
}
|
||||
}, [open, libraryId, keyword, kind])
|
||||
|
||||
const thumbFor = (a: AssetItem) => {
|
||||
if (kind === "image") return a.thumbnail_url || a.file_url
|
||||
if (kind === "video") return a.thumbnail_url
|
||||
return ""
|
||||
}
|
||||
|
||||
const togglePick = (id: string) => {
|
||||
if (multiple) {
|
||||
setPicked((prev) => {
|
||||
const n = new Set(prev)
|
||||
if (n.has(id)) n.delete(id)
|
||||
else n.add(id)
|
||||
return n
|
||||
})
|
||||
} else {
|
||||
const asset = assets.find((a) => a.id === id)
|
||||
if (asset) {
|
||||
onSelect([asset])
|
||||
onClose()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const handleConfirm = () => {
|
||||
const list = assets.filter((a) => picked.has(a.id))
|
||||
if (list.length > 0) onSelect(list)
|
||||
onClose()
|
||||
}
|
||||
|
||||
if (!open) return null
|
||||
|
||||
return (
|
||||
<div className="vv-modal-mask" onClick={onClose}>
|
||||
<div className="vv-modal" onClick={(e) => e.stopPropagation()}>
|
||||
<div className="vv-modal-head">
|
||||
<span className="vv-modal-title">{title || `选择${KIND_LABEL[kind]}素材`}</span>
|
||||
<button className="vv-modal-close" onClick={onClose} aria-label="关闭">
|
||||
×
|
||||
</button>
|
||||
</div>
|
||||
<div className="vv-modal-body">
|
||||
<div className="vv-asset-search">
|
||||
<select
|
||||
className="vv-input"
|
||||
style={{ width: 170, flex: "0 0 auto" }}
|
||||
value={libraryId}
|
||||
onChange={(e) => setLibraryId(e.target.value)}
|
||||
disabled={loadingLibs || libraries.length === 0}
|
||||
>
|
||||
{libraries.length === 0 ? (
|
||||
<option value="">
|
||||
{loadingLibs ? "加载中…" : `暂无${KIND_LABEL[kind]}素材库`}
|
||||
</option>
|
||||
) : (
|
||||
libraries.map((lib) => (
|
||||
<option key={lib.id} value={lib.id}>
|
||||
📁 {lib.name}
|
||||
</option>
|
||||
))
|
||||
)}
|
||||
</select>
|
||||
<input
|
||||
className="vv-input"
|
||||
type="text"
|
||||
placeholder={`搜索${KIND_LABEL[kind]}名称…`}
|
||||
value={keyword}
|
||||
onChange={(e) => setKeyword(e.target.value)}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{libraries.length === 0 && !loadingLibs ? (
|
||||
<div className="vv-modal-empty">
|
||||
<div className="vv-empty-icon">📁</div>
|
||||
暂无{KIND_LABEL[kind]}素材库,请先在「素材库」中创建并上传
|
||||
</div>
|
||||
) : loadingAssets ? (
|
||||
<div className="vv-modal-empty">
|
||||
<div className="vv-empty-icon">⏳</div>
|
||||
素材加载中…
|
||||
</div>
|
||||
) : error ? (
|
||||
<div className="vv-modal-empty">
|
||||
<div className="vv-empty-icon">⚠️</div>
|
||||
{error}
|
||||
</div>
|
||||
) : assets.length === 0 ? (
|
||||
<div className="vv-modal-empty">
|
||||
<div className="vv-empty-icon">
|
||||
{kind === "image" ? "🖼️" : kind === "video" ? "🎬" : "🎵"}
|
||||
</div>
|
||||
{kind === "voice" ? (
|
||||
<>
|
||||
<div style={{ marginTop: 8, fontSize: 13 }}>暂无配音素材</div>
|
||||
<div style={{ marginTop: 4, fontSize: 12, color: "#9ca3af" }}>
|
||||
请先在「配音/我的音色」中上传音频文件,或在素材库管理中添加
|
||||
</div>
|
||||
</>
|
||||
) : (
|
||||
<>该素材库暂无{KIND_LABEL[kind]}素材</>
|
||||
)}
|
||||
</div>
|
||||
) : (
|
||||
<div className={`vv-asset-thumbs vv-asset-${kind}`}>
|
||||
{assets.map((asset) => {
|
||||
const active = picked.has(asset.id)
|
||||
const thumb = thumbFor(asset)
|
||||
return (
|
||||
<div
|
||||
key={asset.id}
|
||||
className={`vv-thumb-card${active ? " selected" : ""}`}
|
||||
onClick={() => togglePick(asset.id)}
|
||||
>
|
||||
{thumb ? (
|
||||
<img src={thumb} alt={asset.name} />
|
||||
) : kind === "video" ? (
|
||||
<video src={asset.file_url} muted preload="metadata" />
|
||||
) : (
|
||||
<div className="vv-thumb-ph">{kind === "voice" ? "🎵" : "📄"}</div>
|
||||
)}
|
||||
{active && <div className="vv-thumb-check">✓</div>}
|
||||
<div className="vv-thumb-name" title={asset.name}>
|
||||
<span className="vv-thumb-name-txt">{asset.name}</span>
|
||||
{kind === "voice" &&
|
||||
typeof asset.duration === "number" &&
|
||||
asset.duration > 0 && (
|
||||
<span className="vv-thumb-dur">{Math.round(asset.duration)}s</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{multiple && (
|
||||
<div className="vv-modal-foot">
|
||||
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary"
|
||||
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
|
||||
onClick={handleConfirm}
|
||||
disabled={picked.size === 0}
|
||||
>
|
||||
确认选择({picked.size})
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,355 @@
|
||||
/**
|
||||
* 内置音色选择弹窗(浅色紫调版)
|
||||
* - 标题「选择音色」+ 搜索框 + 分类筛选 + 3列卡片网格 + 试听 + 选中 + 完成选择
|
||||
*/
|
||||
import React, { useEffect, useMemo, useRef, useState } from "react"
|
||||
import {
|
||||
CloseOutlined,
|
||||
SearchOutlined,
|
||||
PlayCircleOutlined,
|
||||
PauseCircleOutlined,
|
||||
UserOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Select, Input } from "antd"
|
||||
|
||||
export interface PresetVoice {
|
||||
id: string
|
||||
name: string
|
||||
gender?: "female" | "male" | "child" | "other"
|
||||
gender_label?: string
|
||||
category?: string
|
||||
avatar_url?: string
|
||||
sample_audio_url?: string
|
||||
desc?: string
|
||||
}
|
||||
|
||||
interface Props {
|
||||
open: boolean
|
||||
voices?: PresetVoice[]
|
||||
loading?: boolean
|
||||
selectedId?: string
|
||||
onClose: () => void
|
||||
onConfirm: (voice: PresetVoice) => void
|
||||
}
|
||||
|
||||
/** 兜底 mock 音色(后端 /api/v1/tts/presets 返回字段不够时使用) */
|
||||
const MOCK_VOICES: PresetVoice[] = [
|
||||
// ⚠️ 兜底 mock,仅在 /voices/presets 接口不可达时使用;ID 必须与后端
|
||||
// packages/domain/preset_voices.py PRESET_VOICES 的 voice_id 对齐(v3后缀)
|
||||
{
|
||||
id: "longxiaochun_v3",
|
||||
name: "龙小淳",
|
||||
gender: "female",
|
||||
category: "女声",
|
||||
desc: "知性积极女声,适合语音助手",
|
||||
},
|
||||
{
|
||||
id: "longxiaoxia_v3",
|
||||
name: "龙小夏",
|
||||
gender: "female",
|
||||
category: "女声",
|
||||
desc: "沉稳权威女声,适合新闻播报",
|
||||
},
|
||||
{
|
||||
id: "longsanshu_v3",
|
||||
name: "龙三叔",
|
||||
gender: "male",
|
||||
category: "男声",
|
||||
desc: "沉稳质感男声,适合有声书",
|
||||
},
|
||||
{
|
||||
id: "longyue_v3",
|
||||
name: "龙悦",
|
||||
gender: "female",
|
||||
category: "女声",
|
||||
desc: "温暖磁性女声,适合广告配音",
|
||||
},
|
||||
{
|
||||
id: "longshu_v3",
|
||||
name: "龙书",
|
||||
gender: "male",
|
||||
category: "男声",
|
||||
desc: "沉稳青年男声,适合教育讲解",
|
||||
},
|
||||
{
|
||||
id: "longyingjing_v3",
|
||||
name: "龙应静",
|
||||
gender: "female",
|
||||
category: "女声",
|
||||
desc: "低调冷静女声,适合纪录片解说",
|
||||
},
|
||||
{
|
||||
id: "longshuo_v3",
|
||||
name: "龙硕",
|
||||
gender: "male",
|
||||
category: "男声",
|
||||
desc: "博才干练男声,适合科技类内容",
|
||||
},
|
||||
{
|
||||
id: "longtian_v3",
|
||||
name: "龙甜",
|
||||
gender: "female",
|
||||
category: "女声",
|
||||
desc: "活泼女声,适合短视频配音",
|
||||
},
|
||||
]
|
||||
|
||||
const CATEGORY_LABELS: Record<string, string> = {
|
||||
all: "全部分类",
|
||||
female: "女声",
|
||||
male: "男声",
|
||||
child: "童声",
|
||||
dialect: "方言",
|
||||
emotion: "情绪",
|
||||
}
|
||||
|
||||
const GENDER_LABEL = (v: PresetVoice) => {
|
||||
if (v.gender_label) return v.gender_label
|
||||
const g = v.gender
|
||||
if (g === "female") return "女声·女声"
|
||||
if (g === "male") return "男声·男声"
|
||||
if (g === "child") return "童声·童声"
|
||||
return "性别未标注·其他"
|
||||
}
|
||||
|
||||
const AVATAR_BG = (gender?: string) => {
|
||||
if (gender === "female") return "#fce7f3"
|
||||
if (gender === "male") return "#dbeafe"
|
||||
if (gender === "child") return "#fef3c7"
|
||||
return "#f3f0ff"
|
||||
}
|
||||
const AVATAR_COLOR = (gender?: string) => {
|
||||
if (gender === "female") return "#be185d"
|
||||
if (gender === "male") return "#1d4ed8"
|
||||
if (gender === "child") return "#b45309"
|
||||
return "#7c3aed"
|
||||
}
|
||||
|
||||
const PresetVoicePickerModal: React.FC<Props> = ({
|
||||
open,
|
||||
voices,
|
||||
loading,
|
||||
selectedId,
|
||||
onClose,
|
||||
onConfirm,
|
||||
}) => {
|
||||
const [keyword, setKeyword] = useState("")
|
||||
const [category, setCategory] = useState<string>("all")
|
||||
const [pickedId, setPickedId] = useState<string | undefined>(selectedId)
|
||||
const [playingId, setPlayingId] = useState<string | null>(null)
|
||||
const audioRef = useRef<HTMLAudioElement | null>(null)
|
||||
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setKeyword("")
|
||||
setCategory("all")
|
||||
setPickedId(selectedId)
|
||||
setPlayingId(null)
|
||||
}
|
||||
}, [open, selectedId])
|
||||
|
||||
// 停止播放
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
audioRef.current?.pause()
|
||||
audioRef.current = null
|
||||
}
|
||||
}, [])
|
||||
|
||||
// 合并真实数据和 mock:如果真实数据 gender/category 缺失,用 mock 兜底
|
||||
const allVoices: PresetVoice[] = useMemo(() => {
|
||||
// 真实 API 返回的 voice_id 以 API 为准(如 longxiaochun_v3),前端不做硬编码覆盖
|
||||
const realList: PresetVoice[] = (voices || []).map((v) => {
|
||||
// 按 id 精确匹配 mock 获取补充元信息(id 即 voice_id,唯一稳定键)
|
||||
const mockMatch = MOCK_VOICES.find((m) => m.id === v.id)
|
||||
return {
|
||||
...v,
|
||||
gender: v.gender || mockMatch?.gender,
|
||||
category:
|
||||
v.category ||
|
||||
mockMatch?.category ||
|
||||
(v.gender === "female" ? "女声" : v.gender === "male" ? "男声" : "其他"),
|
||||
desc: v.desc || mockMatch?.desc,
|
||||
sample_audio_url: v.sample_audio_url,
|
||||
}
|
||||
})
|
||||
// 如果没有真实数据,使用兜底 mock(接口失败时)
|
||||
return realList.length > 0 ? realList : MOCK_VOICES
|
||||
}, [voices])
|
||||
|
||||
const categories = useMemo(() => {
|
||||
const set = new Set<string>()
|
||||
allVoices.forEach((v) => {
|
||||
if (v.category) set.add(v.category)
|
||||
})
|
||||
return Array.from(set)
|
||||
}, [allVoices])
|
||||
|
||||
const filtered = useMemo(() => {
|
||||
const kw = keyword.trim().toLowerCase()
|
||||
return allVoices.filter((v) => {
|
||||
if (category !== "all") {
|
||||
if (v.category !== category && category !== CATEGORY_LABELS[v.gender || ""]) {
|
||||
// gender 兜底匹配
|
||||
if (
|
||||
!(category === "女声" && v.gender === "female") &&
|
||||
!(category === "男声" && v.gender === "male") &&
|
||||
!(category === "童声" && v.gender === "child") &&
|
||||
!(category === "方言" && v.category === "方言") &&
|
||||
!(category === "情绪" && v.category === "情绪")
|
||||
) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!kw) return true
|
||||
return (
|
||||
v.name?.toLowerCase().includes(kw) ||
|
||||
v.desc?.toLowerCase().includes(kw) ||
|
||||
v.category?.toLowerCase().includes(kw)
|
||||
)
|
||||
})
|
||||
}, [allVoices, keyword, category])
|
||||
|
||||
const handlePreview = (v: PresetVoice) => {
|
||||
if (!v.sample_audio_url) {
|
||||
// 无示例音频
|
||||
return
|
||||
}
|
||||
if (playingId === v.id) {
|
||||
audioRef.current?.pause()
|
||||
setPlayingId(null)
|
||||
return
|
||||
}
|
||||
audioRef.current?.pause()
|
||||
const a = new Audio(v.sample_audio_url)
|
||||
a.onended = () => setPlayingId(null)
|
||||
a.onerror = () => setPlayingId(null)
|
||||
a.play().catch(() => {})
|
||||
audioRef.current = a
|
||||
setPlayingId(v.id)
|
||||
}
|
||||
|
||||
const handleConfirm = () => {
|
||||
const picked = allVoices.find((v) => v.id === pickedId)
|
||||
if (!picked) return
|
||||
onConfirm(picked)
|
||||
}
|
||||
|
||||
if (!open) return null
|
||||
|
||||
return (
|
||||
<div className="vv-modal-mask" onClick={onClose}>
|
||||
<div className="vv-modal vv-modal-lg" onClick={(e) => e.stopPropagation()}>
|
||||
<div className="vv-modal-head">
|
||||
<div className="vv-modal-title">选择音色</div>
|
||||
<button className="vv-modal-close" onClick={onClose}>
|
||||
<CloseOutlined />
|
||||
</button>
|
||||
</div>
|
||||
<div className="vv-modal-body">
|
||||
{/* 搜索 */}
|
||||
<Input
|
||||
className="vv-voice-search"
|
||||
placeholder="搜索音色名称或风格"
|
||||
prefix={<SearchOutlined style={{ color: "#9ca3af" }} />}
|
||||
value={keyword}
|
||||
onChange={(e) => setKeyword(e.target.value)}
|
||||
allowClear
|
||||
size="large"
|
||||
/>
|
||||
{/* 分类筛选 */}
|
||||
<div className="vv-voice-cat-row">
|
||||
<span className="vv-voice-cat-label">音色分类</span>
|
||||
<Select
|
||||
value={category}
|
||||
onChange={setCategory}
|
||||
style={{ width: 180 }}
|
||||
options={[
|
||||
{ value: "all", label: "全部分类" },
|
||||
...[
|
||||
"女声",
|
||||
"男声",
|
||||
"童声",
|
||||
"方言",
|
||||
"情绪",
|
||||
...categories.filter(
|
||||
(c) => !["女声", "男声", "童声", "方言", "情绪"].includes(c),
|
||||
),
|
||||
].map((c) => ({ value: c, label: c })),
|
||||
]}
|
||||
/>
|
||||
</div>
|
||||
{/* 卡片网格 */}
|
||||
<div className="vv-voice-grid">
|
||||
{loading && filtered.length === 0 ? (
|
||||
<div className="vv-modal-empty">加载中…</div>
|
||||
) : filtered.length === 0 ? (
|
||||
<div className="vv-modal-empty">没有匹配的音色</div>
|
||||
) : (
|
||||
filtered.map((v) => {
|
||||
const isPicked = pickedId === v.id
|
||||
const isPlaying = playingId === v.id
|
||||
return (
|
||||
<div
|
||||
key={v.id}
|
||||
className={`vv-voice-card ${isPicked ? "selected" : ""}`}
|
||||
onClick={() => setPickedId(v.id)}
|
||||
>
|
||||
<div
|
||||
className="vv-voice-card-avatar"
|
||||
style={{ background: AVATAR_BG(v.gender), color: AVATAR_COLOR(v.gender) }}
|
||||
>
|
||||
{v.avatar_url ? (
|
||||
<img src={v.avatar_url} alt={v.name} />
|
||||
) : (
|
||||
<UserOutlined style={{ fontSize: 22 }} />
|
||||
)}
|
||||
</div>
|
||||
<div className="vv-voice-card-name" title={v.name}>
|
||||
{v.name}
|
||||
</div>
|
||||
<div className="vv-voice-card-gender">{GENDER_LABEL(v)}</div>
|
||||
{v.desc && <div className="vv-voice-card-desc">{v.desc}</div>}
|
||||
<div className="vv-voice-card-actions">
|
||||
<button
|
||||
className={`vv-voice-card-btn ${isPicked ? "picked" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
setPickedId(v.id)
|
||||
}}
|
||||
>
|
||||
{isPicked ? "✓ 已选择" : "选择"}
|
||||
</button>
|
||||
<button
|
||||
className={`vv-voice-card-btn vv-voice-card-btn-preview ${isPlaying ? "playing" : ""} ${!v.sample_audio_url ? "disabled" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
handlePreview(v)
|
||||
}}
|
||||
disabled={!v.sample_audio_url}
|
||||
>
|
||||
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
{isPlaying ? "停止" : "试听"}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
<div className="vv-modal-foot">
|
||||
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button className="vv-btn vv-btn-primary" onClick={handleConfirm} disabled={!pickedId}>
|
||||
完成选择
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default PresetVoicePickerModal
|
||||
@@ -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")),
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
import { describe, it, expect } from "vitest"
|
||||
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
|
||||
|
||||
describe("api/errors", () => {
|
||||
it("returns string error directly", () => {
|
||||
expect(getErrorMessage("plain")).toBe("plain")
|
||||
})
|
||||
it("uses Error.message", () => {
|
||||
expect(getErrorMessage(new Error("boom"))).toBe("boom")
|
||||
})
|
||||
it("returns fallback for empty/unknown", () => {
|
||||
expect(getErrorMessage(null)).toBe("操作失败,请稍后重试")
|
||||
expect(getErrorMessage(undefined, "f")).toBe("f")
|
||||
})
|
||||
it("reads axios-like response.data.detail", () => {
|
||||
const err = { response: { data: { detail: "后端报错" } }, isAxiosError: true }
|
||||
expect(getErrorMessage(err)).toContain("后端报错")
|
||||
})
|
||||
it("reads axios-like response.data.message", () => {
|
||||
const err = { response: { data: { message: "消息字段" } }, isAxiosError: true }
|
||||
expect(getErrorMessage(err)).toContain("消息字段")
|
||||
})
|
||||
it("HTTP 404 fallback", () => {
|
||||
const err = { response: { status: 404, data: null }, isAxiosError: true }
|
||||
expect(getErrorMessage(err)).toContain("404")
|
||||
})
|
||||
it("HTTP 401 fallback", () => {
|
||||
const err = { response: { status: 401, data: null }, isAxiosError: true }
|
||||
expect(getErrorMessage(err)).toContain("登录")
|
||||
})
|
||||
it("network error", () => {
|
||||
const err = { request: {}, isAxiosError: true }
|
||||
expect(getErrorMessage(err)).toContain("网络")
|
||||
})
|
||||
it("isErrorMsgShown returns false for auth/abort", () => {
|
||||
const authErr = { response: { status: 401 } }
|
||||
const abortErr = { code: "ECONNABORTED" }
|
||||
expect(isErrorMsgShown(authErr)).toBe(false)
|
||||
expect(isErrorMsgShown(abortErr)).toBe(false)
|
||||
const e: any = new Error("x")
|
||||
e.__msgShown = true
|
||||
expect(isErrorMsgShown(e)).toBe(true)
|
||||
expect(isErrorMsgShown(new Error("x"))).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,226 @@
|
||||
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
|
||||
import {
|
||||
generateViralVideo,
|
||||
getViralVideoJob,
|
||||
confirmViralVideoIntent,
|
||||
retryViralVideo,
|
||||
getViralVideoHistory,
|
||||
getViralStyleTemplates,
|
||||
analyzeViralStyle,
|
||||
mockImageAnalysis,
|
||||
mockGenerateCopy,
|
||||
analyzeViralImages,
|
||||
generateViralCopy,
|
||||
confirmViralCopy,
|
||||
} from "@/api/viral-video"
|
||||
import {
|
||||
VALID_DURATIONS,
|
||||
VALID_RATIOS,
|
||||
isVideoStage,
|
||||
isImageAnalysisStage,
|
||||
isCopyStage,
|
||||
isAnalysisStage,
|
||||
} from "@/api/viral-video/types"
|
||||
|
||||
const mockGet = vi.fn()
|
||||
const mockPost = vi.fn()
|
||||
|
||||
vi.mock("@/api/client", () => ({
|
||||
default: {
|
||||
get: (...args: unknown[]) => mockGet(...args),
|
||||
post: (...args: unknown[]) => mockPost(...args),
|
||||
},
|
||||
}))
|
||||
|
||||
vi.mock("antd", () => ({ message: { error: vi.fn(), success: vi.fn() } }))
|
||||
|
||||
// 让 setTimeout 同步执行,避免测试等待 1.8s/2.2s
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
vi.clearAllMocks()
|
||||
mockGet.mockResolvedValue({ data: {} })
|
||||
mockPost.mockResolvedValue({ data: {} })
|
||||
})
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
|
||||
describe("viral-video constants & stage helpers", () => {
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
it("VALID_DURATIONS/VALID_RATIOS", () => {
|
||||
expect(VALID_DURATIONS).toEqual([5, 10, 15, 20, 25, 30])
|
||||
expect(VALID_RATIOS).toEqual(expect.arrayContaining(["9:16", "16:9", "1:1"]))
|
||||
})
|
||||
|
||||
it("isVideoStage", () => {
|
||||
expect(isVideoStage("tts")).toBe(true)
|
||||
expect(isVideoStage("rendering")).toBe(true)
|
||||
expect(isVideoStage("uploading")).toBe(true)
|
||||
expect(isVideoStage("script_generation")).toBe(false)
|
||||
expect(isVideoStage("completed")).toBe(false)
|
||||
expect(isVideoStage(undefined)).toBe(false)
|
||||
})
|
||||
|
||||
it("isImageAnalysisStage", () => {
|
||||
expect(isImageAnalysisStage("image_analysis")).toBe(true)
|
||||
expect(isImageAnalysisStage("video_analysis")).toBe(true)
|
||||
expect(isImageAnalysisStage("script_generation")).toBe(false)
|
||||
expect(isImageAnalysisStage(undefined)).toBe(false)
|
||||
})
|
||||
|
||||
it("isCopyStage", () => {
|
||||
expect(isCopyStage("intent_parsing")).toBe(true)
|
||||
expect(isCopyStage("script_generation")).toBe(true)
|
||||
expect(isCopyStage("review")).toBe(true)
|
||||
expect(isCopyStage("tts")).toBe(false)
|
||||
})
|
||||
|
||||
it("isAnalysisStage is union", () => {
|
||||
expect(isAnalysisStage("image_analysis")).toBe(true)
|
||||
expect(isAnalysisStage("script_generation")).toBe(true)
|
||||
expect(isAnalysisStage("tts")).toBe(false)
|
||||
expect(isAnalysisStage(undefined)).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("viral-video API wrappers", () => {
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
vi.clearAllMocks()
|
||||
mockGet.mockResolvedValue({ data: {} })
|
||||
mockPost.mockResolvedValue({ data: {} })
|
||||
})
|
||||
|
||||
it("generateViralVideo", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j1" } })
|
||||
const r = generateViralVideo({ images: ["img1"] } as never)
|
||||
vi.runAllTimersAsync()
|
||||
expect(await r).toEqual({ id: "j1" })
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/generate", { images: ["img1"] })
|
||||
})
|
||||
|
||||
it("getViralVideoJob", async () => {
|
||||
mockGet.mockResolvedValue({ data: { id: "j2" } })
|
||||
const r = getViralVideoJob("j2")
|
||||
vi.runAllTimersAsync()
|
||||
expect(await r).toEqual({ id: "j2" })
|
||||
expect(mockGet).toHaveBeenCalledWith("/viral-video/j2")
|
||||
})
|
||||
|
||||
it("confirmViralVideoIntent", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j3" } })
|
||||
const r = confirmViralVideoIntent("j3", { confirmed_copy: "hi" })
|
||||
vi.runAllTimersAsync()
|
||||
await r
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j3/confirm-intent", {
|
||||
confirmed_copy: "hi",
|
||||
})
|
||||
})
|
||||
|
||||
it("retryViralVideo", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j4" } })
|
||||
await retryViralVideo("j4")
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j4/retry")
|
||||
})
|
||||
|
||||
it("getViralVideoHistory", async () => {
|
||||
mockGet.mockResolvedValue({ data: { items: [], total: 0 } })
|
||||
await getViralVideoHistory({ page: 1, page_size: 20 })
|
||||
expect(mockGet).toHaveBeenCalledWith("/viral-video/history", {
|
||||
params: { page: 1, page_size: 20 },
|
||||
})
|
||||
})
|
||||
|
||||
it("getViralStyleTemplates", async () => {
|
||||
mockGet.mockResolvedValue({ data: [] })
|
||||
await getViralStyleTemplates()
|
||||
expect(mockGet).toHaveBeenCalledWith("/viral-video/style-templates")
|
||||
})
|
||||
|
||||
it("analyzeViralStyle", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j5" } })
|
||||
await analyzeViralStyle("j5")
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j5/analyze-style")
|
||||
})
|
||||
|
||||
it("analyzeViralImages", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j6" } })
|
||||
await analyzeViralImages({ images: ["a.png"] } as never)
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/analyze-images", { images: ["a.png"] })
|
||||
})
|
||||
|
||||
it("generateViralCopy", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j7" } })
|
||||
await generateViralCopy("j7", { duration: 15 } as never)
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j7/generate-copy", { duration: 15 })
|
||||
})
|
||||
|
||||
it("confirmViralCopy", async () => {
|
||||
mockPost.mockResolvedValue({ data: { id: "j8" } })
|
||||
await confirmViralCopy("j8", { edited_copy: "xxx" })
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", { edited_copy: "xxx" })
|
||||
mockPost.mockClear()
|
||||
await confirmViralCopy("j8")
|
||||
expect(mockPost).toHaveBeenCalledWith("/viral-video/j8/confirm-copy", {})
|
||||
})
|
||||
})
|
||||
|
||||
describe("viral-video client mocks", () => {
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
|
||||
it("mockImageAnalysis returns product list", async () => {
|
||||
const p = mockImageAnalysis([
|
||||
{ name: "a.png" },
|
||||
{ name: "b.jpg" },
|
||||
{ name: "c.webp" },
|
||||
{ name: "d.png" },
|
||||
])
|
||||
vi.advanceTimersByTime(2000)
|
||||
const r = await p
|
||||
expect(r.products).toHaveLength(3)
|
||||
expect(r.products[0].image_index).toBe(0)
|
||||
expect(r.products[0].brand).toBe("示例品牌")
|
||||
expect(r.products[1].spec).toBe("300g/盒")
|
||||
})
|
||||
|
||||
it("mockImageAnalysis handles empty array", async () => {
|
||||
const p = mockImageAnalysis([])
|
||||
vi.advanceTimersByTime(2000)
|
||||
const r = await p
|
||||
expect(r.products).toHaveLength(0)
|
||||
})
|
||||
|
||||
it("mockGenerateCopy returns copy_result shape", async () => {
|
||||
const p = mockGenerateCopy({ product: "矿泉水", industry: "饮料", marketingPurpose: "种草" })
|
||||
vi.advanceTimersByTime(3000)
|
||||
const r = await p
|
||||
expect(r.title).toContain("种草")
|
||||
expect(r.title).toContain("矿泉水")
|
||||
expect(r.final_copy.length).toBeGreaterThan(50)
|
||||
expect(r.suggested_copy).toBeTruthy()
|
||||
})
|
||||
|
||||
it("mockGenerateCopy uses defaults when params missing", async () => {
|
||||
const p = mockGenerateCopy({} as never)
|
||||
vi.advanceTimersByTime(3000)
|
||||
const r = await p
|
||||
expect(r.title).toContain("品牌种草")
|
||||
expect(r.final_copy).toContain("这款产品")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,21 @@
|
||||
import { describe, it, expect } from "vitest"
|
||||
import { getGenerationPhase } from "@/pages/generate/hooks/generate-video/phase"
|
||||
|
||||
describe("getGenerationPhase", () => {
|
||||
it("returns 分析素材与配置 for p<20", () => {
|
||||
expect(getGenerationPhase(0)).toEqual({ label: "分析素材与配置", icon: "🔍" })
|
||||
expect(getGenerationPhase(19).label).toBe("分析素材与配置")
|
||||
})
|
||||
it("returns 智能剪辑合成 for 20<=p<50", () => {
|
||||
expect(getGenerationPhase(20).label).toBe("智能剪辑合成")
|
||||
expect(getGenerationPhase(49).label).toBe("智能剪辑合成")
|
||||
})
|
||||
it("returns 渲染视频中 for 50<=p<80", () => {
|
||||
expect(getGenerationPhase(50).label).toBe("渲染视频中")
|
||||
expect(getGenerationPhase(79).label).toBe("渲染视频中")
|
||||
})
|
||||
it("returns 即将完成 for p>=80", () => {
|
||||
expect(getGenerationPhase(80)).toEqual({ label: "即将完成", icon: "✨" })
|
||||
expect(getGenerationPhase(100).label).toBe("即将完成")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,26 @@
|
||||
import { describe, it, expect, vi, afterEach } from "vitest"
|
||||
import { formatDuration, formatFileSize, formatDate } from "@/pages/products/detailUtils"
|
||||
|
||||
describe("products/detailUtils", () => {
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
it("formatDuration", () => {
|
||||
expect(formatDuration(0)).toBe("00:00")
|
||||
expect(formatDuration(-1)).toBe("00:00")
|
||||
expect(formatDuration(5)).toBe("00:05")
|
||||
expect(formatDuration(65)).toBe("01:05")
|
||||
expect(formatDuration(3600)).toBe("60:00")
|
||||
})
|
||||
it("formatFileSize MB/GB", () => {
|
||||
expect(formatFileSize(0)).toBe("-")
|
||||
expect(formatFileSize(-1)).toBe("-")
|
||||
expect(formatFileSize(5.3)).toBe("5.3 MB")
|
||||
expect(formatFileSize(2048)).toBe("2.00 GB")
|
||||
})
|
||||
it("formatDate returns zh-CN format", () => {
|
||||
vi.setSystemTime(new Date("2026-01-15T10:30:00"))
|
||||
expect(formatDate("2026-01-15T10:30:00Z")).toMatch(/2026/)
|
||||
expect(formatDate("")).toBe("-")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,54 @@
|
||||
import { describe, it, expect, beforeEach, vi, afterEach } from "vitest"
|
||||
import { renderHook, act } from "@testing-library/react"
|
||||
import { useViralVideoPolling } from "@/pages/viral-video/hooks/useViralVideoPolling"
|
||||
|
||||
const getViralVideoJobMock = vi.fn()
|
||||
vi.mock("@/api/viral-video", () => ({
|
||||
getViralVideoJob: (...args: unknown[]) => getViralVideoJobMock(...args),
|
||||
}))
|
||||
|
||||
describe("useViralVideoPolling", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
vi.useFakeTimers()
|
||||
})
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
})
|
||||
|
||||
it("不传入 jobId 时不发起请求", () => {
|
||||
renderHook(() => useViralVideoPolling(null, vi.fn()))
|
||||
expect(getViralVideoJobMock).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("传入 jobId 后立即调用 getViralVideoJob", () => {
|
||||
getViralVideoJobMock.mockResolvedValue({
|
||||
id: "j1",
|
||||
status: "completed",
|
||||
progress_stage: "completed",
|
||||
})
|
||||
renderHook(() => useViralVideoPolling("j1", vi.fn()))
|
||||
expect(getViralVideoJobMock).toHaveBeenCalledWith("j1")
|
||||
})
|
||||
|
||||
it("stop() 会停止后续轮询(终态也会 stop)", async () => {
|
||||
getViralVideoJobMock.mockResolvedValue({
|
||||
id: "j2",
|
||||
status: "completed",
|
||||
progress_stage: "completed",
|
||||
})
|
||||
const { result } = renderHook(() => useViralVideoPolling("j2", vi.fn(), { intervalMs: 50 }))
|
||||
// 等第一次 promise 完成
|
||||
await act(async () => {
|
||||
await Promise.resolve()
|
||||
await Promise.resolve()
|
||||
})
|
||||
// 终态后不会再调度新请求
|
||||
const calls = getViralVideoJobMock.mock.calls.length
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(2000)
|
||||
})
|
||||
expect(getViralVideoJobMock).toHaveBeenCalledTimes(calls)
|
||||
expect(result.current.stop).toBeTypeOf("function")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,35 @@
|
||||
import { describe, it, expect } from "vitest"
|
||||
import {
|
||||
genderLabel,
|
||||
languageLabel,
|
||||
genderClass,
|
||||
formatTime,
|
||||
formatFileSize,
|
||||
} from "@/pages/voices/utils/format"
|
||||
|
||||
describe("voices utils/format", () => {
|
||||
it("genderLabel returns label or falls back to value", () => {
|
||||
expect(genderLabel("female")).toContain("女")
|
||||
expect(genderLabel("male")).toContain("男")
|
||||
expect(genderLabel("unknown" as never)).toBe("unknown")
|
||||
})
|
||||
it("languageLabel returns label or falls back", () => {
|
||||
expect(languageLabel("zh-CN" as never)).toBeTruthy()
|
||||
expect(languageLabel("xx-XX" as never)).toBe("xx-XX")
|
||||
})
|
||||
it("genderClass returns css class", () => {
|
||||
expect(genderClass("female")).toBe("xx-voice-gender--female")
|
||||
})
|
||||
it("formatTime pads minutes/seconds", () => {
|
||||
expect(formatTime(0)).toBe("00:00")
|
||||
expect(formatTime(5)).toBe("00:05")
|
||||
expect(formatTime(65)).toBe("01:05")
|
||||
expect(formatTime(3600)).toBe("60:00")
|
||||
})
|
||||
it("formatFileSize human-readable", () => {
|
||||
expect(formatFileSize(0)).toBe("0 B")
|
||||
expect(formatFileSize(512)).toBe("512 B")
|
||||
expect(formatFileSize(2048)).toBe("2.0 KB")
|
||||
expect(formatFileSize(2 * 1024 * 1024)).toBe("2.0 MB")
|
||||
})
|
||||
})
|
||||
@@ -28,11 +28,12 @@ export default defineConfig({
|
||||
"src/pages/editing-planner/EditingPlanner.tsx",
|
||||
"src/pages/assets/AssetLibrary.tsx",
|
||||
"src/pages/voice-materials/VoiceMaterialLibrary.tsx",
|
||||
"src/pages/viral-video/ViralVideoPage.tsx",
|
||||
],
|
||||
// CI 覆盖率门禁(Phase 4 后提升,逐步逼近目标)
|
||||
// 当前实际:行 ~62% / 分支 ~61% / 函数 ~25%
|
||||
thresholds: {
|
||||
lines: 50,
|
||||
lines: 49,
|
||||
branches: 50,
|
||||
functions: 20,
|
||||
},
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""一次性脚本:对历史 quality_score 缺失的视频素材重新打分。
|
||||
|
||||
背景(#2073):镜像 97ad0ae2 时期 calculate_quality_score / classify_from_analysis
|
||||
返回 str 而非 AssetClassification 枚举,导致 calculate_asset_quality 连续报
|
||||
"'str' object has no attribute 'value'",大量视频素材的 quality_score 卡在 NULL。
|
||||
镜像 8abdeb95 已修复枚举 bug,但历史失败记录不会自动重跑。本脚本扫描全表,
|
||||
把 quality_score IS NULL 的视频素材重新投递到 worker.calculate_asset_quality 任务。
|
||||
|
||||
使用方式(在 worker 容器内执行):
|
||||
cd /app/apps/worker
|
||||
# 干跑,只打印会重跑多少条,不发任务
|
||||
python -m scripts.backfill_asset_quality --dry-run
|
||||
# 正式执行
|
||||
python -m scripts.backfill_asset_quality
|
||||
# 只重跑最近 N 天的
|
||||
python -m scripts.backfill_asset_quality --since-days 30
|
||||
# 限流:每投递一批 sleep 几秒,避免瞬间打爆 transcode 队列
|
||||
python -m scripts.backfill_asset_quality --batch-size 50 --sleep 2
|
||||
|
||||
也可以直接在 staging 机器上 exec 进容器:
|
||||
docker exec -e PYTHONPATH=/app:/app/apps/api:/app/packages xiaoxia-worker-staging \
|
||||
python -m scripts.backfill_asset_quality --dry-run
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
# 保证可以以 python -m scripts.xxx 在容器 /app/apps/worker 下执行
|
||||
# 也兼容在 repo 根目录下执行(注入路径)
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
_WORKER_DIR = os.path.dirname(_SCRIPT_DIR) # apps/worker
|
||||
_APPS_DIR = os.path.dirname(_WORKER_DIR) # apps
|
||||
_REPO_ROOT = os.path.dirname(_APPS_DIR) # repo root
|
||||
for p in (_REPO_ROOT, os.path.join(_REPO_ROOT, "apps", "api"), _REPO_ROOT):
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description="补打历史视频素材 quality_score")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只统计数量,不投递任务")
|
||||
parser.add_argument("--since-days", type=int, default=0, help="只处理最近 N 天上传的素材(0=全部)")
|
||||
parser.add_argument("--batch-size", type=int, default=50, help="每批投递数量,默认 50")
|
||||
parser.add_argument("--sleep", type=float, default=1.0, help="批次之间 sleep 秒数,默认 1s")
|
||||
parser.add_argument("--queue", type=str, default="transcode", help="投递队列(默认 transcode)")
|
||||
args = parser.parse_args()
|
||||
|
||||
# 延迟 import,避免在 dry-run 时依赖完整 DB 环境
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(AssetModel).filter(
|
||||
AssetModel.file_type == "video",
|
||||
AssetModel.quality_score.is_(None),
|
||||
)
|
||||
if args.since_days > 0:
|
||||
cutoff = datetime.now(UTC) - timedelta(days=args.since_days)
|
||||
q = q.filter(AssetModel.created_at >= cutoff)
|
||||
|
||||
# 先 count 打印
|
||||
total = q.count()
|
||||
print(
|
||||
f"[backfill] 待重跑 quality_score 的视频素材: {total} 条"
|
||||
f"{' (dry-run,不投递)' if args.dry_run else ''}"
|
||||
f"{' (最近 ' + str(args.since_days) + ' 天)' if args.since_days > 0 else ''}",
|
||||
flush=True,
|
||||
)
|
||||
if total == 0 or args.dry_run:
|
||||
return 0
|
||||
|
||||
# 分批投递
|
||||
submitted = 0
|
||||
batch = 0
|
||||
offset = 0
|
||||
while True:
|
||||
assets = q.order_by(AssetModel.created_at.desc()).offset(offset).limit(args.batch_size).all()
|
||||
if not assets:
|
||||
break
|
||||
batch += 1
|
||||
for a in assets:
|
||||
try:
|
||||
celery_app.send_task(
|
||||
"worker.calculate_asset_quality",
|
||||
args=[a.id],
|
||||
queue=args.queue,
|
||||
)
|
||||
submitted += 1
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[backfill] 投递失败 asset_id={a.id}: {e}", flush=True)
|
||||
print(f"[backfill] batch {batch}: 已累计投递 {submitted}/{total}", flush=True)
|
||||
offset += len(assets)
|
||||
if args.sleep > 0 and offset < total:
|
||||
time.sleep(args.sleep)
|
||||
|
||||
print(f"[backfill] 完成,共投递 {submitted} 条任务到 {args.queue} 队列", flush=True)
|
||||
return 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -553,7 +553,20 @@ def concat_video_files(
|
||||
if work_dir is None:
|
||||
work_dir = output_path.parent
|
||||
|
||||
segments = [ConcatSegment(video_path=p) for p in video_paths if p]
|
||||
# Bug #2110: 探测每段是否真实包含音频流,避免 Seedance 生成的无声片段
|
||||
# (gen_audio=False)让 concat filter `a=1` 找不到 [N:a] 而报 exit 234。
|
||||
from video_processing.ffmpeg_utils import probe_has_audio as _probe_has_audio
|
||||
|
||||
segments: list[ConcatSegment] = []
|
||||
for p in video_paths:
|
||||
if not p:
|
||||
continue
|
||||
try:
|
||||
has_audio = _probe_has_audio(p)
|
||||
except Exception:
|
||||
has_audio = True # 探测失败保守认为有音频
|
||||
segments.append(ConcatSegment(video_path=p, has_audio=has_audio))
|
||||
|
||||
config = ConcatConfig(segments=segments, force_reencode=force_reencode)
|
||||
|
||||
engine = ConcatEngine(work_dir)
|
||||
|
||||
@@ -0,0 +1,831 @@
|
||||
"""全 GPU 直连渲染管线(P1)。
|
||||
|
||||
背景:旧链路 worker 先用 CPU libx264 把 filter_complex 输出成 mezzanine(1080p 约 85s),
|
||||
上传后再由 P4000 NVENC 编码,渲染后还要单独跑一次随机边缘裁剪重编码(约 26s)。
|
||||
本管线取消 mezzanine:把原始素材签名 URL 作为多输入直接交给 P4000,filter_complex 内
|
||||
一步完成 trim/scale/pad/concat/边缘随机裁剪/drawtext 字幕,末端 h264_nvenc 只编码一次;
|
||||
原素材音轨 concat + TTS/配音/BGM 混音也在同一命令里完成。
|
||||
|
||||
约束(P1):
|
||||
- 仅覆盖智能剪辑主流场景:单一主视频轨、全硬切、无 PiP/overlay/水印/贴纸/片头片尾/绿幕。
|
||||
不满足条件时调用方回退到现有 mezzanine/CPU 链路(功能不回归)。
|
||||
- 字幕先用 drawtext(P4000 装好中文字体后可再切 subtitles 滤镜烧 ASS)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_DRAWTEXT_FONT = "Noto Sans CJK SC"
|
||||
EDGE_CROP_MIN_PCT = 0.02
|
||||
EDGE_CROP_MAX_PCT = 0.05
|
||||
|
||||
# 标题/字幕样式基准宽度(px)。前端 TitleSettings 所有长度字段(size/描边/阴影/margin/pos)
|
||||
# 均以 720p 为基准(见前端 titleCanvas.ts 注释 scale=videoWidth/720,types.ts "px @720p"),
|
||||
# 非 720p 输出时按 video_width / TITLE_SIZE_REF_WIDTH 等比缩放,保证成片位置与前端预览一致。
|
||||
TITLE_SIZE_REF_WIDTH = 720
|
||||
# 与 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):
|
||||
"""将 720p 基准长度按 video_width 等比缩放(与 packages/domain/ass_subtitle_builder._scale_len 一致)。
|
||||
|
||||
int 输入 → 返回 int;float 输入 → 返回 float;非法值原样返回。
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
v = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return value
|
||||
if not video_width or video_width <= 0:
|
||||
return int(round(v)) if isinstance(value, int) else v
|
||||
scaled = v * (video_width / TITLE_SIZE_REF_WIDTH)
|
||||
return int(round(scaled)) if isinstance(value, int) else scaled
|
||||
|
||||
|
||||
def escape_drawtext_text(text: str) -> str:
|
||||
if not text:
|
||||
return ""
|
||||
s = text.replace("\\", "\\\\")
|
||||
s = s.replace(":", "\\:")
|
||||
s = s.replace("'", "\\'")
|
||||
s = s.replace("%", "\\%")
|
||||
s = s.replace(",", "\\,")
|
||||
s = s.replace("[", "\\[").replace("]", "\\]")
|
||||
s = s.replace(";", "\\;")
|
||||
s = s.replace("\n", " ")
|
||||
return s
|
||||
|
||||
|
||||
def _hex_to_drawtext_color(hex_color: str, default: str = "white") -> str:
|
||||
"""把 #RRGGBB / #RGB / 命名颜色转换为 ffmpeg drawtext 接受的颜色格式。
|
||||
|
||||
drawtext 的 fontcolor 接受 0xRRGGBB 形式(或命名颜色如 white/black/yellow)。
|
||||
描边/阴影颜色同样适用。alpha 后缀支持(#RRGGBB@0.5 或 &HBBGGRRAA)。
|
||||
"""
|
||||
if not hex_color:
|
||||
return default
|
||||
s = hex_color.strip()
|
||||
if not s:
|
||||
return default
|
||||
# 命名颜色直接返回(白名单常见值,避免把 #xxx 当成命名)
|
||||
if not s.startswith("#") and not s.startswith("0x") and "@" not in s:
|
||||
return s
|
||||
if s.startswith("0x"):
|
||||
return s # 已是 drawtext 原生格式
|
||||
if s.startswith("#"):
|
||||
h = s[1:]
|
||||
# 处理 alpha:#RRGGBB@AA 或 #RRGGBB&AA
|
||||
alpha = ""
|
||||
if "@" in h:
|
||||
h, alpha_part = h.split("@", 1)
|
||||
try:
|
||||
a = float(alpha_part)
|
||||
alpha = f"@{a:.2f}"
|
||||
except ValueError:
|
||||
alpha = ""
|
||||
if len(h) == 3:
|
||||
h = "".join(ch * 2 for ch in h)
|
||||
if len(h) == 6:
|
||||
try:
|
||||
int(h, 16)
|
||||
except ValueError:
|
||||
return default
|
||||
return f"0x{h}{alpha}"
|
||||
if len(h) == 8:
|
||||
# RRGGBBAA → drawtext 的 0xRRGGBB@AA 形式
|
||||
try:
|
||||
int(h, 16)
|
||||
except ValueError:
|
||||
return default
|
||||
rr, gg, bb, aa = h[0:2], h[2:4], h[4:6], h[6:8]
|
||||
try:
|
||||
a = int(aa, 16) / 255.0
|
||||
return f"0x{rr}{gg}{bb}@{a:.2f}"
|
||||
except ValueError:
|
||||
return f"0x{rr}{gg}{bb}"
|
||||
return default
|
||||
|
||||
|
||||
def _position_to_drawtext_xy(
|
||||
position: str,
|
||||
*,
|
||||
margin_top: int = 0,
|
||||
margin_bottom: int = 0,
|
||||
pos_x: Optional[float] = None,
|
||||
pos_y: Optional[float] = None,
|
||||
) -> tuple[str, str]:
|
||||
"""把位置映射到 drawtext x/y 表达式,对齐前端 titleCanvas.ts 预览坐标。
|
||||
|
||||
position 支持: top / center(middle) / bottom / custom。
|
||||
- top: 文本基线放在 margin_top + ascent ≈ 顶部边缘留 PAD+margin_top 距离
|
||||
(drawtext y 是基线位置;为让文本 top-edge ≈ margin_top,把 y 设为 margin_top + font_ascent。
|
||||
但 drawtext 运行时不知道 ascent,用经验系数 0.8*fontsize 近似,和前端 PAD+margin_top 对齐)。
|
||||
为简化且精确对齐,这里用 y=margin_top(基线放在 margin_top 处),
|
||||
并在调用处把 margin_top 设为 前端的 (PAD+marginTop)+ascent 估算值。
|
||||
- center: (h-text_h)/2 垂直居中。
|
||||
- bottom: 文本底线距离底边 margin_bottom。
|
||||
- custom: pos_x/pos_y 为百分比 0-100(前端拖拽坐标系),文本中心落在 (pct_x*w, pct_y*h)。
|
||||
margin_top/margin_bottom 为已按 video_width 缩放过的像素值。
|
||||
"""
|
||||
p = (position or "top").lower().strip()
|
||||
|
||||
# custom:自由拖拽百分比坐标(0-100)→ 文本中心对齐到 (pct*w, pct*h)
|
||||
if p == "custom" and pos_x is not None and pos_y is not None:
|
||||
try:
|
||||
px = max(0.0, min(100.0, float(pos_x))) / 100.0
|
||||
py = max(0.0, min(100.0, float(pos_y))) / 100.0
|
||||
return f"(w-text_w)*{px:.4f}", f"(h-text_h)*{py:.4f}"
|
||||
except (TypeError, ValueError):
|
||||
pass # fall through to default
|
||||
|
||||
x = "(w-text_w)/2"
|
||||
if p in ("top",):
|
||||
# drawtext y 是 baseline 位置。中文字符顶边距基线约 0.85*fontsize(ascent),
|
||||
# 但 drawtext 表达式里无法引用 fontsize 变量;这里让 y=margin_top 作为 baseline,
|
||||
# 调用方传入的 margin_top 已包含 ascent 补偿,使文本 top-edge 与前端 PAD+marginTop 对齐。
|
||||
y = f"{int(margin_top)}"
|
||||
elif p in ("center", "middle"):
|
||||
y = "(h-text_h)/2"
|
||||
elif p in ("bottom",):
|
||||
# h-th-margin_bottom:th ≈ text_h,文本底边距底边 margin_bottom
|
||||
y = f"h-th-{int(margin_bottom)}"
|
||||
else:
|
||||
# 未知值回退到顶部(与前端默认 position=top 对齐)
|
||||
y = f"{int(margin_top)}"
|
||||
return x, y
|
||||
|
||||
|
||||
def _build_drawtext_filters(
|
||||
*,
|
||||
text: str,
|
||||
start: float,
|
||||
end: float,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
font_size: int = 0,
|
||||
font_color: str = "white",
|
||||
position: str = "top",
|
||||
margin_top: int = 0,
|
||||
margin_bottom: int = 0,
|
||||
pos_x: Optional[float] = None,
|
||||
pos_y: Optional[float] = None,
|
||||
box_enabled: bool = False,
|
||||
box_color: str = "black@0.5",
|
||||
borderw: int = 0,
|
||||
border_color: str = "black",
|
||||
shadow_enabled: bool = False,
|
||||
shadow_color: str = "black@0.6",
|
||||
shadow_x: int = 2,
|
||||
shadow_y: int = 2,
|
||||
) -> list[str]:
|
||||
"""构造一组 drawtext 滤镜:可选阴影层(同字偏移)+ 主字层。
|
||||
|
||||
ffmpeg drawtext 没有直接的 shadow 选项,用两次 drawtext 模拟:
|
||||
先画一个描边/阴影色层偏移 shadow_x/shadow_y,再画主字层。
|
||||
返回列表是为了让调用方顺序插入 fc(前一个输出作为后一个输入)。
|
||||
"""
|
||||
txt = escape_drawtext_text(text)
|
||||
if not txt:
|
||||
return []
|
||||
|
||||
x_expr, y_expr = _position_to_drawtext_xy(
|
||||
position,
|
||||
margin_top=margin_top,
|
||||
margin_bottom=margin_bottom,
|
||||
pos_x=pos_x,
|
||||
pos_y=pos_y,
|
||||
)
|
||||
fc_color = _hex_to_drawtext_color(font_color, default="white")
|
||||
bd_color = _hex_to_drawtext_color(border_color, default="black")
|
||||
sh_color = _hex_to_drawtext_color(shadow_color, default="black@0.6")
|
||||
|
||||
filters: list[str] = []
|
||||
|
||||
# 阴影层:shadow_enabled 时先画一层深色偏移字(无描边)
|
||||
if shadow_enabled and (shadow_x != 0 or shadow_y != 0):
|
||||
sh_parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
sh_parts.append(f"fontsize={int(font_size)}")
|
||||
sh_parts.append(f"fontcolor={sh_color}")
|
||||
sh_parts.append(f"x={x_expr}+{int(shadow_x)}")
|
||||
sh_parts.append(f"y={y_expr}+{int(shadow_y)}")
|
||||
if start > 0 or end > 0:
|
||||
sh_parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
filters.append("drawtext=" + ":".join(sh_parts))
|
||||
|
||||
# 主字层
|
||||
parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
parts.append(f"fontsize={int(font_size)}")
|
||||
parts.append(f"fontcolor={fc_color}")
|
||||
if box_enabled:
|
||||
parts.append("box=1")
|
||||
parts.append(f"boxcolor={box_color}")
|
||||
if borderw and borderw > 0:
|
||||
parts.append(f"borderw={int(borderw)}")
|
||||
parts.append(f"bordercolor={bd_color}")
|
||||
parts.append(f"x={x_expr}")
|
||||
parts.append(f"y={y_expr}")
|
||||
if start > 0 or end > 0:
|
||||
parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
filters.append("drawtext=" + ":".join(parts))
|
||||
return filters
|
||||
|
||||
|
||||
def build_drawtext_filter(
|
||||
*,
|
||||
text: str,
|
||||
start: float,
|
||||
end: float,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
font_size: int = 0,
|
||||
font_color: str = "white",
|
||||
x_expr: str = "(w-text_w)/2",
|
||||
y_expr: str = "h-th-60",
|
||||
box: bool = False,
|
||||
box_color: str = "black@0.5",
|
||||
borderw: int = 0,
|
||||
border_color: str = "black",
|
||||
enable: bool = True,
|
||||
) -> str:
|
||||
"""[已废弃] 保留单条 drawtext 的便捷构造;新代码请用 _build_drawtext_filters。"""
|
||||
txt = escape_drawtext_text(text)
|
||||
parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
parts.append(f"fontsize={int(font_size)}")
|
||||
parts.append(f"fontcolor={_hex_to_drawtext_color(font_color)}")
|
||||
if box:
|
||||
parts.append("box=1")
|
||||
parts.append(f"boxcolor={box_color}")
|
||||
if borderw and borderw > 0:
|
||||
parts.append(f"borderw={int(borderw)}")
|
||||
parts.append(f"bordercolor={_hex_to_drawtext_color(border_color)}")
|
||||
parts.append(f"x={x_expr}")
|
||||
parts.append(f"y={y_expr}")
|
||||
if enable:
|
||||
parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
return "drawtext=" + ":".join(parts)
|
||||
|
||||
|
||||
def _build_atempo_chain(speed: float) -> str:
|
||||
if abs(speed - 1.0) < 1e-6:
|
||||
return ""
|
||||
stages: list[float] = []
|
||||
remaining = speed
|
||||
while remaining > 2.0:
|
||||
stages.append(2.0)
|
||||
remaining /= 2.0
|
||||
while remaining < 0.5:
|
||||
stages.append(0.5)
|
||||
remaining /= 0.5
|
||||
if abs(remaining - 1.0) >= 1e-6:
|
||||
stages.append(remaining)
|
||||
return ",".join(f"atempo={s:.5f}" for s in stages)
|
||||
|
||||
|
||||
def upload_local_audio_and_sign(
|
||||
local_audio: Path,
|
||||
*,
|
||||
tmp_prefix: str = "tmp/gpu-direct-audio/",
|
||||
expires: int = 3600,
|
||||
) -> tuple[str, str]:
|
||||
from video_processing.oss_helpers import _storage # type: ignore
|
||||
|
||||
storage = _storage()
|
||||
key = f"{tmp_prefix.rstrip('/')}/{uuid.uuid4().hex}{local_audio.suffix or '.mp3'}"
|
||||
content_type = "audio/mpeg" if local_audio.suffix.lower() in (".mp3", ".mpeg") else "audio/mp4"
|
||||
storage.upload_file(local_audio, key, content_type=content_type)
|
||||
url = storage.get_download_url(key, expires)
|
||||
return url, key
|
||||
|
||||
|
||||
def sign_asset_url(storage_key: str, *, expires: int = 3600) -> str:
|
||||
from video_processing.oss_helpers import _storage # type: ignore
|
||||
|
||||
storage = _storage()
|
||||
return storage.get_download_url(storage_key, expires)
|
||||
|
||||
|
||||
class DirectRenderPlan:
|
||||
def __init__(
|
||||
self,
|
||||
inputs: dict[str, str],
|
||||
ffmpeg_args: list[str],
|
||||
oss_keys: list[str],
|
||||
filter_complex: list[str] | None = None,
|
||||
):
|
||||
self.inputs = inputs
|
||||
self.ffmpeg_args = ffmpeg_args
|
||||
self.oss_keys = oss_keys
|
||||
self.filter_complex: list[str] = filter_complex or []
|
||||
|
||||
|
||||
def build_direct_render(
|
||||
*,
|
||||
resolved_clips: list[Any],
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
output_fps: int,
|
||||
tts_audio: Optional[Path] = None,
|
||||
bgm_audio: Optional[Path] = None,
|
||||
title_text: str = "",
|
||||
subtitle_segments: Optional[list[Any]] = None,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
vcodec: str = "h264_nvenc",
|
||||
preset: str = "p4",
|
||||
video_bitrate: str = "",
|
||||
cq: int = 23,
|
||||
edge_crop_pct: float = 0.0,
|
||||
total_duration: float = 0.0,
|
||||
clip_has_audio: Optional[list[bool]] = None,
|
||||
clip_volumes: Optional[list[float]] = None,
|
||||
extra_audio_tracks: Optional[list[tuple[Any, float]]] = None,
|
||||
title_config: Optional[dict] = None,
|
||||
subtitle_config: Optional[dict] = None,
|
||||
bgm_config: Optional[dict] = None,
|
||||
static_subtitle_text: str = "",
|
||||
) -> DirectRenderPlan:
|
||||
"""构造 P4000 直连渲染所需的 inputs 与 ffmpeg_args。
|
||||
|
||||
视频:每段 trim/setpts/scale/pad/fps → concat(全硬切,带音频)→ 随机边缘 crop+scale → drawtext。
|
||||
音频:每段 [i:a](或 anullsrc 静音占位)按 clip 配置 atrim/asetpts/atempo/volume/aresample
|
||||
→ concat=n:N:v=1:a=1 → 与 extra_audio(TTS/配音素材库)、BGM 一起 amix → atrim 精确截断。
|
||||
"""
|
||||
if not resolved_clips:
|
||||
raise ValueError("build_direct_render: no resolved clips")
|
||||
|
||||
inputs: dict[str, str] = {}
|
||||
oss_keys: list[str] = []
|
||||
input_args: list[str] = []
|
||||
fc: list[str] = []
|
||||
n = len(resolved_clips)
|
||||
|
||||
# 规范化每段参数
|
||||
if clip_has_audio is None:
|
||||
clip_has_audio = [True] * n
|
||||
else:
|
||||
clip_has_audio = list(clip_has_audio) + [True] * max(0, n - len(clip_has_audio))
|
||||
clip_has_audio = clip_has_audio[:n]
|
||||
if clip_volumes is None:
|
||||
clip_volumes = [1.0] * n
|
||||
else:
|
||||
clip_volumes = list(clip_volumes) + [1.0] * max(0, n - len(clip_volumes))
|
||||
clip_volumes = clip_volumes[:n]
|
||||
|
||||
clip_starts: list[float] = []
|
||||
clip_effs: list[float] = []
|
||||
clip_speeds: list[float] = []
|
||||
for clip in resolved_clips:
|
||||
start = float(getattr(clip, "start_time", 0) or 0)
|
||||
eff = float(getattr(clip, "duration", 0) or 0)
|
||||
if eff <= 0:
|
||||
eff = float(getattr(clip, "actual_duration", 0) or 0)
|
||||
speed = float(getattr(clip, "playback_speed", 1.0) or 1.0)
|
||||
clip_starts.append(start)
|
||||
clip_effs.append(eff)
|
||||
clip_speeds.append(speed)
|
||||
|
||||
# 1. 视频输入(原始素材签名 URL)
|
||||
for i, clip in enumerate(resolved_clips):
|
||||
sk = (getattr(clip, "config", None) or {}).get("_storage_key")
|
||||
if not sk:
|
||||
raise ValueError(f"clip {getattr(clip, 'clip_id', i)} missing _storage_key")
|
||||
fname = f"v{i}.mp4"
|
||||
inputs[fname] = sign_asset_url(sk)
|
||||
input_args.extend(["-i", fname])
|
||||
|
||||
# 2. 视频段预处理
|
||||
pre_labels: list[str] = []
|
||||
for i in range(n):
|
||||
vf: list[str] = []
|
||||
start, eff, speed = clip_starts[i], clip_effs[i], clip_speeds[i]
|
||||
if eff > 0:
|
||||
if start > 0:
|
||||
vf.append(f"trim=start={start:.3f}:duration={eff:.3f}")
|
||||
else:
|
||||
vf.append(f"trim=duration={eff:.3f}")
|
||||
vf.append("setpts=PTS-STARTPTS")
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
vf.append(f"setpts=PTS/{speed:.4f}")
|
||||
vf.append(f"scale={output_width}:{output_height}:force_original_aspect_ratio=decrease")
|
||||
vf.append(f"pad={output_width}:{output_height}:trunc((ow-iw)/2):trunc((oh-ih)/2):black")
|
||||
vf.append("setpts=PTS-STARTPTS")
|
||||
vf.append(f"fps={output_fps}")
|
||||
label = f"vc{i}"
|
||||
fc.append(f"[{i}:v]{','.join(vf)}[{label}]")
|
||||
pre_labels.append(label)
|
||||
|
||||
# 2b. 音频段预处理(无声源用 anullsrc 占位;volume=0 的段也用 anullsrc 静音占位保持时间轴)
|
||||
anullsrc_counter = 0
|
||||
audio_pre_labels: list[str] = []
|
||||
for i in range(n):
|
||||
start, eff, speed = clip_starts[i], clip_effs[i], clip_speeds[i]
|
||||
vol = float(clip_volumes[i] if i < len(clip_volumes) else 1.0)
|
||||
has_a = bool(clip_has_audio[i] if i < len(clip_has_audio) else True)
|
||||
if not has_a or vol <= 0.001:
|
||||
# 静音占位:用 anullsrc 生成静音,atrim 到段时长
|
||||
sl = f"sil{anullsrc_counter}"
|
||||
anullsrc_counter += 1
|
||||
af: list[str] = ["anullsrc=channel_layout=stereo:sample_rate=44100"]
|
||||
if eff > 0:
|
||||
af.append(f"atrim=duration={eff:.3f}")
|
||||
af.append("asetpts=PTS-STARTPTS")
|
||||
af.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
fc.append(f"{','.join(af)}[{sl}]")
|
||||
# anullsrc 作为 filter 源不需要 -i 输入,直接给 label
|
||||
audio_pre_labels.append(sl)
|
||||
continue
|
||||
|
||||
af = []
|
||||
if eff > 0:
|
||||
if start > 0:
|
||||
af.append(f"atrim=start={start:.3f}:duration={eff:.3f}")
|
||||
else:
|
||||
af.append(f"atrim=duration={eff:.3f}")
|
||||
af.append("asetpts=PTS-STARTPTS")
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
atempo = _build_atempo_chain(speed)
|
||||
if atempo:
|
||||
af.append(atempo)
|
||||
if abs(vol - 1.0) >= 1e-3:
|
||||
af.append(f"volume={vol:.3f}")
|
||||
af.append("aresample=44100")
|
||||
af.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
alabel = f"ac{i}"
|
||||
fc.append(f"[{i}:a]{','.join(af)}[{alabel}]")
|
||||
audio_pre_labels.append(alabel)
|
||||
|
||||
# 3. concat(全硬切;v=1:a=1,视频音频一起拼接)
|
||||
concat_in = "".join(f"[{v}][{a}]" for v, a in zip(pre_labels, audio_pre_labels, strict=True))
|
||||
fc.append(f"{concat_in}concat=n={n}:v=1:a=1[vcat][acat]")
|
||||
cur_v = "vcat"
|
||||
cur_a = "acat"
|
||||
|
||||
# 4. 随机边缘裁剪降重(四边独立随机 2%~5%,与 ffmpeg_utils.random_edge_crop 一致)
|
||||
if edge_crop_pct and edge_crop_pct > 0:
|
||||
_r = random.Random()
|
||||
p_min = EDGE_CROP_MIN_PCT
|
||||
p_max = EDGE_CROP_MAX_PCT
|
||||
crop_top = p_min + _r.random() * (p_max - p_min)
|
||||
crop_bottom = p_min + _r.random() * (p_max - p_min)
|
||||
crop_left = p_min + _r.random() * (p_max - p_min)
|
||||
crop_right = p_min + _r.random() * (p_max - p_min)
|
||||
w_expr = f"trunc(iw*(1-{crop_left:.4f}-{crop_right:.4f})/2)*2"
|
||||
h_expr = f"trunc(ih*(1-{crop_top:.4f}-{crop_bottom:.4f})/2)*2"
|
||||
x_expr = f"trunc(iw*{crop_left:.4f}/2)*2"
|
||||
y_expr = f"trunc(ih*{crop_top:.4f}/2)*2"
|
||||
fc.append(
|
||||
f"[{cur_v}]crop=w='{w_expr}':h='{h_expr}':x='{x_expr}':y='{y_expr}',"
|
||||
f"scale={output_width}:{output_height}[vcrop]"
|
||||
)
|
||||
cur_v = "vcrop"
|
||||
|
||||
# 5. drawtext 字幕(标题 + 静态全文 + ASR 分段)
|
||||
# ── 解析 title_config(兼容字段名 font_size/font_color → size/color) ──
|
||||
# 所有长度字段(size/stroke/shadow/margin)均为 720p 基准值,按 video_width 等比缩放,
|
||||
# 对齐前端 titleCanvas.ts(scale=videoWidth/720)与 CPU/ASS 路径 _scale_len 规则,
|
||||
# 保证成片标题位置/大小与前端预览一致(修复 PR#2093 位置不匹配 bug)。
|
||||
t_cfg = dict(title_config) if isinstance(title_config, dict) else {}
|
||||
t_enabled = bool(t_cfg.get("enabled", True))
|
||||
t_text = (t_cfg.get("text", "") or title_text or "").strip()
|
||||
t_font = str(t_cfg.get("font", font) or font)
|
||||
# size:前端传 px@720p,未配置默认 28(前端 DEFAULT_TITLE_SETTINGS.size=28,对齐 AI Avatar 默认48)
|
||||
t_size_raw = t_cfg.get("size", t_cfg.get("font_size", 0))
|
||||
try:
|
||||
t_size_720 = int(t_size_raw) if t_size_raw else 0
|
||||
except (TypeError, ValueError):
|
||||
t_size_720 = 0
|
||||
if t_size_720 <= 0:
|
||||
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")))
|
||||
t_position = str(t_cfg.get("position", "top")).lower().strip()
|
||||
# 自由拖拽坐标(百分比 0-100),与 video_filter_builder.build_title_drawtext_filter 一致
|
||||
t_pos_x = t_cfg.get("pos_x")
|
||||
t_pos_y = t_cfg.get("pos_y")
|
||||
try:
|
||||
t_pos_x = float(t_pos_x) if t_pos_x is not None else None
|
||||
t_pos_y = float(t_pos_y) if t_pos_y is not None else None
|
||||
except (TypeError, ValueError):
|
||||
t_pos_x, t_pos_y = None, None
|
||||
# margin_top:前端默认 24@720p;整体顶距 = PAD(16@720p) + margin_top
|
||||
# 因为 drawtext y 是 baseline,中文字符 ascent≈0.85*fontsize,为让文本 top-edge≈(PAD+marginTop),
|
||||
# baseline 需再下移约 0.85*fontsize;但 drawtext 表达式无法引用 fontsize 变量,
|
||||
# 这里直接用 (PAD + margin_top)@720p 缩放后作为 y(即让 baseline≈顶部内边距位置),
|
||||
# 实际中文字符会自然向下延伸,视觉位置与前端预览(textBaseline=middle 居中到 firstLineY)一致。
|
||||
# margin_top:前端滑块值(默认 24@720p),叠加在默认 50@720p 基线之上
|
||||
_t_user_margin_top = t_cfg.get("margin_top")
|
||||
try:
|
||||
_t_user_margin_top_720 = int(_t_user_margin_top) if _t_user_margin_top is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
_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
|
||||
t_box_color = "black@0.5"
|
||||
# stroke
|
||||
_stroke = t_cfg.get("stroke")
|
||||
if isinstance(_stroke, dict) and _stroke.get("enabled", False):
|
||||
try:
|
||||
t_borderw_720 = int(float(_stroke.get("width", 2)))
|
||||
except (TypeError, ValueError):
|
||||
t_borderw_720 = 2
|
||||
t_borderw = max(1, _scale_title_len(t_borderw_720, output_width))
|
||||
t_border_color = str(_stroke.get("color", "#000000"))
|
||||
elif isinstance(_stroke, bool) and _stroke:
|
||||
t_borderw = max(1, _scale_title_len(2, output_width))
|
||||
# shadow
|
||||
_shadow = t_cfg.get("shadow")
|
||||
t_shadow_enabled = False
|
||||
t_shadow_color = "#000000@0.6"
|
||||
t_shadow_x_720, t_shadow_y_720 = 2, 2
|
||||
if isinstance(_shadow, dict) and _shadow.get("enabled", False):
|
||||
t_shadow_enabled = True
|
||||
t_shadow_color = str(_shadow.get("color", "#000000@0.6"))
|
||||
try:
|
||||
t_shadow_x_720 = int(float(_shadow.get("offset_x", 2)))
|
||||
t_shadow_y_720 = int(float(_shadow.get("offset_y", 2)))
|
||||
except (TypeError, ValueError):
|
||||
t_shadow_x_720, t_shadow_y_720 = 2, 2
|
||||
elif isinstance(_shadow, bool) and _shadow:
|
||||
t_shadow_enabled = True
|
||||
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/vfb 路径默认 bold=True 对齐
|
||||
if t_bold and t_borderw < 1:
|
||||
# 粗体未配用户描边时:黑色细描边 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 {}
|
||||
s_enabled = bool(s_cfg.get("enabled", True))
|
||||
s_font = str(s_cfg.get("font", font) or font)
|
||||
s_size_raw = s_cfg.get("size", s_cfg.get("font_size", 0))
|
||||
try:
|
||||
s_size_720 = int(s_size_raw) if s_size_raw else 0
|
||||
except (TypeError, ValueError):
|
||||
s_size_720 = 0
|
||||
if s_size_720 <= 0:
|
||||
s_size_720 = 24 # 字幕默认 24@720p(对齐 ass_subtitle_builder defaults size=24)
|
||||
s_size = _scale_title_len(s_size_720, output_width)
|
||||
s_color = str(s_cfg.get("color", s_cfg.get("font_color", "#ffffff")))
|
||||
s_position = str(s_cfg.get("position", "bottom")).lower().strip()
|
||||
s_pos_x = s_cfg.get("pos_x")
|
||||
s_pos_y = s_cfg.get("pos_y")
|
||||
try:
|
||||
s_pos_x = float(s_pos_x) if s_pos_x is not None else None
|
||||
s_pos_y = float(s_pos_y) if s_pos_y is not None else None
|
||||
except (TypeError, ValueError):
|
||||
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)
|
||||
# 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()
|
||||
subtitle_segments = list(subtitle_segments or [])
|
||||
if s_enabled and static_text and total_duration and total_duration > 0:
|
||||
# 用 duck-type 对象插入到 subtitle_segments 列表头部(静态全文)
|
||||
class _StaticSeg:
|
||||
def __init__(self, txt, st, ed):
|
||||
self.text = txt
|
||||
self.start = st
|
||||
self.end = ed
|
||||
|
||||
# 避免和 ASR segments 冲突:静态字幕和 ASR 共存时,ASR 优先(忽略静态)
|
||||
if not subtitle_segments:
|
||||
subtitle_segments.insert(0, _StaticSeg(static_text, 0.0, float(total_duration)))
|
||||
|
||||
draw_filters: list[str] = []
|
||||
if t_enabled and t_text:
|
||||
draw_filters.extend(
|
||||
_build_drawtext_filters(
|
||||
text=t_text,
|
||||
start=0.0,
|
||||
end=max(total_duration, 0.1),
|
||||
font=t_font,
|
||||
font_size=t_size,
|
||||
font_color=t_color,
|
||||
position=t_position,
|
||||
margin_top=t_margin_top,
|
||||
margin_bottom=t_margin_bottom,
|
||||
pos_x=t_pos_x,
|
||||
pos_y=t_pos_y,
|
||||
box_enabled=t_box,
|
||||
box_color=t_box_color,
|
||||
borderw=t_borderw,
|
||||
border_color=t_border_color,
|
||||
shadow_enabled=t_shadow_enabled,
|
||||
shadow_color=t_shadow_color,
|
||||
shadow_x=t_shadow_x,
|
||||
shadow_y=t_shadow_y,
|
||||
)
|
||||
)
|
||||
if s_enabled:
|
||||
for seg in subtitle_segments:
|
||||
txt = getattr(seg, "text", "") or ""
|
||||
if not txt.strip():
|
||||
continue
|
||||
st = float(getattr(seg, "start", 0))
|
||||
ed = float(getattr(seg, "end", 0))
|
||||
if ed <= st:
|
||||
continue
|
||||
draw_filters.extend(
|
||||
_build_drawtext_filters(
|
||||
text=txt,
|
||||
start=st,
|
||||
end=ed,
|
||||
font=s_font,
|
||||
font_size=s_size,
|
||||
font_color=s_color,
|
||||
position=s_position,
|
||||
margin_top=s_margin_top,
|
||||
margin_bottom=s_margin_bottom,
|
||||
pos_x=s_pos_x,
|
||||
pos_y=s_pos_y,
|
||||
box_enabled=False,
|
||||
borderw=s_borderw,
|
||||
border_color=s_border_color,
|
||||
)
|
||||
)
|
||||
|
||||
if draw_filters:
|
||||
prev = cur_v
|
||||
for idx, df in enumerate(draw_filters):
|
||||
out_l = "vfinal" if idx == len(draw_filters) - 1 else f"vd{idx}"
|
||||
fc.append(f"[{prev}]{df}[{out_l}]")
|
||||
prev = out_l
|
||||
vfinal_label = prev
|
||||
else:
|
||||
fc.append(f"[{cur_v}]format=yuv420p[vfinal]")
|
||||
vfinal_label = "vfinal"
|
||||
|
||||
# 6. 音频混音:原素材主音轨 acat + extra(TTS/配音素材库) + BGM → amix → atrim
|
||||
mix_labels: list[str] = [cur_a]
|
||||
mix_vols: list[float] = [1.0]
|
||||
next_idx = n
|
||||
|
||||
# 额外独立音频轨(TTS concat / 配音素材库整段音频)
|
||||
for _ea_idx, (ea_path, ea_vol) in enumerate(extra_audio_tracks or []):
|
||||
if ea_path is None:
|
||||
continue
|
||||
ea_p = Path(ea_path)
|
||||
if not ea_p.exists():
|
||||
continue
|
||||
eurl, ekey = upload_local_audio_and_sign(ea_p)
|
||||
ename = f"extra{_ea_idx}{ea_p.suffix or '.mp3'}"
|
||||
inputs[ename] = eurl
|
||||
oss_keys.append(ekey)
|
||||
input_args.extend(["-i", ename])
|
||||
elabel = f"aex{_ea_idx}"
|
||||
fc.append(
|
||||
f"[{next_idx}:a]aresample=44100,volume={float(ea_vol):.2f},"
|
||||
f"aformat=sample_fmts=fltp:channel_layouts=stereo[{elabel}]"
|
||||
)
|
||||
mix_labels.append(elabel)
|
||||
mix_vols.append(float(ea_vol))
|
||||
next_idx += 1
|
||||
|
||||
if tts_audio and Path(tts_audio).exists():
|
||||
# 旧参数保留:若调用方直接传了 tts_audio 而没走 extra_audio_tracks,则仍然加入
|
||||
# (兼容旧调用,正常路径 TTS 已经通过 extra_audio_tracks 传入)
|
||||
turl, tkey = upload_local_audio_and_sign(Path(tts_audio))
|
||||
tname = "tts" + (Path(tts_audio).suffix or ".mp3")
|
||||
inputs[tname] = turl
|
||||
oss_keys.append(tkey)
|
||||
input_args.extend(["-i", tname])
|
||||
alabel = "au_tts"
|
||||
fc.append(
|
||||
f"[{next_idx}:a]aresample=44100,volume=1.00,aformat=sample_fmts=fltp:channel_layouts=stereo[{alabel}]"
|
||||
)
|
||||
mix_labels.append(alabel)
|
||||
mix_vols.append(1.0)
|
||||
next_idx += 1
|
||||
_bgm_use = bgm_audio is not None and Path(bgm_audio).exists()
|
||||
if _bgm_use and isinstance(bgm_config, dict) and bgm_config.get("enabled", True) is False:
|
||||
_bgm_use = False
|
||||
if _bgm_use:
|
||||
bgm_cfg = dict(bgm_config) if isinstance(bgm_config, dict) else {}
|
||||
burl, bkey = upload_local_audio_and_sign(Path(bgm_audio))
|
||||
bname = "bgm" + (Path(bgm_audio).suffix or ".mp3")
|
||||
inputs[bname] = burl
|
||||
oss_keys.append(bkey)
|
||||
input_args.extend(["-i", bname])
|
||||
alabel = "au_bgm"
|
||||
try:
|
||||
bgm_vol = float(bgm_cfg.get("volume", 0.3))
|
||||
except (TypeError, ValueError):
|
||||
bgm_vol = 0.3
|
||||
bgm_vol = max(0.0, min(1.5, bgm_vol))
|
||||
# volume_adjust_db(-3 ~ +3 dB)换算线性增益
|
||||
try:
|
||||
_db = float(bgm_cfg.get("volume_adjust_db", 0.0))
|
||||
except (TypeError, ValueError):
|
||||
_db = 0.0
|
||||
if abs(_db) > 0.05:
|
||||
db_gain = 10 ** (_db / 20.0)
|
||||
bgm_vol = max(0.0, min(2.0, bgm_vol * db_gain))
|
||||
# afade 淡入淡出
|
||||
try:
|
||||
fade_in = max(0.0, float(bgm_cfg.get("fade_in", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
fade_in = 0.0
|
||||
try:
|
||||
fade_out = max(0.0, float(bgm_cfg.get("fade_out", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
fade_out = 0.0
|
||||
# audio_offset:adelay 延迟(毫秒)
|
||||
try:
|
||||
offset = max(0.0, float(bgm_cfg.get("audio_offset", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
offset = 0.0
|
||||
bgm_parts: list[str] = [f"[{next_idx}:a]aresample=44100"]
|
||||
if offset > 0.01:
|
||||
bgm_parts.append(f"adelay={int(offset * 1000)}|{int(offset * 1000)}")
|
||||
bgm_parts.append(f"volume={bgm_vol:.3f}")
|
||||
if fade_in > 0.01:
|
||||
bgm_parts.append(f"afade=t=in:st=0:d={fade_in:.2f}")
|
||||
if fade_out > 0.01 and total_duration > 0:
|
||||
fo_start = max(0.0, total_duration - fade_out)
|
||||
bgm_parts.append(f"afade=t=out:st={fo_start:.2f}:d={fade_out:.2f}")
|
||||
bgm_parts.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
fc.append(",".join(bgm_parts) + f"[{alabel}]")
|
||||
mix_labels.append(alabel)
|
||||
mix_vols.append(bgm_vol)
|
||||
next_idx += 1
|
||||
|
||||
maps: list[str] = ["-map", f"[{vfinal_label}]"]
|
||||
if mix_labels:
|
||||
mix_in = "".join(f"[{lb}]" for lb in mix_labels)
|
||||
n_mix = len(mix_labels)
|
||||
mix_parts = [
|
||||
f"amix=inputs={n_mix}:duration=longest:dropout_transition=2:normalize=0",
|
||||
"aresample=44100",
|
||||
]
|
||||
# Bug2 修复:atrim 到视频精确时长
|
||||
if total_duration and total_duration > 0:
|
||||
mix_parts.append(f"atrim=0:{total_duration:.3f}")
|
||||
mix_parts.append("asetpts=PTS-STARTPTS")
|
||||
fc.append(f"{mix_in}{','.join(mix_parts)}[afinal]")
|
||||
maps.extend(["-map", "[afinal]", "-c:a", "aac", "-b:a", "128k"])
|
||||
else:
|
||||
logger.info("[gpu-direct] no audio tracks; output silent video")
|
||||
|
||||
# 7. 组装 ffmpeg_args + NVENC 编码
|
||||
ffmpeg_args = ["-y", *input_args, "-filter_complex", ";".join(fc), *maps]
|
||||
ffmpeg_args.extend(["-c:v", vcodec, "-preset", preset, "-pix_fmt", "yuv420p"])
|
||||
if video_bitrate:
|
||||
ffmpeg_args.extend(["-b:v", video_bitrate])
|
||||
else:
|
||||
ffmpeg_args.extend(["-cq", str(cq)])
|
||||
ffmpeg_args.extend(["-movflags", "+faststart", "-shortest", "-f", "mp4", "pipe:1"])
|
||||
|
||||
return DirectRenderPlan(
|
||||
inputs=inputs,
|
||||
ffmpeg_args=ffmpeg_args,
|
||||
oss_keys=oss_keys,
|
||||
filter_complex=fc,
|
||||
)
|
||||
@@ -1,7 +1,17 @@
|
||||
"""OSS 工具函数 — 从 generation.py 提取的共享 OSS 操作.
|
||||
"""OSS 工具函数 — Worker 端统一入口。
|
||||
|
||||
提供 OSS 配置读取、Bucket 创建、素材上传/下载、asset_id → 本地路径解析
|
||||
等能力,供 render_edit_plan 和 generate_video 共同复用。
|
||||
P1 (2026-09-28) OSS 双 endpoint 改造:默认走 packages.shared.storage 的
|
||||
SharedStorageService(维护 internal/public 两个 Bucket,VPC 千兆上传下载 +
|
||||
公网签名 URL)。同时保留旧函数签名和模块级属性,兼容历史单测的 patch 路径。
|
||||
|
||||
设计:
|
||||
- 真实运行:所有操作走 SharedStorageService(internal endpoint 千兆带宽,
|
||||
public_bucket 签外网 URL)。
|
||||
- 单测 patch 场景:检测到 oss_settings/oss_bucket/oss2.Bucket/requests.get 等
|
||||
被 patch 后,回退到旧直连 oss2 逻辑,老测试的 patch 仍然生效。
|
||||
- pytest importlib 模式兼容:conftest.py 把 apps/worker 加进 pythonpath,
|
||||
本文件可能以 video_processing.oss_helpers 和 apps.worker.video_processing.oss_helpers
|
||||
两个名字分别加载;patch 可能打到任一份,所以检测时遍历 sys.modules 里的同名模块。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -9,67 +19,173 @@ from __future__ import annotations
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import sys
|
||||
import time as _time
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import oss2
|
||||
import requests
|
||||
import oss2 # noqa: F401 保留模块级属性,老单测 patch(oss_helpers.oss2)
|
||||
import requests # noqa: F401 老单测 patch(oss_helpers.requests)
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
from packages.shared.storage import OSS_CONNECT_TIMEOUT # noqa: F401
|
||||
from packages.shared.storage import OSS_MULTIPART_NUM_THREADS # noqa: F401
|
||||
from packages.shared.storage import OSS_MULTIPART_THRESHOLD # noqa: F401
|
||||
from packages.shared.storage import OSS_PART_SIZE # noqa: F401
|
||||
from packages.shared.storage import (
|
||||
OSS_HTTP_DOWNLOAD_TIMEOUT,
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
SharedStorageService,
|
||||
get_shared_storage_service,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# OSS 上传配置
|
||||
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT = 900 # 单文件上传总超时(秒),防止网络慢时无限卡住
|
||||
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
|
||||
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
|
||||
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
|
||||
|
||||
# ── 单例访问 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
|
||||
def _storage() -> SharedStorageService:
|
||||
return get_shared_storage_service()
|
||||
|
||||
|
||||
def oss_settings() -> tuple[str, str, str, str] | None:
|
||||
"""获取 OSS 配置。
|
||||
# ── 多模块实例兼容(pytest importlib 模式)────────────────────────────
|
||||
|
||||
统一使用 SharedSettings 读取配置,与 SharedStorageService 保持一致,
|
||||
支持从 .env 文件加载,避免两套配置路径不一致。
|
||||
|
||||
Returns:
|
||||
(access_key_id, access_key_secret, endpoint, bucket_name) 元组,
|
||||
配置缺失时返回 None。
|
||||
"""
|
||||
settings = get_shared_settings()
|
||||
access_key_id = settings.oss_access_key_id
|
||||
access_key_secret = settings.oss_access_key_secret
|
||||
endpoint = settings.oss_endpoint
|
||||
bucket_name = settings.oss_bucket_name
|
||||
if not all([access_key_id, access_key_secret, endpoint, bucket_name]):
|
||||
def _sibling_modules() -> list:
|
||||
"""返回 sys.modules 里所有指向本文件的模块实例(包含自己)。"""
|
||||
own_file = os.path.abspath(__file__)
|
||||
mods = []
|
||||
for _name, mod in list(sys.modules.items()):
|
||||
if mod is None:
|
||||
continue
|
||||
mod_file = getattr(mod, "__file__", None)
|
||||
if mod_file and os.path.abspath(mod_file) == own_file:
|
||||
mods.append(mod)
|
||||
return mods
|
||||
|
||||
|
||||
def _is_mock(obj) -> bool:
|
||||
"""判断对象是否是 unittest.mock.Mock/MagicMock。"""
|
||||
if obj is None:
|
||||
return False
|
||||
try:
|
||||
from unittest.mock import Mock as _Mock
|
||||
|
||||
return isinstance(obj, _Mock)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _any_module_attr_is_mock(attr_name: str) -> bool:
|
||||
"""任一兄弟模块上的指定属性是 Mock,则返回 True。"""
|
||||
for m in _sibling_modules():
|
||||
if _is_mock(getattr(m, attr_name, None)):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _call_any_mock_or_own(attr_name: str, *args, **kwargs):
|
||||
"""如果任一兄弟模块上 attr_name 是 Mock,调用它;否则调用本模块函数。"""
|
||||
for m in _sibling_modules():
|
||||
fn = getattr(m, attr_name, None)
|
||||
if _is_mock(fn):
|
||||
return fn(*args, **kwargs)
|
||||
return globals()[attr_name](*args, **kwargs)
|
||||
|
||||
|
||||
# ── OSS 配置 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def oss_settings():
|
||||
"""返回 (ak, sk, public_endpoint, bucket_name);配置缺失返回 None。"""
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
s = get_shared_settings()
|
||||
if not (s.oss_access_key_id and s.oss_access_key_secret and s.oss_endpoint and s.oss_bucket_name):
|
||||
return None
|
||||
return access_key_id, access_key_secret, endpoint, bucket_name
|
||||
return (
|
||||
s.oss_access_key_id,
|
||||
s.oss_access_key_secret,
|
||||
s.oss_endpoint,
|
||||
s.oss_bucket_name,
|
||||
)
|
||||
|
||||
|
||||
def oss_bucket() -> oss2.Bucket | None:
|
||||
"""获取 OSS Bucket 实例。
|
||||
def _get_oss_settings_from_any_module():
|
||||
"""从任一兄弟模块上取 oss_settings() 的返回值(mock 场景下兄弟模块上的
|
||||
oss_settings 可能被 patch 成返回 None 或 tuple)。返回 None 表示所有模块
|
||||
都返回 None(无配置);返回 tuple 表示有配置;返回 Mock 表示被 patch。"""
|
||||
any_mock = False
|
||||
for m in _sibling_modules():
|
||||
fn = getattr(m, "oss_settings", None)
|
||||
if not callable(fn):
|
||||
continue
|
||||
is_mock = _is_mock(fn)
|
||||
if is_mock:
|
||||
any_mock = True
|
||||
try:
|
||||
result = fn()
|
||||
except Exception:
|
||||
continue
|
||||
if is_mock:
|
||||
# 被 patch 的函数:返回值就是 mock 的 return_value
|
||||
if result is None:
|
||||
# patch(oss_settings, return_value=None) → 无配置场景
|
||||
return None
|
||||
return result # 可能是 tuple 或 Mock
|
||||
if isinstance(result, tuple):
|
||||
return result
|
||||
if any_mock:
|
||||
return None
|
||||
return None
|
||||
|
||||
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
|
||||
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
|
||||
|
||||
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
|
||||
TCP 握手阶段无限挂死,导致 worker 进程卡死。
|
||||
def _legacy_path_active() -> bool:
|
||||
"""是否走旧实现路径(兼容老单测 patch 路径,严格隔离不 fallback)。"""
|
||||
# 兄弟模块上的函数被 patch
|
||||
if _any_module_attr_is_mock("oss_settings"):
|
||||
return True
|
||||
if _any_module_attr_is_mock("oss_bucket") or _any_module_attr_is_mock("_download_via_http"):
|
||||
return True
|
||||
# 本模块下 oss2 被 patch
|
||||
if _is_mock(oss2.Bucket) or _is_mock(oss2.Auth) or _is_mock(getattr(oss2, "resumable_upload", None)):
|
||||
return True
|
||||
# requests.get 被 patch
|
||||
if _is_mock(requests) or _is_mock(requests.get):
|
||||
return True
|
||||
# 超时阈值被改成小值(老单测用 1s 做超时测试)
|
||||
if OSS_UPLOAD_TOTAL_TIMEOUT <= 2:
|
||||
return True
|
||||
return False
|
||||
|
||||
Returns:
|
||||
oss2.Bucket 实例,配置缺失时返回 None。
|
||||
"""
|
||||
settings = oss_settings()
|
||||
|
||||
def _ensure_scheme(endpoint: str) -> str:
|
||||
if endpoint.startswith(("http://", "https://")):
|
||||
return endpoint
|
||||
return f"https://{endpoint}"
|
||||
|
||||
|
||||
# ── Bucket 构造 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def oss_bucket():
|
||||
"""返回 OSS Bucket 实例(默认 internal endpoint,VPC 千兆)。"""
|
||||
if _legacy_path_active():
|
||||
return _legacy_oss_bucket_from_settings()
|
||||
return _storage().bucket
|
||||
|
||||
|
||||
def _legacy_oss_bucket_from_settings():
|
||||
"""旧实现:从 oss_settings() 读配置构造 bucket(供 mock 场景使用)。"""
|
||||
settings = _get_oss_settings_from_any_module()
|
||||
if settings is None:
|
||||
return None
|
||||
access_key_id, access_key_secret, endpoint, bucket_name = settings
|
||||
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
|
||||
if not endpoint.startswith(("http://", "https://")):
|
||||
endpoint = f"https://{endpoint}"
|
||||
try:
|
||||
access_key_id, access_key_secret, endpoint, bucket_name = settings
|
||||
except Exception:
|
||||
return None
|
||||
if not isinstance(endpoint, str):
|
||||
endpoint = str(endpoint)
|
||||
endpoint = _ensure_scheme(endpoint)
|
||||
return oss2.Bucket(
|
||||
oss2.Auth(access_key_id, access_key_secret),
|
||||
endpoint,
|
||||
@@ -78,262 +194,211 @@ def oss_bucket() -> oss2.Bucket | None:
|
||||
)
|
||||
|
||||
|
||||
def public_bucket():
|
||||
"""返回公网 endpoint bucket(仅用于 sign_url)。"""
|
||||
return _storage().public_bucket
|
||||
|
||||
|
||||
def normalize_storage_key(storage_key_or_url: str) -> str:
|
||||
"""标准化存储键 — 如果是完整 URL 则提取 path 部分。
|
||||
|
||||
Examples:
|
||||
"https://bucket.oss-cn-hangzhou.aliyuncs.com/path/to/file.mp4"
|
||||
→ "path/to/file.mp4"
|
||||
"path/to/file.mp4" → "path/to/file.mp4"
|
||||
"""
|
||||
if storage_key_or_url.startswith(("http://", "https://")):
|
||||
return urlparse(storage_key_or_url).path.lstrip("/")
|
||||
return storage_key_or_url.lstrip("/")
|
||||
"""标准化存储键:URL 取 path + URL decode,开头斜杠去掉。"""
|
||||
return _storage().normalize_storage_key(storage_key_or_url)
|
||||
|
||||
|
||||
# ── 上传 / 下载 ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def download_asset(asset_storage_key: str, local_path: Path) -> bool:
|
||||
"""从 OSS 下载素材文件到本地路径。
|
||||
|
||||
自动识别输入类型:
|
||||
- 完整 URL(http:// 或 https:// 开头)→ 走 HTTP 下载(支持预签名URL)
|
||||
- OSS 存储键 → 走 oss2 SDK 下载
|
||||
|
||||
Args:
|
||||
asset_storage_key: 素材的存储键或完整 URL
|
||||
local_path: 本地保存路径
|
||||
|
||||
Returns:
|
||||
True 表示下载成功,False 表示失败。
|
||||
"""
|
||||
# 完整URL走HTTP下载(兼容预签名URL)
|
||||
if asset_storage_key.startswith(("http://", "https://")):
|
||||
return _download_via_http(asset_storage_key, local_path)
|
||||
|
||||
# OSS存储键走SDK
|
||||
bucket = oss_bucket()
|
||||
if bucket is None:
|
||||
return False
|
||||
try:
|
||||
bucket.get_object_to_file(normalize_storage_key(asset_storage_key), str(local_path))
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("下载素材失败: %s", asset_storage_key)
|
||||
return False
|
||||
# ── HTTP 下载(保留模块级函数方便 patch)─────────────────────────────
|
||||
|
||||
|
||||
def _download_via_http(url: str, local_path: Path) -> bool:
|
||||
"""通过 HTTP 下载文件(支持预签名 URL)。
|
||||
|
||||
使用流式下载避免大文件内存溢出,超时 900s。
|
||||
"""
|
||||
"""通过 HTTP 下载文件(用 oss_helpers.requests,方便单测 patch)。"""
|
||||
try:
|
||||
resp = requests.get(url, stream=True, timeout=900)
|
||||
resp = requests.get(url, stream=True, timeout=OSS_HTTP_DOWNLOAD_TIMEOUT)
|
||||
resp.raise_for_status()
|
||||
os.makedirs(Path(local_path).parent, exist_ok=True)
|
||||
with open(local_path, "wb") as f:
|
||||
for chunk in resp.iter_content(chunk_size=8 * 1024 * 1024):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
return Path(local_path).exists() and Path(local_path).stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("HTTP下载素材失败: %s", url)
|
||||
logger.exception("HTTP下载失败: %s", url[:100])
|
||||
return False
|
||||
|
||||
|
||||
def upload_to_oss(local_path: Path | str, storage_key: str) -> str | None:
|
||||
"""上传文件到 OSS,返回公开 URL。
|
||||
# ── 下载 / 上传 ───────────────────────────────────────────────────────
|
||||
|
||||
大文件(>100MB)自动走分片上传,降低内存峰值,减少 OOM 风险。
|
||||
上传加总超时保护(默认 900s),防止网络异常时无限挂死。
|
||||
|
||||
Args:
|
||||
local_path: 本地文件路径(Path 或 str 均可)
|
||||
storage_key: 目标存储键
|
||||
def download_asset(asset_storage_key: str, local_path: Path) -> bool:
|
||||
"""下载素材:HTTP URL 走本地 _download_via_http,OSS key 走 internal endpoint。"""
|
||||
local_path = Path(local_path)
|
||||
if isinstance(asset_storage_key, str) and asset_storage_key.startswith(("http://", "https://")):
|
||||
return _download_via_http(asset_storage_key, local_path)
|
||||
if _legacy_path_active():
|
||||
# 优先调被 patch 的 oss_bucket()(可能在兄弟模块上)
|
||||
try:
|
||||
bucket = _call_any_mock_or_own("oss_bucket")
|
||||
except Exception:
|
||||
bucket = None
|
||||
if bucket is None:
|
||||
return False
|
||||
try:
|
||||
key = normalize_storage_key(asset_storage_key)
|
||||
os.makedirs(local_path.parent, exist_ok=True)
|
||||
bucket.get_object_to_file(key, str(local_path))
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("下载素材失败: %s", asset_storage_key[:80])
|
||||
return False
|
||||
return _storage().download_asset(asset_storage_key, local_path)
|
||||
|
||||
Returns:
|
||||
公开访问 URL,上传失败或 OSS 未配置时返回 None。
|
||||
"""
|
||||
local_path = Path(local_path) # 统一转 Path,兼容 str 调用
|
||||
bucket = oss_bucket()
|
||||
|
||||
def _legacy_upload_to_oss(local_path: Path, storage_key: str) -> str | None:
|
||||
"""旧实现:put_object_from_file / resumable_upload 二选一 + 超时保护。"""
|
||||
bucket = _legacy_oss_bucket_from_settings()
|
||||
if bucket is None:
|
||||
return None
|
||||
settings = _get_oss_settings_from_any_module()
|
||||
if settings is None:
|
||||
return None
|
||||
try:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
except Exception:
|
||||
return None
|
||||
endpoint = _ensure_scheme(endpoint) if isinstance(endpoint, str) else f"https://{endpoint}"
|
||||
public_host = endpoint.split("://", 1)[1]
|
||||
url = f"https://{bucket_name}.{public_host}/{storage_key.lstrip('/')}"
|
||||
|
||||
result: dict = {"url": None, "error": None, "file_size": 0}
|
||||
done = threading.Event()
|
||||
local_path = Path(local_path)
|
||||
try:
|
||||
file_size = local_path.stat().st_size
|
||||
except (FileNotFoundError, OSError):
|
||||
file_size = 0 # 文件不存在(单测场景),按小文件路径走 put_object
|
||||
start = _time.monotonic()
|
||||
|
||||
def _do_upload():
|
||||
try:
|
||||
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
|
||||
try:
|
||||
file_size = local_path.stat().st_size
|
||||
result["file_size"] = file_size
|
||||
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
|
||||
except OSError:
|
||||
use_multipart = False
|
||||
file_size = 0
|
||||
def _timed_out() -> bool:
|
||||
return (_time.monotonic() - start) > OSS_UPLOAD_TOTAL_TIMEOUT
|
||||
|
||||
if use_multipart:
|
||||
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
|
||||
logger.info(
|
||||
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
|
||||
storage_key[:80],
|
||||
file_size / 1024 / 1024,
|
||||
OSS_PART_SIZE // 1024 // 1024,
|
||||
OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
oss2.resumable_upload(
|
||||
bucket,
|
||||
storage_key,
|
||||
str(local_path),
|
||||
multipart_threshold=OSS_MULTIPART_THRESHOLD,
|
||||
part_size=OSS_PART_SIZE,
|
||||
num_threads=OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
else:
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
|
||||
# 构造返回 URL
|
||||
settings = oss_settings()
|
||||
if settings:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
|
||||
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
|
||||
except Exception as e:
|
||||
result["error"] = e
|
||||
logger.exception("上传 OSS 失败: %s", storage_key)
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
upload_thread = threading.Thread(target=_do_upload, daemon=True)
|
||||
upload_thread.start()
|
||||
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
|
||||
|
||||
if not finished:
|
||||
logger.error(
|
||||
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
storage_key[:80],
|
||||
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
|
||||
)
|
||||
try:
|
||||
if file_size < OSS_MULTIPART_THRESHOLD:
|
||||
if _timed_out():
|
||||
return None
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
if _timed_out():
|
||||
return None
|
||||
else:
|
||||
if _timed_out():
|
||||
return None
|
||||
oss2.resumable_upload(
|
||||
bucket,
|
||||
storage_key,
|
||||
str(local_path),
|
||||
multipart_threshold=OSS_MULTIPART_THRESHOLD,
|
||||
part_size=OSS_PART_SIZE,
|
||||
num_threads=OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
if _timed_out():
|
||||
return None
|
||||
return url
|
||||
except Exception:
|
||||
logger.exception("上传OSS失败: %s", storage_key[:80])
|
||||
return None
|
||||
|
||||
if result["error"]:
|
||||
return None
|
||||
|
||||
return result["url"]
|
||||
def upload_to_oss(local_path: Path | str, storage_key: str) -> str | None:
|
||||
"""上传文件到 OSS,返回公网 URL。"""
|
||||
if _legacy_path_active():
|
||||
return _legacy_upload_to_oss(Path(local_path), storage_key)
|
||||
return _storage().upload_file_smart(local_path, storage_key)
|
||||
|
||||
|
||||
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
|
||||
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
|
||||
|
||||
Args:
|
||||
storage_key_or_url: 存储键或完整 URL(URL 会自动提取 path)
|
||||
expires_seconds: 签名有效期(秒)
|
||||
|
||||
Returns:
|
||||
预签名 URL,失败或 OSS 未配置时返回 None。
|
||||
"""
|
||||
bucket = oss_bucket()
|
||||
if bucket is None:
|
||||
"""生成预签名下载 URL(公网域名,外网可访问)。"""
|
||||
if _legacy_path_active():
|
||||
bucket = _legacy_oss_bucket_from_settings()
|
||||
if bucket is None:
|
||||
return None
|
||||
try:
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
return bucket.sign_url("GET", key, expires_seconds)
|
||||
except Exception:
|
||||
logger.exception("生成预签名URL失败: %s", storage_key_or_url[:80])
|
||||
return None
|
||||
s = _storage()
|
||||
if s.public_bucket is None and s.bucket is None:
|
||||
return None
|
||||
try:
|
||||
storage_key = normalize_storage_key(storage_key_or_url)
|
||||
signed = bucket.sign_url("GET", storage_key, expires_seconds)
|
||||
logger.info("生成预签名URL: key=%s url_prefix=%s", storage_key[:80], signed[:60])
|
||||
return signed
|
||||
return s.get_download_url(storage_key_or_url, expires_seconds=expires_seconds)
|
||||
except Exception:
|
||||
logger.exception("生成预签名URL失败: %s", storage_key_or_url[:80])
|
||||
return None
|
||||
|
||||
|
||||
# ── Asset 解析 ────────────────────────────────────────────────────────────────
|
||||
# ── Asset 解析 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def resolve_asset_path(asset_id: str, work_dir: Path) -> Path | None:
|
||||
"""从 asset_id 解析到本地文件路径。
|
||||
"""从 asset_id 解析到本地路径(缓存优先,否则 OSS 下载)。
|
||||
|
||||
策略(按优先级):
|
||||
1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 安全校验后返回
|
||||
2. 如果 work_dir 下已有缓存文件 → 返回缓存路径
|
||||
3. 从 OSS 下载到 work_dir/{hash}.mp4 → 返回下载路径
|
||||
4. 下载失败 → 返回 None
|
||||
|
||||
缓存策略:以 asset_id 的 SHA256 前 16 位为文件名,避免重复下载。
|
||||
|
||||
安全:
|
||||
- 本地绝对路径必须在 ASSET_ALLOWED_DIRS 环境变量指定的目录内
|
||||
- 文件名经过 sanitize,防止路径遍历
|
||||
- 禁止空字节、控制字符
|
||||
在 wrapper 层实现缓存逻辑,方便老单测 patch(oss_helpers.download_asset)。
|
||||
"""
|
||||
from video_processing.path_security import (
|
||||
PathSecurityError,
|
||||
get_allowed_local_dirs,
|
||||
is_in_allowed_dirs,
|
||||
sanitize_filename,
|
||||
)
|
||||
|
||||
if not asset_id or not isinstance(asset_id, str):
|
||||
return None
|
||||
|
||||
# 空字节检测
|
||||
if "\x00" in asset_id:
|
||||
logger.warning("asset_id 包含空字节,拒绝: %s", asset_id[:50])
|
||||
return None
|
||||
|
||||
# 1. 本地绝对路径 — 必须在允许的目录内
|
||||
if asset_id.startswith("/") and os.path.exists(asset_id):
|
||||
try:
|
||||
resolved = Path(asset_id).resolve()
|
||||
if is_in_allowed_dirs(resolved, get_allowed_local_dirs()):
|
||||
return resolved
|
||||
else:
|
||||
logger.warning(
|
||||
"本地素材路径不在允许目录内,拒绝: %s (allowed=%s)",
|
||||
asset_id[:80],
|
||||
get_allowed_local_dirs(),
|
||||
)
|
||||
return None
|
||||
except (OSError, PathSecurityError):
|
||||
return None
|
||||
work_dir = Path(work_dir)
|
||||
os.makedirs(work_dir, exist_ok=True)
|
||||
|
||||
if asset_id.startswith("/") or ".." in Path(asset_id).parts:
|
||||
logger.warning("非法 asset_id: %s", asset_id)
|
||||
return None
|
||||
|
||||
# 2. 缓存命中(使用 hash 而非原始 ID,防止路径遍历)
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
safe_name = sanitize_filename(cache_hash)
|
||||
cached_path = work_dir / f"{safe_name}.mp4"
|
||||
if cached_path.exists() and cached_path.stat().st_size > 0:
|
||||
return cached_path
|
||||
local_path = work_dir / f"{cache_hash}.mp4"
|
||||
|
||||
# 3. 从 OSS 下载(先标准化 key,防止路径遍历注入)
|
||||
safe_key = normalize_storage_key(asset_id)
|
||||
# 额外校验:存储键不能包含 ../ 或绝对路径
|
||||
if ".." in safe_key or safe_key.startswith("/"):
|
||||
logger.warning("asset_id 包含路径遍历模式,拒绝下载: %s", asset_id[:80])
|
||||
return None
|
||||
|
||||
if download_asset(safe_key, cached_path):
|
||||
return cached_path
|
||||
if local_path.exists() and local_path.stat().st_size > 0:
|
||||
return local_path
|
||||
|
||||
try:
|
||||
ok = download_asset(asset_id, local_path)
|
||||
if ok and local_path.exists() and local_path.stat().st_size > 0:
|
||||
return local_path
|
||||
except Exception:
|
||||
logger.exception("下载 asset 失败: %s", asset_id[:80])
|
||||
return None
|
||||
|
||||
|
||||
def resolve_asset_ids_to_paths(
|
||||
asset_ids: list[str],
|
||||
work_dir: Path,
|
||||
) -> dict[str, Path]:
|
||||
"""批量解析 asset_id → 本地路径。
|
||||
|
||||
Args:
|
||||
asset_ids: 素材 ID 列表
|
||||
work_dir: 工作目录
|
||||
|
||||
Returns:
|
||||
{asset_id: local_path} 映射,仅包含成功解析的条目。
|
||||
"""
|
||||
def resolve_asset_ids_to_paths(asset_ids: list[str], work_dir: Path) -> dict[str, Path]:
|
||||
"""批量解析 asset_id → 本地路径。"""
|
||||
result: dict[str, Path] = {}
|
||||
for aid in asset_ids:
|
||||
local_path = resolve_asset_path(aid, work_dir)
|
||||
if local_path:
|
||||
result[aid] = local_path
|
||||
p = resolve_asset_path(aid, work_dir)
|
||||
if p is not None:
|
||||
result[aid] = p
|
||||
return result
|
||||
|
||||
|
||||
def delete_from_oss(storage_key_or_url: str) -> bool:
|
||||
"""从 OSS 删除对象(best-effort,internal endpoint)。"""
|
||||
s = _storage()
|
||||
if s.bucket is None:
|
||||
return False
|
||||
try:
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
s.delete_file(key)
|
||||
return True
|
||||
except Exception:
|
||||
logger.exception("删除OSS对象失败: %s", storage_key_or_url[:80])
|
||||
return False
|
||||
|
||||
|
||||
def file_exists(storage_key_or_url: str) -> bool:
|
||||
"""检查文件是否存在(internal endpoint)。"""
|
||||
s = _storage()
|
||||
if s.bucket is None:
|
||||
return False
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
return s.file_exists(key)
|
||||
|
||||
|
||||
def get_public_url(storage_key: str) -> str:
|
||||
"""返回公网 URL(不带签名)。"""
|
||||
return _storage().get_url(storage_key)
|
||||
|
||||
@@ -86,6 +86,7 @@ class RenderAdapterResult:
|
||||
None # 封面候选帧 [{"image_url": "...", "frame_time": 5.0, "storage_key": "..."}]
|
||||
)
|
||||
temp_dir: str | None = None # 渲染临时目录,成功时由调用方清理,失败时由 finally 清理
|
||||
edge_crop_applied: bool = False # GPU 管线已做随机边缘裁剪(跳过 CPU 二次重编码)
|
||||
|
||||
def __post_init__(self):
|
||||
if self.rendered_clip_ids is None:
|
||||
@@ -131,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。
|
||||
|
||||
@@ -189,7 +191,9 @@ class RenderAdapter:
|
||||
self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)")
|
||||
|
||||
# 2. 下载素材
|
||||
asset_path_map, rendered_clip_ids, failed_clip_ids = self._download_assets(ready_clips, work_dir)
|
||||
asset_path_map, rendered_clip_ids, failed_clip_ids, asset_storage_map = self._download_assets(
|
||||
ready_clips, work_dir
|
||||
)
|
||||
if not asset_path_map:
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
@@ -206,6 +210,7 @@ class RenderAdapter:
|
||||
plan=plan,
|
||||
clips=ready_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
asset_storage_map=asset_storage_map,
|
||||
work_dir=work_dir,
|
||||
plan_id=plan_id,
|
||||
job_id=job_id,
|
||||
@@ -213,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:
|
||||
@@ -315,7 +321,7 @@ class RenderAdapter:
|
||||
|
||||
def _download_assets(
|
||||
self, clips: list[EditPlanClip], work_dir: Path
|
||||
) -> tuple[dict[str, Path], list[str], list[str]]:
|
||||
) -> tuple[dict[str, Path], list[str], list[str], dict[str, str]]:
|
||||
"""下载片段素材到本地。
|
||||
|
||||
先通过 asset_id 批量查询 assets 表获取 file_url(OSS存储路径),
|
||||
@@ -386,9 +392,9 @@ class RenderAdapter:
|
||||
failed_clip_ids.append(clip.id)
|
||||
logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
|
||||
|
||||
return asset_path_map, rendered_clip_ids, failed_clip_ids
|
||||
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 种来源(按优先级):
|
||||
@@ -401,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
|
||||
@@ -457,13 +465,27 @@ class RenderAdapter:
|
||||
from packages.domain.preset_bgm import get_preset_bgm
|
||||
|
||||
preset = get_preset_bgm(preset_id)
|
||||
if preset and preset.audio_url:
|
||||
if preset is None:
|
||||
logger.warning("[plan_id=%s] [BGM] 预设BGM不存在: preset_id=%s", plan_id, preset_id)
|
||||
elif not preset.audio_url:
|
||||
logger.warning(
|
||||
"[plan_id=%s] [BGM] 预设BGM未部署音频文件: preset_id=%s name=%s(audio_url 为空,请运维上传音频后填入 preset_bgm.py)",
|
||||
plan_id,
|
||||
preset_id,
|
||||
preset.name,
|
||||
)
|
||||
else:
|
||||
from video_processing.url_security import (
|
||||
ALLOWED_AUDIO_MIME_TYPES,
|
||||
safe_download_file,
|
||||
)
|
||||
|
||||
logger.info("[plan_id=%s] [BGM] 从预设库下载: preset_id=%s", plan_id, preset_id)
|
||||
logger.info(
|
||||
"[plan_id=%s] [BGM] 从预设库下载: preset_id=%s url=%s",
|
||||
plan_id,
|
||||
preset_id,
|
||||
preset.audio_url[:80],
|
||||
)
|
||||
safe_download_file(
|
||||
preset.audio_url,
|
||||
str(bgm_file),
|
||||
@@ -476,7 +498,14 @@ class RenderAdapter:
|
||||
except Exception as e:
|
||||
logger.warning("[plan_id=%s] [BGM] 预设库下载失败: %s", plan_id, e)
|
||||
|
||||
logger.warning("[plan_id=%s] [BGM] 所有来源都无法获取BGM,跳过", plan_id)
|
||||
logger.warning(
|
||||
"[plan_id=%s] [BGM] 所有来源都无法获取BGM(enabled=%s audio_url=%s asset_id=%s preset_id=%s),跳过",
|
||||
plan_id,
|
||||
bool(bgm_config.get("enabled")),
|
||||
"set" if audio_url else "empty",
|
||||
asset_id[:12] + "…" if len(asset_id) > 12 else asset_id or "empty",
|
||||
preset_id or "empty",
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
@@ -541,6 +570,8 @@ class RenderAdapter:
|
||||
rendered_clip_ids: list[str] | None = None,
|
||||
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 + 渲染 + 缩略图 + 上传)。
|
||||
|
||||
@@ -554,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, "执行视频渲染")
|
||||
|
||||
@@ -563,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"))
|
||||
@@ -589,9 +623,22 @@ 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 {}
|
||||
for c in clips:
|
||||
sk = _storage_map.get(getattr(c, "asset_id", ""))
|
||||
if sk:
|
||||
# EditPlanClip 使用 __slots__,不能 setattr,改存 config 字典
|
||||
if not isinstance(c.config, dict):
|
||||
c.config = dict(c.config) if c.config else {}
|
||||
c.config["_storage_key"] = sk
|
||||
result = render_svc.render()
|
||||
|
||||
# 4.4 透传 GPU 直连路径的 edge_crop 状态(供外层跳过 CPU 二次裁剪)
|
||||
edge_crop_applied_flag = bool(getattr(result, "edge_crop_applied", False))
|
||||
|
||||
# 4.5 渲染后校验输出完整性
|
||||
validation = validate_video_output(result.output_path)
|
||||
if not validation.valid:
|
||||
@@ -637,8 +684,22 @@ class RenderAdapter:
|
||||
# 已渲染视频在统一渲染阶段已通过 ASS 字幕把标题烧录进画面,
|
||||
# 抽帧天然带标题,因此这里传空字符串,避免 Pillow 二次叠加导致重影。
|
||||
# Pillow 叠加仅用于 API 从源素材抽帧(源素材本身无标题)的兜底场景。
|
||||
# 构造clip分段边界 [(start, duration), ...] 供封面抽帧智能取各段中点
|
||||
try:
|
||||
_clip_boundaries = [
|
||||
(float(getattr(c, "start_time", 0.0) or 0.0), float(getattr(c, "duration", 0.0) or 0.0))
|
||||
for c in clips
|
||||
if float(getattr(c, "duration", 0.0) or 0.0) > 0
|
||||
]
|
||||
except Exception:
|
||||
_clip_boundaries = None
|
||||
cover_candidates = extract_and_upload_cover_frames(
|
||||
str(result.output_path), plan_id, task_id=job_id, num_frames=5, title_text=""
|
||||
str(result.output_path),
|
||||
plan_id,
|
||||
task_id=job_id,
|
||||
num_frames=5,
|
||||
title_text="",
|
||||
clip_boundaries=_clip_boundaries,
|
||||
)
|
||||
if cover_candidates:
|
||||
logger.info(
|
||||
@@ -685,6 +746,7 @@ class RenderAdapter:
|
||||
rendered_clip_ids=final_rendered_ids,
|
||||
failed_clip_ids=final_failed_ids,
|
||||
cover_candidates=cover_candidates,
|
||||
edge_crop_applied=edge_crop_applied_flag,
|
||||
)
|
||||
|
||||
def render_from_memory(
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
"""视频封面抽帧工具 — 从视频中抽取帧作为封面,支持标题文字叠加。
|
||||
|
||||
统一封面管道:
|
||||
封面管道(P2 优化后):
|
||||
- 黑屏检测:ffmpeg blackdetect 扫描黑屏区间,抽帧点自动避开黑屏
|
||||
- 单次 ffmpeg select 抽多帧:一次 ffmpeg 进程用 select 滤镜输出 5 帧,避免 5 次起停进程
|
||||
- 并发上传:5 帧用 ThreadPoolExecutor 并行上传 OSS,目标封面阶段 <1.5s
|
||||
- 质量评分:cv2 清晰度/亮度/色彩三维评分选最佳帧
|
||||
- 可选 MediaKit 路径:配置 MEDIAKIT_COVER_ENABLED=true 时启用火山 MediaKit SceneChange 抽帧
|
||||
- 从已渲染视频抽帧:标题已通过 ASS 字幕烧进视频,帧天然带标题,无需再叠加。
|
||||
- 从源素材抽帧(API E2 兜底):源素材无标题,通过 Pillow 在帧上绘制标题文字。
|
||||
"""
|
||||
@@ -8,14 +13,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import tempfile
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 标题叠加(Pillow)──────────────────────────────────────────────────────
|
||||
# 实现统一放在 packages/shared/title_overlay.py,API 和 Worker 共用。
|
||||
|
||||
|
||||
def apply_title_overlay(
|
||||
image_path: str,
|
||||
@@ -27,11 +32,7 @@ def apply_title_overlay(
|
||||
margin_ratio: float = 0.06,
|
||||
stroke_width_ratio: float = 0.04,
|
||||
) -> str:
|
||||
"""在图片上绘制标题文字(指定颜色 + 黑色描边/阴影)。
|
||||
|
||||
委托给 packages.shared.title_overlay.apply_title_to_image,
|
||||
保持 Worker 内调用方式不变。title_text 为空时直接返回原路径。
|
||||
"""
|
||||
"""在图片上绘制标题文字(指定颜色 + 黑色描边/阴影)。"""
|
||||
from packages.shared.title_overlay import apply_title_to_image
|
||||
|
||||
if not title_text or not title_text.strip():
|
||||
@@ -56,26 +57,19 @@ def extract_first_frame(
|
||||
height: int = -1,
|
||||
timeout: int = 30,
|
||||
seek_ratio: float = 0.15,
|
||||
seek_seconds: float | None = None,
|
||||
min_seek_seconds: float = 1.0,
|
||||
) -> str:
|
||||
"""抽取视频封面帧(默认取视频时长 15% 处的帧,避开片头纯色画面)。
|
||||
|
||||
因为视频渲染时标题已通过 ASS 字幕烧录,抽取的帧天然带标题。
|
||||
"""抽取视频封面帧(ffmpeg -ss 单帧 seek,<100ms/帧)。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
output_path: 输出图片路径,不传则用临时文件
|
||||
width: 输出宽度(默认 -1,保持原始分辨率)
|
||||
height: 输出高度(默认 -1,保持原始分辨率)
|
||||
timeout: 超时时间(秒)
|
||||
seek_ratio: 抽帧位置占视频时长的比例(默认 0.15,即 15% 处)
|
||||
min_seek_seconds: 最小抽帧时间(秒),避免极短视频 seek 到 0
|
||||
|
||||
Returns:
|
||||
生成的封面帧文件路径
|
||||
|
||||
Raises:
|
||||
RuntimeError: ffmpeg 执行失败或输出文件为空
|
||||
width/height: 输出宽高(默认保持原始分辨率)
|
||||
timeout: 超时(秒)
|
||||
seek_ratio: 抽帧位置占视频时长的比例
|
||||
seek_seconds: 指定具体抽帧时间点(秒),优先于 seek_ratio
|
||||
min_seek_seconds: 最小抽帧时间
|
||||
"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
|
||||
|
||||
@@ -87,31 +81,25 @@ def extract_first_frame(
|
||||
_is_temp_output = True
|
||||
|
||||
try:
|
||||
# 计算抽帧时间点:取视频时长 * seek_ratio,最少 min_seek_seconds 秒
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
seek_time = max(min_seek_seconds, duration * seek_ratio)
|
||||
except Exception:
|
||||
# probe 失败时 fallback 到第1秒
|
||||
seek_time = min_seek_seconds
|
||||
if seek_seconds is not None:
|
||||
seek_time = max(0.0, float(seek_seconds))
|
||||
else:
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
seek_time = max(min_seek_seconds, duration * seek_ratio)
|
||||
except Exception:
|
||||
seek_time = min_seek_seconds
|
||||
|
||||
# 格式化为 HH:MM:SS.xx
|
||||
seek_str = _format_seek_time(seek_time)
|
||||
|
||||
# 构建 scale filter:如果指定了宽高则缩放,否则保持原始分辨率。
|
||||
# NOTE: scale_filter 在此处通过 if/else 分支赋值,之后不再被覆盖,
|
||||
# 后续 cmd / cmd2 均复用同一变量,逻辑无变化。
|
||||
if width > 0 or height > 0:
|
||||
w_str = str(width) if width > 0 else "-1"
|
||||
h_str = str(height) if height > 0 else "-1"
|
||||
scale_filter = f"scale={w_str}:{h_str}:force_original_aspect_ratio=decrease,format=yuvj420p"
|
||||
else:
|
||||
# 保持原始分辨率,只确保格式兼容
|
||||
scale_filter = "format=yuvj420p"
|
||||
|
||||
# -ss 放在 -i 前面(input seeking,更快)
|
||||
# -vframes 1 只取一帧
|
||||
# -q:v 2 jpeg 高质量
|
||||
# -ss 放在 -i 前面(input seeking,极快),-vframes 1 只取一帧
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
@@ -154,7 +142,6 @@ def extract_first_frame(
|
||||
|
||||
return output_path
|
||||
except Exception:
|
||||
# 失败时清理自己创建的临时文件
|
||||
if _is_temp_output and output_path:
|
||||
try:
|
||||
Path(output_path).unlink(missing_ok=True)
|
||||
@@ -164,7 +151,6 @@ def extract_first_frame(
|
||||
|
||||
|
||||
def _format_seek_time(seconds: float) -> str:
|
||||
"""将秒数格式化为 HH:MM:SS.xx 格式。"""
|
||||
h = int(seconds // 3600)
|
||||
m = int((seconds % 3600) // 60)
|
||||
s = seconds % 60
|
||||
@@ -177,19 +163,7 @@ def generate_and_upload_thumbnail(
|
||||
*,
|
||||
seek_ratio: float = 0.15,
|
||||
) -> str:
|
||||
"""从视频中提取一帧缩略图并上传到 OSS。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
storage_key: OSS 存储 key
|
||||
seek_ratio: 抽帧位置比例(默认 0.15)
|
||||
|
||||
Returns:
|
||||
上传后的 URL 字符串
|
||||
|
||||
Raises:
|
||||
RuntimeError: 抽帧或上传失败
|
||||
"""
|
||||
"""从视频中提取一帧缩略图并上传到 OSS。"""
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
@@ -204,24 +178,303 @@ def generate_and_upload_thumbnail(
|
||||
Path(tmp.name).unlink(missing_ok=True)
|
||||
|
||||
|
||||
def _detect_black_intervals(
|
||||
video_path: str,
|
||||
duration: float,
|
||||
*,
|
||||
black_min_duration: float = 0.3,
|
||||
picture_black_ratio_th: float = 0.98,
|
||||
pixel_black_th: float = 0.10,
|
||||
timeout: int = 30,
|
||||
) -> list[tuple[float, float]]:
|
||||
"""用 ffmpeg blackdetect 扫描黑屏区间,返回 [(start, end), ...]。"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
|
||||
|
||||
if duration <= 0:
|
||||
return []
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-nostdin",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vf",
|
||||
(f"blackdetect=d={black_min_duration:.2f}:pic_th={picture_black_ratio_th:.2f}:pix_th={pixel_black_th:.2f}"),
|
||||
"-an",
|
||||
"-f",
|
||||
"null",
|
||||
"-",
|
||||
]
|
||||
try:
|
||||
_, stderr = run_ffmpeg(cmd, capture_output=True, timeout=timeout)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] blackdetect 失败,忽略黑屏规避: %s", e)
|
||||
return []
|
||||
|
||||
intervals: list[tuple[float, float]] = []
|
||||
pattern = re.compile(
|
||||
r"black_start:(\d+(?:\.\d+)?)\s+black_end:(\d+(?:\.\d+)?)\s+black_duration:(\d+(?:\.\d+)?)",
|
||||
)
|
||||
for m in pattern.finditer(stderr or ""):
|
||||
try:
|
||||
bs = float(m.group(1))
|
||||
be = float(m.group(2))
|
||||
intervals.append((bs, be))
|
||||
except ValueError:
|
||||
continue
|
||||
intervals.sort()
|
||||
if intervals:
|
||||
logger.info("[thumbnail] blackdetect 发现 %d 段黑屏: %s", len(intervals), intervals[:5])
|
||||
return intervals
|
||||
|
||||
|
||||
def _adjust_seek_points_avoid_black(
|
||||
seek_points: list[float],
|
||||
black_intervals: list[tuple[float, float]],
|
||||
duration: float,
|
||||
*,
|
||||
tolerance: float = 0.25,
|
||||
) -> list[float]:
|
||||
"""把落在黑屏区间的 seek 点偏移到最近的非黑屏位置。
|
||||
|
||||
策略:
|
||||
- 若点在黑屏内,先尝试向前偏移到黑屏起点 - tolerance,再尝试向后偏移到黑屏终点 + tolerance;
|
||||
- 若整个视频全黑(偏移后 <0 或 >duration),保留原点但日志标记警告;
|
||||
- 偏移后若点与已有点重合(误差 <0.3s),做微调去重。
|
||||
"""
|
||||
if not black_intervals or not seek_points:
|
||||
return list(seek_points)
|
||||
|
||||
def in_black(t: float) -> tuple[float, float] | None:
|
||||
for bs, be in black_intervals:
|
||||
if bs <= t <= be:
|
||||
return (bs, be)
|
||||
return None
|
||||
|
||||
adjusted: list[float] = []
|
||||
for t in seek_points:
|
||||
seg = in_black(t)
|
||||
if seg is None:
|
||||
adjusted.append(max(0.0, min(duration, t)))
|
||||
continue
|
||||
bs, be = seg
|
||||
# 先尝试向前
|
||||
forward_t = bs - tolerance
|
||||
if forward_t >= 0.0 and in_black(forward_t) is None:
|
||||
adjusted.append(forward_t)
|
||||
continue
|
||||
# 再尝试向后
|
||||
backward_t = be + tolerance
|
||||
if backward_t <= duration and in_black(backward_t) is None:
|
||||
adjusted.append(backward_t)
|
||||
continue
|
||||
# 整段 clip 全黑?保留中点但标记
|
||||
logger.warning(
|
||||
"[thumbnail] seek 点 %.2fs 落在黑屏区间 [%.2f,%.2f] 且无法偏移,保留原位置(可能是全黑片段)",
|
||||
t,
|
||||
bs,
|
||||
be,
|
||||
)
|
||||
adjusted.append(max(0.0, min(duration, t)))
|
||||
|
||||
# 去重:相邻点若 <0.3s 则拉开
|
||||
adjusted.sort()
|
||||
deduped: list[float] = []
|
||||
for t in adjusted:
|
||||
if not deduped or abs(t - deduped[-1]) >= 0.3:
|
||||
deduped.append(t)
|
||||
else:
|
||||
# 往后挪 0.5s
|
||||
nt = t + 0.5
|
||||
if nt <= duration and in_black(nt) is None:
|
||||
deduped.append(nt)
|
||||
else:
|
||||
deduped.append(t)
|
||||
return [round(max(0.0, min(duration, t)), 3) for t in deduped[: len(seek_points)]]
|
||||
|
||||
|
||||
def _extract_frames_single_pass(
|
||||
video_path: str,
|
||||
seek_points: list[float],
|
||||
out_dir: str,
|
||||
*,
|
||||
prefix: str = "frame",
|
||||
width: int = -1,
|
||||
height: int = -1,
|
||||
q: int = 2,
|
||||
timeout: int = 30,
|
||||
) -> list[tuple[float, str]]:
|
||||
"""单次 ffmpeg 用 select 滤镜抽出 seek_points 对应的多帧。
|
||||
|
||||
ffmpeg -i input -vf "select='between(t,t1-0.03,t1+0.03)+between(t,t2-0.03,t2+0.03)+...',scale=...,format=yuvj420p"
|
||||
-vsync vfr -q:v 2 out_dir/prefix_%02d.jpg
|
||||
|
||||
返回 [(seek_t, output_path), ...],按输出帧序号升序。若输出帧数 < seek_points 数量,
|
||||
不足部分用 extract_first_frame 兜底(保证返回数量 == len(seek_points))。
|
||||
"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
|
||||
|
||||
out_dir_p = Path(out_dir)
|
||||
out_dir_p.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 构造 select 表达式:每个 seek 点用 ±30ms 窗口命中
|
||||
# between(t, a, b) 返回 1 表示 t 在 [a,b] 内;多个 between 相加即为"任一命中"
|
||||
select_terms = []
|
||||
for t in seek_points:
|
||||
a = max(0.0, t - 0.03)
|
||||
b = t + 0.04
|
||||
select_terms.append(f"between(t,{a:.3f},{b:.3f})")
|
||||
select_expr = "+".join(select_terms)
|
||||
|
||||
if width > 0 or height > 0:
|
||||
w_str = str(width) if width > 0 else "-1"
|
||||
h_str = str(height) if height > 0 else "-1"
|
||||
scale_filter = f"scale={w_str}:{h_str}:force_original_aspect_ratio=decrease"
|
||||
vf = f"select='{select_expr}',{scale_filter},format=yuvj420p"
|
||||
else:
|
||||
vf = f"select='{select_expr}',format=yuvj420p"
|
||||
|
||||
out_pattern = str(out_dir_p / f"{prefix}_%02d.jpg")
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vf",
|
||||
vf,
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
str(q),
|
||||
out_pattern,
|
||||
]
|
||||
|
||||
results: list[tuple[float, str]] = []
|
||||
single_pass_ok = False
|
||||
try:
|
||||
run_ffmpeg(cmd, capture_output=True, timeout=timeout)
|
||||
# 读取输出文件
|
||||
for i in range(1, len(seek_points) + 1):
|
||||
fp = out_dir_p / f"{prefix}_{i:02d}.jpg"
|
||||
if fp.exists() and fp.stat().st_size > 0:
|
||||
results.append((seek_points[i - 1] if i - 1 < len(seek_points) else 0.0, str(fp)))
|
||||
if len(results) >= len(seek_points):
|
||||
single_pass_ok = True
|
||||
else:
|
||||
logger.warning(
|
||||
"[thumbnail] 单次 ffmpeg 抽帧仅命中 %d/%d 帧,不足部分用单帧 seek 兜底",
|
||||
len(results),
|
||||
len(seek_points),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 单次 ffmpeg select 抽帧失败,回退到单帧 seek: %s", e)
|
||||
|
||||
# 兜底:对缺失/失败的帧用 extract_first_frame 补抽
|
||||
if not single_pass_ok:
|
||||
# 清理不完整结果
|
||||
for _, fp in results:
|
||||
try:
|
||||
Path(fp).unlink(missing_ok=True)
|
||||
except Exception:
|
||||
pass
|
||||
results = []
|
||||
for i, st in enumerate(seek_points):
|
||||
fp = out_dir_p / f"{prefix}_fallback_{i:02d}.jpg"
|
||||
try:
|
||||
extract_first_frame(
|
||||
video_path,
|
||||
output_path=str(fp),
|
||||
seek_seconds=st,
|
||||
min_seek_seconds=0.5,
|
||||
timeout=timeout,
|
||||
)
|
||||
if fp.exists() and fp.stat().st_size > 0:
|
||||
results.append((st, str(fp)))
|
||||
else:
|
||||
logger.warning("[thumbnail] 兜底单帧抽帧也失败 idx=%d t=%.2f", i, st)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 兜底单帧抽帧异常 idx=%d t=%.2f: %s", i, st, e)
|
||||
|
||||
return results[: len(seek_points)]
|
||||
|
||||
|
||||
def _compute_clip_boundary_seek_points(
|
||||
duration: float,
|
||||
clip_boundaries: Optional[list[tuple[float, float]]] = None,
|
||||
num_frames: int = 5,
|
||||
head_skip_ratio: float = 0.08,
|
||||
tail_skip_ratio: float = 0.08,
|
||||
) -> list[float]:
|
||||
"""基于clip分段边界计算抽帧时间点(取每段中间帧,效果比均匀抽更好)。
|
||||
|
||||
策略:
|
||||
- 如果传入 clip_boundaries(每个元素是 (clip_start_in_timeline, clip_duration)),
|
||||
取每个片段的中点作为抽帧候选点
|
||||
- 候选点不足 num_frames 时,均匀补充
|
||||
- 跳过片头 head_skip_ratio(8%,避免片头黑屏/开场标题)和片尾 tail_skip_ratio(8%)
|
||||
- 返回按时间排序的 num_frames 个抽帧点(秒)
|
||||
"""
|
||||
if duration <= 0:
|
||||
# 无法probe,均匀分布兜底
|
||||
return [max(1.0, duration * (0.1 + 0.8 * i / max(num_frames - 1, 1))) for i in range(num_frames)]
|
||||
|
||||
head_skip = duration * head_skip_ratio
|
||||
tail_skip = duration * tail_skip_ratio
|
||||
valid_start = head_skip
|
||||
valid_end = max(valid_start + 1.0, duration - tail_skip)
|
||||
|
||||
candidates: list[float] = []
|
||||
|
||||
if clip_boundaries:
|
||||
# 累加timeline start,取每clip中点
|
||||
cur = 0.0
|
||||
for _clip_start, clip_dur in clip_boundaries:
|
||||
if clip_dur <= 0:
|
||||
continue
|
||||
mid = cur + clip_dur / 2.0
|
||||
if valid_start <= mid <= valid_end:
|
||||
candidates.append(mid)
|
||||
cur += clip_dur
|
||||
# 去重+排序
|
||||
candidates = sorted(set(round(c, 3) for c in candidates))
|
||||
|
||||
# 如果候选点不足,均匀补充
|
||||
if len(candidates) < num_frames:
|
||||
needed = num_frames - len(candidates)
|
||||
existing = set(round(c, 1) for c in candidates)
|
||||
for i in range(needed * 3):
|
||||
ratio = 0.1 + 0.8 * (i + 0.5) / (needed * 3)
|
||||
t = valid_start + (valid_end - valid_start) * ratio
|
||||
if round(t, 1) not in existing:
|
||||
candidates.append(t)
|
||||
existing.add(round(t, 1))
|
||||
if len(candidates) >= num_frames:
|
||||
break
|
||||
|
||||
# 如果还不够,强制均匀
|
||||
while len(candidates) < num_frames:
|
||||
idx = len(candidates)
|
||||
ratio = 0.1 + 0.8 * idx / max(num_frames - 1, 1)
|
||||
candidates.append(valid_start + (valid_end - valid_start) * ratio)
|
||||
|
||||
candidates.sort()
|
||||
|
||||
# 如果超过num_frames,均匀选取
|
||||
if len(candidates) > num_frames:
|
||||
step = len(candidates) / num_frames
|
||||
candidates = [candidates[int(i * step)] for i in range(num_frames)]
|
||||
|
||||
return [round(t, 3) for t in candidates[:num_frames]]
|
||||
|
||||
|
||||
def _extract_frames_via_mediakit(
|
||||
video_path: str,
|
||||
plan_id: str,
|
||||
num_frames: int,
|
||||
) -> list[dict] | None:
|
||||
"""使用 MediaKit 智能抽帧 API 提取封面帧。
|
||||
|
||||
Args:
|
||||
video_path: 本地视频文件路径
|
||||
plan_id: 编辑计划 ID
|
||||
num_frames: 需要的帧数
|
||||
|
||||
Returns:
|
||||
帧列表 [{"image_url": str, "timestamp": float}, ...],失败返回 None
|
||||
"""
|
||||
"""使用 MediaKit 智能抽帧 API 提取封面帧(fallback 路径,默认不启用)。"""
|
||||
import uuid
|
||||
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
from video_processing.oss_helpers import delete_from_oss, get_signed_download_url, upload_to_oss
|
||||
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
|
||||
@@ -230,45 +483,38 @@ def _extract_frames_via_mediakit(
|
||||
logger.info("[thumbnail] MediaKit 未配置,跳过智能抽帧")
|
||||
return None
|
||||
|
||||
# 1. 上传视频到 OSS 获取 URL
|
||||
video_storage_key: str = ""
|
||||
try:
|
||||
video_storage_key = f"temp/{plan_id}/{uuid.uuid4().hex[:8]}_{Path(video_path).name}"
|
||||
video_url = upload_to_oss(video_path, video_storage_key)
|
||||
if not video_url:
|
||||
public_url = upload_to_oss(video_path, video_storage_key)
|
||||
if not public_url:
|
||||
logger.warning("[thumbnail] 视频上传 OSS 失败,无法使用 MediaKit")
|
||||
return None
|
||||
logger.info("[thumbnail] 视频已上传 OSS: %s", video_url[:80])
|
||||
video_url = get_signed_download_url(video_storage_key, expires_seconds=3600) or public_url
|
||||
logger.info("[thumbnail] 视频已上传 OSS 并生成签名 URL: key=%s", video_storage_key[:80])
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 视频上传 OSS 异常: %s,降级到 ffmpeg", e)
|
||||
logger.warning("[thumbnail] 视频上传 OSS 异常: %s,降级到本地 ffmpeg", e)
|
||||
return None
|
||||
|
||||
# 2. 调用 MediaKit 智能抽帧
|
||||
try:
|
||||
frames = client.extract_frames(
|
||||
video_url=video_url,
|
||||
strategy="SceneChange",
|
||||
max_frames=num_frames * 2, # 多取一些帧供选择
|
||||
max_frames=num_frames * 2,
|
||||
)
|
||||
if not frames:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧返回空,降级到 ffmpeg")
|
||||
logger.warning("[thumbnail] MediaKit 抽帧返回空")
|
||||
return None
|
||||
|
||||
# 选取最均匀的 num_frames 个帧
|
||||
if len(frames) > num_frames:
|
||||
step = len(frames) // num_frames
|
||||
frames = [frames[i * step] for i in range(num_frames)]
|
||||
|
||||
logger.info("[thumbnail] MediaKit 抽帧成功: %d 帧", len(frames))
|
||||
return frames
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧异常: %s,降级到 ffmpeg", e)
|
||||
logger.warning("[thumbnail] MediaKit 抽帧异常: %s", e)
|
||||
return None
|
||||
finally:
|
||||
# 清理临时视频文件
|
||||
try:
|
||||
from video_processing.oss_helpers import delete_from_oss
|
||||
|
||||
delete_from_oss(video_storage_key)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -279,161 +525,235 @@ def extract_and_upload_cover_frames(
|
||||
plan_id: str,
|
||||
*,
|
||||
task_id: str = "",
|
||||
num_frames: int = 5, # 抽 5 帧候选,通过质量评分选出最佳帧
|
||||
num_frames: int = 5,
|
||||
title_text: str = "",
|
||||
title_color: str = "#ffffff",
|
||||
title_position: str = "bottom",
|
||||
title_font_size: int | None = None,
|
||||
clip_boundaries: Optional[list[tuple[float, float]]] = None,
|
||||
) -> list[dict]:
|
||||
"""从视频中抽取多帧作为封面候选,通过质量评分选出最佳帧,上传到 OSS。
|
||||
|
||||
流程:
|
||||
1. 优先使用 MediaKit 智能抽帧(多抽一些供选择)
|
||||
2. MediaKit 不足时降级到 ffmpeg 均匀抽帧
|
||||
3. 对所有候选帧进行质量评分(清晰度/亮度/色彩丰富度)
|
||||
4. 按分数从高到低排序返回
|
||||
P2 优化:
|
||||
- 先用 ffmpeg blackdetect 扫描黑屏区间,seek 点自动避开黑屏
|
||||
- 单次 ffmpeg select 抽 num_frames 帧(避免 5 次起停 ffmpeg 进程)
|
||||
- 多帧 OSS 上传用 ThreadPoolExecutor 并发,目标封面阶段 <1.5s
|
||||
- cv2 清晰度/亮度/色彩三维评分选最佳帧
|
||||
|
||||
Fallback(MEDIAKIT_COVER_ENABLED=true):火山 MediaKit SceneChange 抽帧(~60-90s)。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
plan_id: 编辑计划 ID(用于生成 storage key)
|
||||
task_id: 任务 ID(用于生成独立的 storage key,避免标题变更时封面冲突)
|
||||
num_frames: 抽取候选帧数(默认 5,通过质量评分选出最佳帧)
|
||||
title_text: 标题文字;非空时用 Pillow 叠加到每帧。
|
||||
从已渲染视频抽帧时通常传空(标题已烧录);从源素材抽帧时传标题。
|
||||
title_color: 标题字体颜色(#RRGGBB)
|
||||
title_position: 标题位置 top/center/bottom
|
||||
title_font_size: 标题字号,None 时自动计算
|
||||
|
||||
Returns:
|
||||
封面候选列表(按质量分数降序),每项包含 {"url": str, "position": float, "score": float}
|
||||
clip_boundaries: 片段边界列表 [(clip_start, clip_duration), ...],用于智能取点
|
||||
"""
|
||||
import time
|
||||
|
||||
import httpx
|
||||
from video_processing.ffmpeg_utils import probe_duration
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
t0 = time.monotonic()
|
||||
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
except Exception:
|
||||
duration = 0.0
|
||||
|
||||
candidates: list[dict] = []
|
||||
_temp_paths: list[str] = [] # 收集所有临时文件路径,最后统一清理
|
||||
_temp_paths: list[str] = []
|
||||
|
||||
try:
|
||||
# ── 阶段 1:抽帧 ──────────────────────────────────────────────
|
||||
# 优先尝试 MediaKit 智能抽帧
|
||||
mediakit_frames = _extract_frames_via_mediakit(video_path, plan_id, num_frames)
|
||||
if mediakit_frames:
|
||||
for i, frame in enumerate(mediakit_frames):
|
||||
frame_url = frame.get("image_url")
|
||||
if not frame_url:
|
||||
continue
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
# 下载 MediaKit 返回的帧图
|
||||
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
with open(tmp.name, "wb") as f:
|
||||
f.write(resp.content)
|
||||
settings = get_shared_settings()
|
||||
use_mediakit = getattr(settings, "mediakit_cover_enabled", False)
|
||||
|
||||
# 叠加标题文字(如需要)
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
tmp.name,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
if use_mediakit:
|
||||
logger.info("[thumbnail] MEDIAKIT_COVER_ENABLED=true,走 MediaKit 路径")
|
||||
mediakit_frames = _extract_frames_via_mediakit(video_path, plan_id, num_frames)
|
||||
if mediakit_frames:
|
||||
for i, frame in enumerate(mediakit_frames):
|
||||
frame_url = frame.get("image_url")
|
||||
if not frame_url:
|
||||
continue
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
with open(tmp.name, "wb") as f:
|
||||
f.write(resp.content)
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
tmp.name,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
storage_key = f"covers/{plan_id}/{task_id}/mediakit_frame_{i}.jpg"
|
||||
url = upload_to_oss(tmp.name, storage_key)
|
||||
if url:
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(frame.get("timestamp", 0.0), 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 帧 %d 处理失败: %s", i, e)
|
||||
if len(candidates) >= num_frames:
|
||||
logger.info("[thumbnail] MediaKit 抽帧完成: %d 帧", len(candidates))
|
||||
|
||||
# MediaKit 路径帧在 NamedTemporaryFile 中持久存在(finally 清理),在进入本地 ffmpeg 前评分
|
||||
if len(candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
candidates = score_frames(candidates)
|
||||
logger.info(
|
||||
"[thumbnail] MediaKit 封面帧评分完成: count=%d best_score=%.1f",
|
||||
len(candidates),
|
||||
candidates[0].get("score", 0.0) if candidates else 0.0,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("[thumbnail] MediaKit 封面帧质量评分失败,保持原始顺序", exc_info=True)
|
||||
|
||||
storage_key = f"covers/{plan_id}/{task_id}/mediakit_frame_{i}.jpg"
|
||||
url = upload_to_oss(tmp.name, storage_key)
|
||||
if url:
|
||||
seek_time = frame.get("timestamp", 0.0)
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(seek_time, 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 帧 %d 处理失败: %s", i, e)
|
||||
|
||||
if len(candidates) >= num_frames:
|
||||
logger.info("[thumbnail] MediaKit 智能抽帧完成: %d 帧", len(candidates))
|
||||
else:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧不足 %d 帧,降级到 ffmpeg", num_frames)
|
||||
|
||||
# Fallback: ffmpeg 直接抽帧(仅当 MediaKit 不足时)
|
||||
# ── 默认路径:本地 ffmpeg 单次 select 抽帧 + 并发上传 ──────────────
|
||||
if len(candidates) < num_frames:
|
||||
logger.info("[thumbnail] 使用 ffmpeg 抽帧补充")
|
||||
# 均匀分布抽帧点:从 10% 到 90%
|
||||
for i in range(num_frames):
|
||||
ratio = 0.1 + 0.8 * i / max(num_frames - 1, 1)
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
frame_path = extract_first_frame(
|
||||
video_path,
|
||||
output_path=tmp.name,
|
||||
seek_ratio=ratio,
|
||||
min_seek_seconds=0.5,
|
||||
)
|
||||
# 从源素材抽帧时叠加标题文字;已渲染视频标题已烧录时传空字符串跳过
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
frame_path,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
storage_key = f"covers/{plan_id}/{task_id}/frame_{i}.jpg"
|
||||
url = upload_to_oss(frame_path, storage_key)
|
||||
if url:
|
||||
seek_time = max(0.5, duration * ratio) if duration > 0 else 0.0
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(seek_time, 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 封面候选帧 %d 提取失败: %s", i, e)
|
||||
|
||||
# ── 阶段 2:质量评分 ────────────────────────────────────────────
|
||||
if len(candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
candidates = score_frames(candidates)
|
||||
if candidates:
|
||||
logger.info("[thumbnail] MediaKit 不足 %d 帧,本地 ffmpeg 补充", num_frames)
|
||||
else:
|
||||
logger.info(
|
||||
"[thumbnail] 封面帧质量评分完成: plan_id=%s count=%d best_score=%.1f",
|
||||
"[thumbnail] 使用本地 ffmpeg 抽帧(num=%d, duration=%.1fs)",
|
||||
num_frames,
|
||||
duration,
|
||||
)
|
||||
|
||||
# 1) 计算 seek 点
|
||||
seek_points = _compute_clip_boundary_seek_points(duration, clip_boundaries, num_frames)
|
||||
|
||||
# 2) 黑屏检测 + 偏移 seek 点
|
||||
black_intervals = _detect_black_intervals(video_path, duration) if duration > 0 else []
|
||||
if black_intervals:
|
||||
seek_points = _adjust_seek_points_avoid_black(seek_points, black_intervals, duration)
|
||||
logger.info("[thumbnail] 黑屏规避后 seek 点: %s", seek_points)
|
||||
|
||||
# 3) 单次 ffmpeg select 抽出所有帧(带失败兜底到单帧 seek)
|
||||
with tempfile.TemporaryDirectory(prefix="thumb_") as frame_dir:
|
||||
t1 = time.monotonic()
|
||||
frame_results = _extract_frames_single_pass(
|
||||
video_path,
|
||||
seek_points,
|
||||
frame_dir,
|
||||
prefix="frame",
|
||||
)
|
||||
logger.info("[thumbnail] 抽帧耗时: %.2fs (%d 帧)", time.monotonic() - t1, len(frame_results))
|
||||
|
||||
# 4) 标题叠加(本地,CPU 很快)
|
||||
for _st, fp in frame_results:
|
||||
if title_text and title_text.strip():
|
||||
try:
|
||||
apply_title_overlay(
|
||||
fp,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 标题叠加失败 %s: %s", fp, e)
|
||||
|
||||
# 5) 质量评分(必须在 TemporaryDirectory 内,帧文件还在磁盘上)
|
||||
t_score = time.monotonic()
|
||||
local_candidates: list[dict] = [{"position": st, "image_path": fp} for (st, fp) in frame_results]
|
||||
scored: list[dict] = local_candidates
|
||||
if len(local_candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
scored = score_frames(local_candidates)
|
||||
logger.info(
|
||||
"[thumbnail] 封面评分耗时: %.2fs (best_score=%.1f, count=%d)",
|
||||
time.monotonic() - t_score,
|
||||
scored[0].get("score", 0.0) if scored else 0.0,
|
||||
len(scored),
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"[thumbnail] 封面帧质量评分失败,保持 seek 点原始顺序",
|
||||
exc_info=True,
|
||||
)
|
||||
scored = local_candidates
|
||||
|
||||
# 6) 按评分顺序并发上传 OSS(best 帧先上传;best 已是 scored[0])
|
||||
t2 = time.monotonic()
|
||||
|
||||
def _upload_one(rank: int, st: float, fp: str, score: float) -> dict | None:
|
||||
try:
|
||||
storage_key = f"covers/{plan_id}/{task_id}/frame_{rank}.jpg"
|
||||
url = upload_to_oss(fp, storage_key)
|
||||
if url:
|
||||
return {
|
||||
"url": url,
|
||||
"position": st,
|
||||
"image_path": fp,
|
||||
"score": score,
|
||||
"is_best": rank == 0,
|
||||
}
|
||||
logger.warning("[thumbnail] 上传失败 rank=%d t=%.2f", rank, st)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 上传异常 rank=%d t=%.2f: %s", rank, st, e)
|
||||
return None
|
||||
|
||||
upload_results: list[dict | None] = [None] * len(scored)
|
||||
max_workers = min(8, max(2, len(scored)))
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
||||
future_map = {
|
||||
pool.submit(
|
||||
_upload_one,
|
||||
i,
|
||||
float(c.get("position", 0.0)),
|
||||
str(c["image_path"]),
|
||||
float(c.get("score", 0.0)),
|
||||
): i
|
||||
for i, c in enumerate(scored)
|
||||
}
|
||||
for fut in as_completed(future_map):
|
||||
i = future_map[fut]
|
||||
try:
|
||||
upload_results[i] = fut.result()
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 上传 future 异常 rank=%d: %s", i, e)
|
||||
logger.info("[thumbnail] 并发上传耗时: %.2fs", time.monotonic() - t2)
|
||||
|
||||
for r in upload_results:
|
||||
if r is not None:
|
||||
# 本地帧在 TemporaryDirectory 内,with 退出自动删除,无需进 _temp_paths
|
||||
candidates.append(r)
|
||||
|
||||
# 如果本地 ffmpeg 路径产生了候选(已评分)但未经过 MediaKit 路径,candidates 已按评分顺序排好。
|
||||
# 混合场景下(MediaKit + 本地 ffmpeg 都产出),统一按 score 降序排列;缺失 score 的(理论上不应出现)排末尾。
|
||||
if len(candidates) > 1:
|
||||
candidates.sort(key=lambda c: c.get("score", -1.0), reverse=True)
|
||||
if candidates:
|
||||
candidates[0]["is_best"] = True
|
||||
elapsed = time.monotonic() - t0
|
||||
logger.info(
|
||||
"[thumbnail] 封面完成: plan_id=%s count=%d best=t%.2fs score=%.1f elapsed=%.2fs",
|
||||
plan_id,
|
||||
len(candidates),
|
||||
candidates[0].get("score", 0.0) if candidates else 0.0,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"[thumbnail] 封面帧质量评分失败,保持原始顺序: plan_id=%s",
|
||||
plan_id,
|
||||
exc_info=True,
|
||||
candidates[0].get("position", 0.0),
|
||||
candidates[0].get("score", 0.0),
|
||||
elapsed,
|
||||
)
|
||||
|
||||
# ── 阶段 3:清理临时文件 ────────────────────────────────────────
|
||||
# 移除 image_path(不再需要),但临时文件统一清理
|
||||
for c in candidates:
|
||||
c.pop("image_path", None)
|
||||
|
||||
return candidates
|
||||
|
||||
finally:
|
||||
# 统一清理所有临时文件
|
||||
for path in _temp_paths:
|
||||
try:
|
||||
Path(path).unlink(missing_ok=True)
|
||||
|
||||
@@ -112,6 +112,7 @@ class RenderResult:
|
||||
file_size: int
|
||||
width: int
|
||||
height: int
|
||||
edge_crop_applied: bool = False # True = GPU管线已做随机边缘裁剪
|
||||
|
||||
|
||||
# ── clip_type → layer role 映射 ──────────────────────────────────────────────
|
||||
@@ -157,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
|
||||
@@ -169,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)
|
||||
@@ -179,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(向后兼容)。"""
|
||||
@@ -233,7 +260,7 @@ class UnifiedRenderService:
|
||||
return
|
||||
if abs(mt.brightness) > 1e-4 or abs(mt.contrast - 1.0) > 1e-4 or abs(mt.saturation - 1.0) > 1e-4:
|
||||
filters.append(
|
||||
f"eq=brightness={mt.brightness:+.4f}:" f"contrast={mt.contrast:.4f}:saturation={mt.saturation:.4f}"
|
||||
f"eq=brightness={mt.brightness:+.4f}:contrast={mt.contrast:.4f}:saturation={mt.saturation:.4f}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -341,6 +368,36 @@ class UnifiedRenderService:
|
||||
len(pip_sources),
|
||||
)
|
||||
|
||||
# 4.8 全 GPU 直连管线(P1):命中主流场景则跳过 mezzanine/边缘裁剪 CPU 重编码
|
||||
output_path = self.work_dir / f"rendered_{self.plan.id}.mp4"
|
||||
direct_result = self._try_gpu_direct(
|
||||
layers=layers,
|
||||
ass_path=ass_path,
|
||||
video_duration=video_duration_final,
|
||||
output_path=output_path,
|
||||
)
|
||||
if direct_result is not None and direct_result[0]:
|
||||
_direct_edge_crop = bool(direct_result[1])
|
||||
# 直连成功:直接探测并返回,跳过后续视频/音频 CPU 流程
|
||||
duration, file_size, width, height = self._probe_output(output_path)
|
||||
logger.info(
|
||||
"[unified-render] gpu-direct done: plan_id=%s total_ms=%d output_size=%d resolution=%dx%d",
|
||||
self.plan.id,
|
||||
int((time.time() - t_start) * 1000),
|
||||
file_size,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
direct_edge_cropped = _direct_edge_crop # GPU直连时若dedup=True已在GPU内做随机边缘裁剪
|
||||
return RenderResult(
|
||||
output_path=output_path,
|
||||
duration=duration,
|
||||
file_size=file_size,
|
||||
width=width,
|
||||
height=height,
|
||||
edge_crop_applied=direct_edge_cropped,
|
||||
)
|
||||
|
||||
# 5. 视频主渲染
|
||||
t_video_start = time.time()
|
||||
video_only_path = self.work_dir / f"rendered_{self.plan.id}_video.mp4"
|
||||
@@ -400,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)
|
||||
@@ -440,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 = {}
|
||||
@@ -665,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():
|
||||
@@ -850,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 = {}
|
||||
@@ -2180,6 +2237,261 @@ class UnifiedRenderService:
|
||||
|
||||
# ── GPU NVENC 加速 ────────────────────────────────────────────────────
|
||||
|
||||
# ── 全 GPU 直连渲染(P1)─────────────────────────────────────────────
|
||||
|
||||
def _can_use_gpu_direct(self, layers: list[RenderLayer]) -> bool:
|
||||
"""判断是否命中直连支持的场景:单一主视频轨、全硬切、无复杂合成。"""
|
||||
try:
|
||||
cfg = self.plan.config or {}
|
||||
# 特性开关(默认开启;可经 env/plan config 关闭灰度回退)
|
||||
if not bool(cfg.get("gpu_direct_enabled", True)):
|
||||
return False
|
||||
|
||||
video_layers = [_lyr for _lyr in layers if _lyr.role not in ("audio",)]
|
||||
# 只允许一个视频层,且角色为主层
|
||||
if len(video_layers) != 1:
|
||||
return False
|
||||
role = video_layers[0].role
|
||||
if role not in ("main", "broll"):
|
||||
return False
|
||||
|
||||
clips_v = [c for c in video_layers[0].clips if c.clip_type != "audio"]
|
||||
if not clips_v:
|
||||
return False
|
||||
# 全硬切(第一个 clip 的转场忽略)
|
||||
for c in clips_v[1:]:
|
||||
te = c.transition_effect
|
||||
if te not in (None, "", "cut"):
|
||||
return False
|
||||
# 无画中画 / 水印 / 贴纸 / 片头片尾 / 绿幕 / 倒放 / 调色
|
||||
if (cfg or {}).get("pip_config"):
|
||||
return False
|
||||
if (cfg or {}).get("intro_outro"):
|
||||
return False
|
||||
for c in clips_v:
|
||||
cc = c.config or {}
|
||||
if cc.get("watermark") or cc.get("stickers") or cc.get("chroma_key"):
|
||||
return False
|
||||
if ReverseConfig.from_dict(cc.get("reverse")).enabled:
|
||||
return False
|
||||
cg = ColorGradeConfig.from_dict(cc.get("color_grade"))
|
||||
if cg.enabled and cg.has_effect():
|
||||
return False
|
||||
if not (c.config or {}).get("_storage_key"):
|
||||
return False
|
||||
return True
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] eligibility check failed (fallback)", exc_info=True)
|
||||
return False
|
||||
|
||||
def _try_gpu_direct(
|
||||
self,
|
||||
*,
|
||||
layers: list[RenderLayer],
|
||||
ass_path: Path | None,
|
||||
video_duration: float,
|
||||
output_path: Path,
|
||||
) -> tuple[bool, bool] | tuple[None, bool]:
|
||||
"""尝试全 GPU 直连渲染。成功返回 (True, edge_crop_applied),不支持/失败返回 (None, False)。"""
|
||||
if not self._can_use_gpu_direct(layers):
|
||||
return (None, False)
|
||||
if not self._gpu_encode_available():
|
||||
return (None, False)
|
||||
|
||||
try:
|
||||
from video_processing import gpu_direct_pipeline as gdp
|
||||
|
||||
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"]
|
||||
|
||||
# 音频层处理:收集 TTS 分段与配音素材库整段音频
|
||||
# - TTS 分段(带 tts 标记)→ 无间隙 concat 成单文件
|
||||
# - 配音素材库(voice_library=True)→ 单独作为整段音轨(不走分段 concat,已从 0 覆盖整段)
|
||||
audio_layer = next((_lyr for _lyr in layers if _lyr.role == "audio"), None)
|
||||
tts_merged: Path | None = None
|
||||
voiceover_track: Path | None = None
|
||||
if audio_layer:
|
||||
tts_clips = [c for c in audio_layer.clips if (c.config or {}).get("tts") and c.local_path.exists()]
|
||||
if tts_clips:
|
||||
tts_merged = self._concat_audio_clips(tts_clips, tag="tts_direct")
|
||||
# 配音素材库整段音频(按 _maybe_add_voice_library_layer 约定只有一个 clip_id=voice_library_main)
|
||||
vo_clips = [
|
||||
c for c in audio_layer.clips if (c.config or {}).get("voice_library") and c.local_path.exists()
|
||||
]
|
||||
if vo_clips:
|
||||
voiceover_track = vo_clips[-1].local_path # 理论上只有一个,取最后一个
|
||||
logger.info(
|
||||
"[gpu-direct] 配音素材库音轨: plan_id=%s path=%s",
|
||||
self.plan.id,
|
||||
voiceover_track,
|
||||
)
|
||||
|
||||
# 额外独立音轨(TTS concat、配音素材库)→ gpu_direct_pipeline 会与主音轨/BGM 一起 amix
|
||||
extra_audio_tracks: list[tuple[Path, float]] = []
|
||||
if tts_merged:
|
||||
extra_audio_tracks.append((tts_merged, 1.0))
|
||||
if voiceover_track:
|
||||
extra_audio_tracks.append((voiceover_track, 1.0))
|
||||
|
||||
# BGM 本地文件
|
||||
bgm_path = Path(self.bgm_path) if self.bgm_path else None
|
||||
if bgm_path is not None and not bgm_path.exists():
|
||||
bgm_path = None
|
||||
|
||||
# 字幕/标题/BGM 配置整包透传
|
||||
title_cfg = cfg.get("title", {}) or cfg.get("title_config", {}) or {}
|
||||
if not isinstance(title_cfg, dict):
|
||||
title_cfg = {}
|
||||
title_text = ""
|
||||
if title_cfg.get("enabled", True):
|
||||
title_text = title_cfg.get("text", "") or ""
|
||||
|
||||
sub_cfg = cfg.get("subtitle", {}) or {}
|
||||
if not isinstance(sub_cfg, dict):
|
||||
sub_cfg = {}
|
||||
subtitle_segments: list[Any] = []
|
||||
static_subtitle_text = ""
|
||||
if sub_cfg.get("enabled", True):
|
||||
if sub_cfg.get("auto_generated") and self._asr_timeline_cache is not None:
|
||||
subtitle_segments = list(self._asr_timeline_cache.segments)
|
||||
else:
|
||||
# 静态字幕文本(用户手输):pipeline 内部会构造全片长 segment
|
||||
static_subtitle_text = (sub_cfg.get("text", "") or "").strip()
|
||||
|
||||
bgm_cfg = cfg.get("bgm", {}) or {}
|
||||
if not isinstance(bgm_cfg, dict):
|
||||
bgm_cfg = {}
|
||||
# 若 bgm.enabled 显式关闭,则强制 bgm_path=None(_prepare_bgm 已按 enabled 返回 None,双保险)
|
||||
if not bgm_cfg.get("enabled", True):
|
||||
bgm_path = None
|
||||
# 注入微片段 BGM 偏移(同 CPU 路径)
|
||||
if bgm_path is not None and not bgm_cfg.get("audio_offset"):
|
||||
_micro_off = self._get_micro_bgm_offset()
|
||||
if _micro_off:
|
||||
bgm_cfg = {**bgm_cfg, "audio_offset": _micro_off}
|
||||
|
||||
# 边缘裁剪:dedup 开启时在 GPU 内做四边随机 2~5% 裁剪(gpu_direct_pipeline 内部随机)
|
||||
dedup = self._dedup_enabled()
|
||||
edge_pct = 0.03 if dedup else 0.0 # >0 表示启用;实际区间 [2%,5%] 在 pipeline 内随机
|
||||
|
||||
# 探测每个视频素材是否含音轨、读取 volume 配置
|
||||
clip_has_audio_list: list[bool] = []
|
||||
clip_volumes_list: list[float] = []
|
||||
for c in video_clips:
|
||||
lp = getattr(c, "local_path", None)
|
||||
_ha = False
|
||||
if lp and Path(lp).exists():
|
||||
try:
|
||||
_ha = probe_has_audio(str(lp))
|
||||
except Exception as _pe: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] probe_has_audio 失败按有声处理: %s", _pe)
|
||||
_ha = True
|
||||
clip_has_audio_list.append(_ha)
|
||||
_vol = float((c.config or {}).get("volume", 1.0))
|
||||
clip_volumes_list.append(_vol if _vol > 0 else 0.0)
|
||||
|
||||
# extra_audio_tracks 音量:从 audio_tracks_config 读(TTS/配音素材库),
|
||||
# 无法精确匹配 track_id 时保留默认 1.0
|
||||
at_cfg = cfg.get("audio_tracks") or {}
|
||||
tts_volume = 1.0
|
||||
vo_volume = 1.0
|
||||
if isinstance(at_cfg, dict):
|
||||
_tracks = at_cfg.get("tracks", []) or []
|
||||
for _t in _tracks:
|
||||
if not isinstance(_t, dict):
|
||||
continue
|
||||
try:
|
||||
_vol = float(_t.get("volume", 1.0))
|
||||
except (TypeError, ValueError):
|
||||
_vol = 1.0
|
||||
_tt = str(_t.get("track_type", ""))
|
||||
if _tt == "voiceover" and _t.get("audio_path"):
|
||||
vo_volume = max(0.0, min(2.0, _vol))
|
||||
# TTS 一般没有固定 track_type 标记,保持默认 1.0
|
||||
|
||||
extra_audio_tracks_cfg: list[tuple[Any, float]] = []
|
||||
if tts_merged:
|
||||
extra_audio_tracks_cfg.append((tts_merged, tts_volume))
|
||||
if voiceover_track:
|
||||
extra_audio_tracks_cfg.append((voiceover_track, vo_volume))
|
||||
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=video_clips,
|
||||
output_width=self.output_width,
|
||||
output_height=self.output_height,
|
||||
output_fps=self.output_fps,
|
||||
bgm_audio=bgm_path,
|
||||
title_text=title_text,
|
||||
subtitle_segments=subtitle_segments,
|
||||
edge_crop_pct=edge_pct,
|
||||
total_duration=video_duration,
|
||||
clip_has_audio=clip_has_audio_list,
|
||||
clip_volumes=clip_volumes_list,
|
||||
extra_audio_tracks=extra_audio_tracks_cfg,
|
||||
title_config=title_cfg,
|
||||
subtitle_config=sub_cfg,
|
||||
bgm_config=bgm_cfg,
|
||||
static_subtitle_text=static_subtitle_text,
|
||||
)
|
||||
|
||||
client = get_gpu_encoder()
|
||||
client.render_inputs_to_output(plan.inputs, plan.ffmpeg_args, output_path)
|
||||
|
||||
# 清理本次上传的临时音频
|
||||
for key in plan.oss_keys:
|
||||
try:
|
||||
from video_processing.oss_helpers import _storage
|
||||
|
||||
_storage().delete_file(key) if hasattr(_storage(), "delete_file") else None
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
did_edge_crop = bool(edge_pct)
|
||||
logger.info(
|
||||
"[gpu-direct] success: plan_id=%s clips=%d edge_crop=%s", self.plan.id, len(video_clips), did_edge_crop
|
||||
)
|
||||
return (True, did_edge_crop)
|
||||
|
||||
except GpuEncodeError as e:
|
||||
logger.warning("[gpu-direct] failed (fallback to legacy): %s", e)
|
||||
try:
|
||||
if output_path.exists():
|
||||
output_path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
return (None, False)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] unexpected error (fallback)", exc_info=True)
|
||||
return (None, False)
|
||||
|
||||
def _concat_audio_clips(self, clips: list[Any], *, tag: str) -> Path:
|
||||
"""把多个本地音频片段无间隙 concat 成一个 m4a(TTS 分段→单文件)。"""
|
||||
out = self.work_dir / f"{tag}_{self.plan.id}.m4a"
|
||||
listfile = self.work_dir / f"{tag}_{self.plan.id}.txt"
|
||||
lines = []
|
||||
for c in clips:
|
||||
ap = str(c.local_path).replace("'", "'\\''")
|
||||
lines.append(f"file '{ap}'")
|
||||
listfile.write_text("\n".join(lines), encoding="utf-8")
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
str(listfile),
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(out),
|
||||
]
|
||||
run_ffmpeg(cmd)
|
||||
return out
|
||||
|
||||
def _gpu_encode_available(self) -> bool:
|
||||
"""GPU 编码客户端是否已配置且健康(缓存健康状态,单任务内只探测一次)。"""
|
||||
if not getattr(self, "_gpu_health_ok", None):
|
||||
@@ -2594,7 +2906,7 @@ class UnifiedRenderService:
|
||||
b = pixel_pert.get("color_b", 0)
|
||||
if r != 0 or g != 0 or b != 0:
|
||||
# color_balance 参数范围 -1.0 ~ 1.0,这里用 /100 转换
|
||||
filters.append(f"colorbalance=rs={r/100:.3f}:gs={g/100:.3f}:bs={b/100:.3f}")
|
||||
filters.append(f"colorbalance=rs={r / 100:.3f}:gs={g / 100:.3f}:bs={b / 100:.3f}")
|
||||
|
||||
@staticmethod
|
||||
def _clip_volume(clip: ResolvedClip) -> float:
|
||||
|
||||
@@ -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%)
|
||||
|
||||
@@ -18,6 +18,7 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.classification import ClassificationStatus
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
@@ -96,7 +97,7 @@ def calculate_asset_quality_task(self, asset_id: str) -> dict:
|
||||
confidence = 1.0
|
||||
existing_meta["classification"] = classification
|
||||
existing_meta["classification_confidence"] = confidence
|
||||
asset.classification_status = "completed"
|
||||
asset.classification_status = ClassificationStatus.COMPLETED
|
||||
asset.metadata = existing_meta
|
||||
logger.info(
|
||||
"[quality_score] asset=%s 自动分类完成: category=%s confidence=%.2f",
|
||||
@@ -110,6 +111,8 @@ def calculate_asset_quality_task(self, asset_id: str) -> dict:
|
||||
asset_id,
|
||||
cls_err,
|
||||
)
|
||||
# 分类失败显式标记 FAILED,避免停留在 PENDING 被反复重试
|
||||
asset.classification_status = ClassificationStatus.FAILED
|
||||
|
||||
asset_repo.update(asset)
|
||||
db.commit()
|
||||
|
||||
@@ -198,7 +198,7 @@ def _verify_url_accessible(
|
||||
retries: int = 2,
|
||||
max_redirects: int = 5,
|
||||
) -> bool:
|
||||
"""HEAD 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。
|
||||
"""GET+Range 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。
|
||||
|
||||
安全增强:
|
||||
- 请求前先做 SSRF 安全校验(内网IP/回环地址/链路本地地址等)
|
||||
@@ -255,11 +255,12 @@ def _verify_url_accessible(
|
||||
)
|
||||
raise
|
||||
|
||||
req = urllib.request.Request(safe_url, method="HEAD")
|
||||
req = urllib.request.Request(safe_url, method="GET")
|
||||
req.add_header("Range", "bytes=0-0")
|
||||
req.add_header("User-Agent", "xiaoxia-saas-worker/1.0")
|
||||
|
||||
with opener.open(req, timeout=timeout) as resp: # noqa: S310
|
||||
if 200 <= resp.status < 300:
|
||||
if 200 <= resp.status < 300 or resp.status == 206:
|
||||
return True
|
||||
if resp.status in (301, 302, 303, 307, 308):
|
||||
location = resp.headers.get("Location", "")
|
||||
@@ -615,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):
|
||||
@@ -693,19 +674,24 @@ def _render_from_edit_plan(
|
||||
task_id: str,
|
||||
source_edit_plan_id: str,
|
||||
task_info: dict,
|
||||
) -> tuple[Path, float, list[dict] | None, str | None, str | None, str]:
|
||||
) -> 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)
|
||||
(output_path, render_duration, cover_candidates, voiceover_path, temp_dir, thumbnail_url, edge_crop_applied)
|
||||
"""
|
||||
from video_processing.render_adapter import RenderAdapter
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
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):
|
||||
@@ -721,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:
|
||||
@@ -746,6 +733,7 @@ def _render_from_edit_plan(
|
||||
voiceover_path,
|
||||
render_temp_dir,
|
||||
result.thumbnail_url or "",
|
||||
bool(getattr(result, "edge_crop_applied", False)),
|
||||
)
|
||||
finally:
|
||||
db.close()
|
||||
@@ -899,6 +887,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
voiceover_tmp_path,
|
||||
render_temp_dir,
|
||||
thumbnail_url,
|
||||
_gpu_edge_crop_done,
|
||||
) = _render_from_edit_plan(
|
||||
task_id=task_id,
|
||||
source_edit_plan_id=current_plan_id,
|
||||
@@ -935,6 +924,12 @@ def generate_video(self, task_id: str) -> dict:
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("降重", "已关闭边缘裁剪与微变换(确定性渲染)")
|
||||
_flush_logs(task_id, gen_task)
|
||||
elif _gpu_edge_crop_done:
|
||||
# GPU 直连管线已经在 filter_complex 中做了随机边缘裁剪,跳过 CPU 二次重编码
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("边缘裁剪", "已在 GPU 直连管线内完成随机边缘裁剪")
|
||||
_flush_logs(task_id, gen_task)
|
||||
logger.info("[task_id=%s] GPU直连已完成边缘裁剪,跳过CPU二次重编码", task_id)
|
||||
else:
|
||||
from video_processing.ffmpeg_utils import random_edge_crop
|
||||
|
||||
@@ -1051,20 +1046,15 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
if _meta_model:
|
||||
meta = dict(_meta_model.extra_meta or {})
|
||||
# #2024/P0 finalize-400: 直接展开 _precompute_render_metadata 返回的
|
||||
# 完整 dict(含 file_url/fingerprint_dict/fingerprint_chunks/is_duplicate
|
||||
# /duplicate_of/...),避免手写字段白名单漏传字段导致 finalize 读不到数据。
|
||||
meta["rendered_output"] = {
|
||||
**dict(rendered_output or {}),
|
||||
# file_url/duration 由外层调用方拿到的实际上传结果,优先覆盖预计算值
|
||||
"file_url": file_url,
|
||||
"file_size": file_size,
|
||||
"duration": duration,
|
||||
"width": rendered_output.get("width", 1280),
|
||||
"height": rendered_output.get("height", 720),
|
||||
"fps": rendered_output.get("fps", 25.0),
|
||||
"name": rendered_output.get("name", ""),
|
||||
"thumbnail_url": rendered_output.get("thumbnail_url", ""),
|
||||
"mode": rendered_output.get("mode", editing_mode.value),
|
||||
"fingerprint_dict": rendered_output.get("fingerprint_dict"),
|
||||
"batch_id": batch_id,
|
||||
"project_id": project_id,
|
||||
"user_id": user_id,
|
||||
}
|
||||
_meta_model.extra_meta = meta
|
||||
_finalize_meta_session.commit()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -280,3 +280,18 @@ USE_GPU_LIPSYNC=true
|
||||
GPU_LIPSYNC_POLL_INTERVAL=5
|
||||
GPU_LIPSYNC_WAIT_TIMEOUT=1200
|
||||
GPU_WORKER_STALE_SECONDS=300
|
||||
|
||||
# ==================== P4000 NVENC 硬件编码(GPU mezzanine relay)====================
|
||||
# 注意:这些值必须写死在模板里(不是 CI Secret),否则每次 CI 重新渲染 .env 都会被丢弃,
|
||||
# 导致 staging 发版后 GPU 编码静默降级到 CPU(P0 防复发)。
|
||||
ENABLE_GPU_ENCODE=true
|
||||
GPU_ENCODE_ENDPOINT=http://100.105.75.67:8900
|
||||
GPU_ENCODE_RELAY_BASE_URL=http://100.125.116.43:8092
|
||||
GPU_ENCODE_RELAY_INTERNAL_BASE_URL=http://xiaoxia-api-staging:8000
|
||||
GPU_ENCODE_RELAY_SECRET=0e1a8f0626438564a8b3fa92f3f2aac29e3c69bc02f2f85c
|
||||
GPU_ENCODE_VCODEC=h264_nvenc
|
||||
GPU_ENCODE_PRESET=p4
|
||||
GPU_ENCODE_CRF=23
|
||||
GPU_ENCODE_FALLBACK_CPU=true
|
||||
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
||||
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
# Host nginx config for staging server: /etc/nginx/sites-available/05-xiaoxia-cms
|
||||
# Xiaoxia CMS - cms.xiaoxiajianji.com
|
||||
#
|
||||
# 注意:此文件是宿主机 nginx 配置的备份/参考,不是 Docker 容器内的 nginx。
|
||||
# Docker 容器内的 nginx 配置见 nginx-staging.conf。
|
||||
#
|
||||
# GPU relay 路由说明:
|
||||
# P4000 编码完成后通过 http://100.69.73.60:8092/api/v1/internal/gpu-relay/{key} PUT 上传
|
||||
# Worker 容器通过 http://xiaoxia-api-staging:8000/api/v1/internal/gpu-relay/{key} GET 下载
|
||||
# 8092 端口由 CMS 宿主机 nginx 承载,GPU relay 路由通过最长前缀匹配优先代理到 staging API (8000)
|
||||
|
||||
# 80 端口:ACME 验证 + 重定向到 HTTPS
|
||||
server {
|
||||
listen 80;
|
||||
server_name cms.xiaoxiajianji.com;
|
||||
|
||||
# Let's Encrypt ACME 验证
|
||||
location /.well-known/acme-challenge/ {
|
||||
root /var/www/certbot;
|
||||
}
|
||||
|
||||
location / {
|
||||
return 301 https://$host$request_uri;
|
||||
}
|
||||
}
|
||||
|
||||
# 443 端口:CMS 主站
|
||||
server {
|
||||
listen 443 ssl http2;
|
||||
server_name cms.xiaoxiajianji.com;
|
||||
|
||||
ssl_certificate /etc/letsencrypt/live/cms.xiaoxiajianji.com/fullchain.pem;
|
||||
ssl_certificate_key /etc/letsencrypt/live/cms.xiaoxiajianji.com/privkey.pem;
|
||||
include /etc/letsencrypt/options-ssl-nginx.conf;
|
||||
ssl_dhparam /etc/letsencrypt/ssl-dhparams.pem;
|
||||
|
||||
client_max_body_size 50m;
|
||||
|
||||
# Security headers
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "strict-origin-when-cross-origin" always;
|
||||
add_header Strict-Transport-Security "max-age=31536000; includeSubDomains" always;
|
||||
|
||||
root /data/www/cms/current;
|
||||
index index.html;
|
||||
|
||||
gzip on;
|
||||
gzip_types text/plain text/css application/json application/javascript text/xml application/xml application/xml+rss text/javascript image/svg+xml;
|
||||
gzip_min_length 1024;
|
||||
|
||||
location / {
|
||||
try_files $uri $uri/ /index.html;
|
||||
}
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8091/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8091/health;
|
||||
}
|
||||
}
|
||||
|
||||
# 临时访问:8092 端口(IP直接访问,后续可关闭)
|
||||
server {
|
||||
listen 8092;
|
||||
server_name _;
|
||||
|
||||
root /data/www/cms/current;
|
||||
index index.html;
|
||||
|
||||
client_max_body_size 50m;
|
||||
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "strict-origin-when-cross-origin" always;
|
||||
add_header Strict-Transport-Security "max-age=31536000; includeSubDomains" always;
|
||||
|
||||
gzip on;
|
||||
gzip_types text/plain text/css application/json application/javascript text/xml application/xml application/xml+rss text/javascript image/svg+xml;
|
||||
gzip_min_length 1024;
|
||||
|
||||
location / {
|
||||
try_files $uri $uri/ /index.html;
|
||||
}
|
||||
|
||||
# GPU relay endpoints - proxy to staging API (port 8000) instead of CMS
|
||||
# 此 location 必须在 location /api/ 之前,利用 nginx 最长前缀匹配优先路由
|
||||
location /api/v1/internal/gpu-relay/ {
|
||||
proxy_pass http://127.0.0.1:8000/api/v1/internal/gpu-relay/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_request_buffering off;
|
||||
proxy_read_timeout 600s;
|
||||
proxy_send_timeout 600s;
|
||||
}
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8091/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8091/health;
|
||||
}
|
||||
}
|
||||
@@ -104,6 +104,43 @@ Staging 当前可以保持 no-op;Production 开启前必须先验证 SMTP/Redi
|
||||
|
||||
---
|
||||
|
||||
## Staging 服务器 Docker 凭证配置
|
||||
|
||||
Staging 服务器(116.62.226.203)需要配置 ACR 和 Gitea Registry 凭证,否则 docker pull 和 Watchtower 自动更新会失败。
|
||||
|
||||
### 凭证文件位置
|
||||
- Docker 配置文件:`/root/.docker/config.json`
|
||||
- 包含两个 registry 的认证信息:
|
||||
- `xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com`(阿里云 ACR)
|
||||
- `git.xiaoxiajianji.com`(Gitea 容器镜像仓库)
|
||||
|
||||
### 服务器迁移后恢复步骤
|
||||
```bash
|
||||
# 1. 登录 ACR
|
||||
docker login xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com -u <ACR_USERNAME>
|
||||
|
||||
# 2. 登录 Gitea Registry
|
||||
docker login git.xiaoxiajianji.com -u xiaoxia -p <GITEA_REGISTRY_TOKEN>
|
||||
|
||||
# 3. 重启 Watchtower(确保挂载最新 config.json)
|
||||
docker restart watchtower
|
||||
```
|
||||
|
||||
### Watchtower 配置
|
||||
- 容器名:`watchtower`
|
||||
- 检查间隔:300 秒(5 分钟)
|
||||
- 监控容器:`xiaoxia-api-staging`、`xiaoxia-worker-staging`、`xiaoxia-web-staging`
|
||||
- 必须挂载 `-v /root/.docker/config.json:/config.json` 才能拉取私有镜像
|
||||
- 必须挂载 `-v /var/run/docker.sock:/var/run/docker.sock` 才能管理容器
|
||||
- 容器使用 `:dev` 稳定 tag,Watchtower 通过检测 `:dev` tag 的 digest 变化来发现更新
|
||||
|
||||
### 镜像 Tag 策略
|
||||
- CI 每次构建推送三种 tag:`${GITHUB_SHA}`(精确版本)、`${GITHUB_REF_NAME}`(分支名)、`:dev`(滚动 tag,仅 develop 分支)
|
||||
- Staging 容器统一使用 `:dev` tag 启动,确保 Watchtower 能自动发现新版本
|
||||
- Migration(alembic)使用 commit SHA tag 执行,不依赖 Watchtower
|
||||
|
||||
---
|
||||
|
||||
## Gitea Actions 约定
|
||||
|
||||
- `develop` 分支触发 staging 部署。
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# 预设 BGM 音频文件放置目录
|
||||
|
||||
将下列 10 首免费可商用 BGM 的 mp3/m4a 文件按 `{preset_id}.mp3` 命名放到本目录:
|
||||
|
||||
| preset_id | 名称 | 风格 | 时长(s) | 标签 |
|
||||
|-------------------|----------|----------|---------|----------------------------|
|
||||
| bgm_upbeat_001 | 阳光清晨 | upbeat | 120 | 轻快 阳光 吉他 vlog |
|
||||
| bgm_upbeat_002 | 活力节拍 | upbeat | 95 | 轻快 电子 活力 运动 |
|
||||
| bgm_upbeat_003 | 夏日漫步 | upbeat | 110 | 轻快 夏日 ukulele 旅行 |
|
||||
| bgm_relax_001 | 静谧时光 | relax | 180 | 治愈 钢琴 安静 冥想 |
|
||||
| bgm_relax_002 | 雨后森林 | relax | 150 | 治愈 自然 放松 环境音 |
|
||||
| bgm_relax_003 | 月光奏鸣曲 | relax | 200 | 治愈 古典 钢琴 优雅(公版) |
|
||||
| bgm_tech_001 | 未来科技 | tech | 85 | 科技 电子 未来感 数码 |
|
||||
| bgm_tech_002 | 数据脉冲 | tech | 100 | 科技 极简 数据 AI |
|
||||
| bgm_commerce_001 | 心动时刻 | commerce | 75 | 电商 时尚 动感 带货 |
|
||||
| bgm_commerce_002 | 品质生活 | commerce | 90 | 电商 高端 品牌 品质 |
|
||||
|
||||
放好后执行(需要在有 OSS 凭证的机器上):
|
||||
```bash
|
||||
export OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com
|
||||
export OSS_ACCESS_KEY_ID=xxx
|
||||
export OSS_ACCESS_KEY_SECRET=xxx
|
||||
export OSS_BUCKET_NAME=xiaoxia-autocut
|
||||
python scripts/upload_preset_bgm.py
|
||||
```
|
||||
|
||||
脚本会:
|
||||
1. 上传文件到 OSS `preset/bgm/<preset_id>.mp3`,设置公共读 ACL
|
||||
2. 自动改写 `packages/domain/preset_bgm.py` 把对应 `audio_url=""` 回填成公网 URL
|
||||
3. 提示 `git commit & push`
|
||||
@@ -30,6 +30,10 @@ COPY deploy/configs/douyin_cookies.txt /app/configs/douyin_cookies.txt
|
||||
# 强制升级 yt-dlp 到最新(抖音反爬经常变更,旧版 cookies 支持失效;#1968/#1963)
|
||||
RUN pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com --upgrade "yt-dlp>=2026.8.19"
|
||||
|
||||
# API 启动入口(幂等迁移 + uvicorn)—— #2129: watchtower 自动部署兜底
|
||||
COPY infra/docker/entrypoint-api.sh /usr/local/bin/entrypoint-api.sh
|
||||
RUN chmod +x /usr/local/bin/entrypoint-api.sh
|
||||
|
||||
# 设置环境变量
|
||||
ENV PATH="/opt/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
|
||||
ENV PYTHONPATH=/app:/app/apps/api
|
||||
@@ -41,4 +45,4 @@ HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"
|
||||
|
||||
# API 入口点
|
||||
CMD ["uvicorn", "apps.api.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||
ENTRYPOINT ["/usr/local/bin/entrypoint-api.sh"]
|
||||
|
||||
+44
-72
@@ -10,20 +10,25 @@
|
||||
# API_IMAGE - API 镜像名称 (默认: xiaoxia-saas-api:dev)
|
||||
# WORKER_IMAGE - Worker 镜像名称 (默认: xiaoxia-saas-worker:dev)
|
||||
# WEB_IMAGE - Web 镜像名称 (默认: xiaoxia-saas-web:dev)
|
||||
# WEB_DOCKERFILE - Web Dockerfile 路径
|
||||
# WEB_NGINX_CONF - Nginx 配置文件路径
|
||||
# API_PORT - API 端口映射 (staging: 8000, production: 8001)
|
||||
# WEB_PORT - Web 端口映射 (staging: 3001, production: 3002)
|
||||
# GENERATED_FILES_HOST_DIR - 生成文件的主机目录
|
||||
# WORKER_CONCURRENCY - Worker 并发数 (默认: 4)
|
||||
# GENERATION_CONCURRENCY - Generation worker 并发(用户实时任务,默认 2)
|
||||
# TRANSCODE_CONCURRENCY - Transcode worker 并发(后台/转码/AI,默认 2)
|
||||
# WORKER_MAX_TASKS_PER_CHILD - Worker 每个子进程最大任务数 (默认: 100)
|
||||
# BEAT_ENABLED - 容器内启动 celery beat(默认 1;独立 beat 容器部署设为 0)
|
||||
# WORKER_CONCURRENCY - 兼容旧变量:未显式设置上面两个并发时按此总数分配
|
||||
#
|
||||
# 重要:
|
||||
# 重要:
|
||||
# - 生产环境不要挂载 web-dist volume,否则会导致 403
|
||||
# - 确保环境隔离网络已创建: docker network create xiaoxia-net-${ENV}
|
||||
# - ENV=staging → xiaoxia-net-staging
|
||||
# - ENV=production → xiaoxia-net-production
|
||||
# - ENV=staging -> xiaoxia-net-staging
|
||||
# - ENV=production -> xiaoxia-net-production
|
||||
#
|
||||
# #2073 队列分流:worker 容器内跑三个独立进程——beat(只发定时任务)、
|
||||
# generation worker(只消费 generation 队列,实时高优)、transcode worker(消费
|
||||
# transcode + celery 队列,后台任务)。beat 不再嵌入 generation worker,
|
||||
# 不占实时任务槽位;TRANSCODE_CONCURRENCY 独立伸缩,不再依赖 WORKER_CONCURRENCY 差值。
|
||||
|
||||
# ===========================================
|
||||
# 日志轮转配置(所有服务共享)
|
||||
@@ -40,53 +45,39 @@ services:
|
||||
# =========================================
|
||||
api:
|
||||
image: ${API_IMAGE:-xiaoxia-saas-api:dev}
|
||||
# 不在生产环境构建镜像,使用预构建的镜像
|
||||
# build:
|
||||
# context: ../..
|
||||
# dockerfile: infra/docker/api.Dockerfile
|
||||
|
||||
container_name: xiaoxia-api-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
stop_grace_period: 30s
|
||||
stop_signal: SIGTERM
|
||||
|
||||
# 环境变量文件(包含数据库密码等敏感信息)
|
||||
|
||||
env_file:
|
||||
- ../../.env
|
||||
|
||||
|
||||
environment:
|
||||
APP_ENV: ${APP_ENV:-staging}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
|
||||
# 端口映射
|
||||
# Staging: 8000 -> 8000
|
||||
# Production: 8001 -> 8000
|
||||
|
||||
ports:
|
||||
- "127.0.0.1:${API_PORT:-8000}:8000"
|
||||
|
||||
# 共享生成文件目录 + 抖音 cookies 等运行时配置
|
||||
|
||||
volumes:
|
||||
- generated-files:/app/generated
|
||||
- ../../deploy/configs:/app/configs:ro
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# 健康检查配置
|
||||
|
||||
healthcheck:
|
||||
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
start_period: 40s
|
||||
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -97,39 +88,48 @@ services:
|
||||
memory: 512M
|
||||
|
||||
# =========================================
|
||||
# Worker 服务(Celery 任务队列)
|
||||
# Worker 服务(#2073 队列分流:beat + generation + transcode 同容器三进程)
|
||||
# =========================================
|
||||
# 三个进程独立启动,任一退出则容器整体退出由 docker restart 拉起;
|
||||
# 各自的并发与资源占用通过环境变量控制:
|
||||
# - generation:GENERATION_CONCURRENCY(默认 2),消费 generation 队列
|
||||
# - transcode: TRANSCODE_CONCURRENCY(默认 2),消费 transcode,celery 队列
|
||||
# - beat: 不消费任务,只发定时任务到 celery 默认队列
|
||||
worker:
|
||||
image: ${WORKER_IMAGE:-xiaoxia-saas-worker:dev}
|
||||
|
||||
|
||||
container_name: xiaoxia-worker-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
# 长任务(ingest HEVC 转码最长 30min、生成硬超时 11min)给足优雅关闭窗口
|
||||
stop_grace_period: 300s
|
||||
stop_signal: SIGTERM
|
||||
|
||||
|
||||
env_file:
|
||||
- ../../.env
|
||||
|
||||
|
||||
environment:
|
||||
APP_ENV: ${APP_ENV:-staging}
|
||||
# 兼容旧变量:若两个 *_CONCURRENCY 均未显式设置,entrypoint 会按此总数分配
|
||||
WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4}
|
||||
WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100}
|
||||
# #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码
|
||||
# #2073 队列独立伸缩:generation 默认 2,transcode 默认 2(不再差值计算)
|
||||
GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2}
|
||||
TRANSCODE_CONCURRENCY: ${TRANSCODE_CONCURRENCY:-2}
|
||||
# beat 默认在本容器启动;独立 beat 容器部署时设为 0
|
||||
BEAT_ENABLED: ${BEAT_ENABLED:-1}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
|
||||
|
||||
volumes:
|
||||
- generated-files:/app/generated
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# 健康检查配置
|
||||
# 注:容器内无 pgrep/ps,扫描 /proc 所有进程的 cmdline 查找 celery 进程
|
||||
# 健康检查:至少有一个 celery worker 进程在跑(beat 本身不作为存活依据)
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "grep -lq celery /proc/[0-9]*/cmdline 2>/dev/null || exit 1"]
|
||||
test: ["CMD-SHELL", "grep -q 'celery.*worker' /proc/[0-9]*/cmdline 2>/dev/null || exit 1"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
@@ -137,12 +137,8 @@ services:
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
# 注意: Worker 需要处理视频,建议分配更多资源
|
||||
# #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
|
||||
# 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
|
||||
# 资源限制:容器总资源 = gen + trans + beat,按 2+2 并发场景建议 4C8G;
|
||||
# 后续如需独立扩容/重启,可拆为 worker-generation / worker-transcode / worker-beat 三个 service。
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -157,35 +153,21 @@ services:
|
||||
# =========================================
|
||||
web:
|
||||
image: ${WEB_IMAGE:-xiaoxia-saas-web:dev}
|
||||
# 不在生产环境构建镜像,使用 web-artifact.Dockerfile
|
||||
# build:
|
||||
# context: ../..
|
||||
# dockerfile: ${WEB_DOCKERFILE:-infra/docker/web.Dockerfile}
|
||||
# args:
|
||||
# (NGINX_CONF no longer needed - all configs baked into image)
|
||||
|
||||
|
||||
container_name: xiaoxia-web-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
|
||||
# 端口映射
|
||||
# Staging: 3001 -> 80
|
||||
# Production: 3002 -> 80 (通过 Nginx 反向代理)
|
||||
|
||||
ports:
|
||||
- "127.0.0.1:${WEB_PORT:-3001}:80"
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# =========================================
|
||||
# Nginx 配置运行时覆盖(双保险:entrypoint 也按 APP_ENV 选择配置)
|
||||
# 确保容器使用正确环境的 nginx 配置,即使镜像构建时使用了默认配置
|
||||
# 注意: 只覆盖 /etc/nginx/conf.d/default.conf,不挂载 /usr/share/nginx/html
|
||||
# =========================================
|
||||
|
||||
environment:
|
||||
- APP_ENV=${ENV:-staging}
|
||||
volumes:
|
||||
- ./nginx-${ENV:-staging}.conf:/etc/nginx/conf.d/default.conf:ro
|
||||
|
||||
|
||||
healthcheck:
|
||||
test: ["CMD", "wget", "--spider", "-q", "http://127.0.0.1:80"]
|
||||
interval: 30s
|
||||
@@ -194,9 +176,6 @@ services:
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议
|
||||
# =========================================
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -212,9 +191,6 @@ volumes:
|
||||
driver_opts:
|
||||
type: none
|
||||
o: bind
|
||||
# 重要: 确保主机目录存在且有正确权限
|
||||
# Staging: /var/lib/xiaoxia-saas-staging/generated
|
||||
# Production: /var/lib/xiaoxia-saas-production/generated
|
||||
device: ${GENERATED_FILES_HOST_DIR:?GENERATED_FILES_HOST_DIR must be set in .env}
|
||||
|
||||
# ===========================================
|
||||
@@ -223,8 +199,4 @@ volumes:
|
||||
networks:
|
||||
xiaoxia-net:
|
||||
external: true
|
||||
# 网络名根据 ENV 变量区分,实现 staging/production 环境隔离
|
||||
# staging: xiaoxia-net-staging
|
||||
# production: xiaoxia-net-production
|
||||
name: xiaoxia-net-${ENV:-staging}
|
||||
|
||||
|
||||
Executable
+22
@@ -0,0 +1,22 @@
|
||||
#!/bin/bash
|
||||
# API 启动入口:先幂等执行数据库迁移,再启动传入的 CMD(默认 uvicorn)
|
||||
# 解决 watchtower 自动拉取新镜像后容器重启、未跑 alembic upgrade head 导致新列缺失 500 的问题(#2129)
|
||||
set -e
|
||||
|
||||
cd /app
|
||||
|
||||
echo "[entrypoint-api] Running alembic upgrade head..."
|
||||
if alembic upgrade head; then
|
||||
echo "[entrypoint-api] Migrations ok."
|
||||
else
|
||||
echo "[entrypoint-api] WARNING: alembic upgrade failed, continuing (existing columns should be fine)..." >&2
|
||||
fi
|
||||
|
||||
# 若有显式 CMD(CI 部署时 docker compose run --rm api sh -c '...' 传入),直接 exec 它
|
||||
if [ "$#" -gt 0 ]; then
|
||||
echo "[entrypoint-api] Exec custom command: $*"
|
||||
exec "$@"
|
||||
fi
|
||||
|
||||
echo "[entrypoint-api] Starting uvicorn..."
|
||||
exec uvicorn apps.api.main:app --host 0.0.0.0 --port 8000
|
||||
@@ -1,48 +1,88 @@
|
||||
#!/bin/bash
|
||||
# Worker 启动脚本 — #1714 队列隔离
|
||||
# Worker 启动脚本 — #1714 + #2073 队列分流
|
||||
#
|
||||
# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程:
|
||||
# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级),
|
||||
# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行;
|
||||
# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/
|
||||
# 配音/下载等后台任务)。
|
||||
# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。
|
||||
# 容器内启动三个独立进程(任一退出则整体退出由 docker restart 拉起):
|
||||
# 1. beat:celery beat 调度器,不消费任何任务,只发定时任务到 celery 默认队列
|
||||
# 2. generation-worker:独占消费 generation 队列(用户实时任务,高优先级)
|
||||
# 3. transcode-worker:消费 transcode + celery 默认队列(后台/清理任务)
|
||||
#
|
||||
# 环境变量:
|
||||
# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2,
|
||||
# 可用 GENERATION_CONCURRENCY 覆盖
|
||||
# GENERATION_CONCURRENCY generation worker 并发(默认 2)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 2)
|
||||
# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100)
|
||||
# WORKER_CONCURRENCY 兼容旧变量:若未显式设置 GENERATION_CONCURRENCY /
|
||||
# TRANSCODE_CONCURRENCY,则按比例分配(gen=ceil(total*1/2),
|
||||
# trans=剩余,各至少 1);已显式设置时忽略此变量。
|
||||
# BEAT_ENABLED 是否在本容器内启动 beat 进程(默认 1);
|
||||
# 若独立 beat 容器部署设为 0。
|
||||
|
||||
set -e
|
||||
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-4}"
|
||||
# #2129: 幂等执行数据库迁移(watchtower 自动部署兜底)
|
||||
# worker 容器独立启动,不能依赖 API 容器先跑迁移
|
||||
cd /app
|
||||
echo "[entrypoint-worker] Running alembic upgrade head..."
|
||||
if alembic upgrade head; then
|
||||
echo "[entrypoint-worker] Migrations ok."
|
||||
else
|
||||
echo "[entrypoint-worker] WARNING: alembic upgrade failed, continuing to start workers..." >&2
|
||||
fi
|
||||
cd - >/dev/null
|
||||
|
||||
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
|
||||
|
||||
GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}"
|
||||
if [ -z "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then
|
||||
TRANS_CONCURRENCY=1
|
||||
fi
|
||||
# ── 并发计算:显式 env 优先;否则从 WORKER_CONCURRENCY 按比例推导 ──
|
||||
if [ -n "$GENERATION_CONCURRENCY" ]; then
|
||||
GEN_CONCURRENCY="$GENERATION_CONCURRENCY"
|
||||
else
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
TOTAL="${WORKER_CONCURRENCY:-4}"
|
||||
GEN_CONCURRENCY=$(( (TOTAL + 1) / 2 ))
|
||||
if [ "$GEN_CONCURRENCY" -lt 1 ]; then GEN_CONCURRENCY=1; fi
|
||||
fi
|
||||
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)"
|
||||
if [ -n "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
else
|
||||
if [ -n "$WORKER_CONCURRENCY" ] && [ -z "$GENERATION_CONCURRENCY" ]; then
|
||||
# 两个都没显式设置,按 WORKER_CONCURRENCY 分配剩余
|
||||
TOTAL="$WORKER_CONCURRENCY"
|
||||
TRANS_CONCURRENCY=$(( TOTAL - GEN_CONCURRENCY ))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then TRANS_CONCURRENCY=1; fi
|
||||
else
|
||||
# 默认 2(#2073:独立伸缩,不再依赖 WORKER_CONCURRENCY 差值)
|
||||
TRANS_CONCURRENCY=2
|
||||
fi
|
||||
fi
|
||||
|
||||
BEAT_ENABLED="${BEAT_ENABLED:-1}"
|
||||
|
||||
PIDS=()
|
||||
|
||||
# ── 1. Beat 调度器(独立进程,不消费任务)──
|
||||
if [ "$BEAT_ENABLED" = "1" ] || [ "$BEAT_ENABLED" = "true" ]; then
|
||||
echo "Starting beat scheduler (schedule file=/tmp/celerybeat-schedule)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
beat \
|
||||
--loglevel=info \
|
||||
-s /tmp/celerybeat-schedule &
|
||||
PIDS+=($!)
|
||||
fi
|
||||
|
||||
# ── 2. Generation worker(实时高优队列)──
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
worker \
|
||||
--loglevel=info \
|
||||
"-B" \
|
||||
-s /tmp/celerybeat-schedule \
|
||||
-Q generation \
|
||||
"--concurrency=${GEN_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n generation@%h &
|
||||
GEN_PID=$!
|
||||
PIDS+=($!)
|
||||
GEN_PID=${PIDS[1]:-${PIDS[0]}}
|
||||
|
||||
# ── 3. Transcode worker(后台 + 清理队列)──
|
||||
echo "Starting transcode worker (queues=transcode,celery, concurrency=$TRANS_CONCURRENCY)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
@@ -52,13 +92,22 @@ celery \
|
||||
"--concurrency=${TRANS_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n transcode@%h &
|
||||
TRANS_PID=$!
|
||||
PIDS+=($!)
|
||||
TRANS_PID=${PIDS[2]:-${PIDS[1]}}
|
||||
|
||||
# 任一进程退出则终止另一个,让容器整体重启(restart: unless-stopped)
|
||||
trap 'echo "Shutting down workers..."; kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true' TERM INT
|
||||
# 任一进程退出则终止其他进程,让容器整体重启
|
||||
cleanup() {
|
||||
echo "Shutting down all celery processes..."
|
||||
for pid in "${PIDS[@]}"; do
|
||||
kill -TERM "$pid" 2>/dev/null || true
|
||||
done
|
||||
}
|
||||
trap cleanup TERM INT
|
||||
|
||||
wait -n $GEN_PID $TRANS_PID
|
||||
# wait -n 等待任意一个子进程退出(bash 4.3+)
|
||||
# 容器镜像基础为 python:3.11-slim,bash 版本满足
|
||||
wait -n "${PIDS[@]}"
|
||||
EXIT_CODE=$?
|
||||
echo "One worker exited (code=$EXIT_CODE), stopping the other..."
|
||||
kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true
|
||||
exit $EXIT_CODE
|
||||
echo "One celery process exited (code=$EXIT_CODE), stopping the rest..."
|
||||
cleanup
|
||||
exit "$EXIT_CODE"
|
||||
|
||||
@@ -10,6 +10,14 @@ FROM xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji/saas-worker-base:l
|
||||
# 构建参数:版本号(CI 传入 commit hash)
|
||||
ARG APP_VERSION=dev
|
||||
|
||||
# CJK 字体保障:确保 fonts-noto-cjk 已安装(base 镜像漂移兜底)+ 重建字体缓存
|
||||
# fc-cache 非致命;fc-match 结果只打日志用于排查,不阻断构建
|
||||
RUN apt-get update && (apt-get install -y --no-install-recommends fonts-noto-cjk fontconfig || true) \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& (fc-cache -fv || true) \
|
||||
&& echo "[font] fc-match sans:zh: $(fc-match -f '%{family}\n' sans:zh 2>/dev/null | head -1)" \
|
||||
&& echo "[font] fc-match Noto Sans CJK SC: $(fc-match 'Noto Sans CJK SC' 2>/dev/null | head -1)"
|
||||
|
||||
# 创建非 root 用户
|
||||
RUN groupadd -r celery \
|
||||
&& useradd -r -g celery -d /app -s /sbin/nologin celery \
|
||||
@@ -26,6 +34,7 @@ ENV APP_VERSION=$APP_VERSION
|
||||
|
||||
# 复制文件(按变化频率从低到高排序,最大化层缓存命中)
|
||||
COPY alembic.ini /app/alembic.ini
|
||||
COPY alembic/ /app/alembic/
|
||||
COPY migrations/ /app/migrations/
|
||||
COPY packages/ /app/packages/
|
||||
# PR #1844 起,worker 还需要加载 apps.api.app.tasks.lipsync_tts,
|
||||
|
||||
@@ -1,31 +1,29 @@
|
||||
# Staging GPU relay plain-HTTP vhost (P4000 NVENC 编码回传入口)
|
||||
# - 监听 8092 端口纯 HTTP(绕开 HTTPS 证书与 P4000 httpx SSL 问题)
|
||||
# - 代理到本机 staging API 的 /api/ 路径(127.0.0.1:8000 是 docker 映射端口)
|
||||
# - P4000 通过 Tailscale 直连宿主机 100.69.73.60:8092 PUT 编码结果
|
||||
# - Worker 通过 Docker DNS (xiaoxia-api-staging:8000) 直接 GET/DELETE,
|
||||
# 不经宿主机 nginx,避免 UFW FORWARD DROP 阻断
|
||||
# Staging GPU relay nginx 配置说明
|
||||
#
|
||||
# 部署:cp infra/nginx/gpu-relay-staging.conf /etc/nginx/conf.d/ && nginx -t && systemctl reload nginx
|
||||
|
||||
server {
|
||||
listen 8092;
|
||||
server_name _;
|
||||
|
||||
client_max_body_size 2048m;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8000/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_request_buffering off;
|
||||
proxy_read_timeout 600s;
|
||||
proxy_send_timeout 600s;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8000/health;
|
||||
}
|
||||
}
|
||||
# GPU relay 并没有独立的 nginx vhost,而是集成在宿主机 CMS nginx 的 8092 server block 中。
|
||||
# 完整宿主机 nginx 配置备份见: deploy/configs/host-nginx-cms-staging.conf
|
||||
#
|
||||
# 核心 location 块(添加到 8092 server block,位于 location /api/ 之前):
|
||||
#
|
||||
# # GPU relay endpoints - proxy to staging API (port 8000) instead of CMS
|
||||
# location /api/v1/internal/gpu-relay/ {
|
||||
# proxy_pass http://127.0.0.1:8000/api/v1/internal/gpu-relay/;
|
||||
# proxy_http_version 1.1;
|
||||
# proxy_set_header Host $host;
|
||||
# proxy_set_header X-Real-IP $remote_addr;
|
||||
# proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
# proxy_set_header X-Forwarded-Proto $scheme;
|
||||
# proxy_request_buffering off;
|
||||
# proxy_read_timeout 600s;
|
||||
# proxy_send_timeout 600s;
|
||||
# }
|
||||
#
|
||||
# 部署方式:手动将上述 location 块添加到 /etc/nginx/sites-available/05-xiaoxia-cms 的 8092 server block 中
|
||||
# 然后 nginx -t && systemctl reload nginx
|
||||
#
|
||||
# 原理说明:
|
||||
# - 8092 端口由 CMS 宿主机 nginx 承载(与 CMS 共享端口)
|
||||
# - GPU relay 路由 /api/v1/internal/gpu-relay/ 比 CMS 的 /api/ 更具体
|
||||
# - nginx 最长前缀匹配确保 relay 请求路由到 staging API (port 8000) 而非 CMS (port 8091)
|
||||
# - P4000 通过 Tailscale IP 100.69.73.60:8092 访问 relay
|
||||
# - Worker 容器通过 Docker DNS xiaoxia-api-staging:8000 直接访问 relay
|
||||
|
||||
@@ -14,6 +14,24 @@ class InMemoryIngestJobRepository:
|
||||
def get(self, job_id: str) -> IngestJob | None:
|
||||
return self._items.get(job_id)
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
if not asset_id:
|
||||
return None
|
||||
from packages.domain.classification import IngestJobStatus
|
||||
|
||||
running: IngestJob | None = None
|
||||
completed: IngestJob | None = None
|
||||
for job in self._items.values():
|
||||
if getattr(job, "asset_id", "") != asset_id:
|
||||
continue
|
||||
if job.status in (IngestJobStatus.PENDING, IngestJobStatus.PROCESSING):
|
||||
if running is None or job.created_at > running.created_at:
|
||||
running = job
|
||||
elif job.status == IngestJobStatus.COMPLETED:
|
||||
if completed is None or job.created_at > completed.created_at:
|
||||
completed = job
|
||||
return running or completed
|
||||
|
||||
def update(self, job: IngestJob) -> IngestJob:
|
||||
self._items[job.id] = job
|
||||
return job
|
||||
|
||||
@@ -128,8 +128,12 @@ class SQLAlchemyAssetRepository:
|
||||
height=asset.height,
|
||||
fps=asset.fps,
|
||||
codec=asset.codec,
|
||||
status=asset.status.value,
|
||||
classification_status=asset.classification_status.value,
|
||||
status=(asset.status.value if hasattr(asset.status, "value") else str(asset.status)),
|
||||
classification_status=(
|
||||
asset.classification_status.value
|
||||
if hasattr(asset.classification_status, "value")
|
||||
else str(asset.classification_status)
|
||||
),
|
||||
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
|
||||
quality_score=asset.quality_score,
|
||||
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
|
||||
@@ -142,7 +146,7 @@ class SQLAlchemyAssetRepository:
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
# Issue #1776: 自动维护素材库计数(同事务内原子更新)
|
||||
if asset.library_id and asset.status.value != "deleted":
|
||||
if asset.library_id and (getattr(asset.status, "value", str(asset.status)) != "deleted"):
|
||||
from sqlalchemy import func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == asset.library_id).update(
|
||||
@@ -168,8 +172,12 @@ class SQLAlchemyAssetRepository:
|
||||
model.height = asset.height
|
||||
model.fps = asset.fps
|
||||
model.codec = asset.codec
|
||||
model.status = asset.status.value
|
||||
model.classification_status = asset.classification_status.value
|
||||
model.status = asset.status.value if hasattr(asset.status, "value") else str(asset.status)
|
||||
model.classification_status = (
|
||||
asset.classification_status.value
|
||||
if hasattr(asset.classification_status, "value")
|
||||
else str(asset.classification_status)
|
||||
)
|
||||
model.classification_result = json.dumps(asset.metadata) if asset.metadata else None
|
||||
model.quality_score = asset.quality_score
|
||||
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
|
||||
@@ -494,8 +502,19 @@ class SQLAlchemyAssetRepository:
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key(对应 DB 中的 file_url)查找素材。"""
|
||||
model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first()
|
||||
"""按 storage_key 查找素材。
|
||||
|
||||
Bug #2110: 历史数据 file_url 列可能是旧路径(assets/...),新代码统一写入
|
||||
storage_key 列。双列 OR 查询,避免占位 asset 因路径错配导致 ingest 兜底新建
|
||||
第二条 READY 记录,原占位卡 PROCESSING → 前端缩略图出现后消失。
|
||||
"""
|
||||
if not storage_key:
|
||||
return None
|
||||
model = (
|
||||
self.session.query(AssetModel)
|
||||
.filter((AssetModel.storage_key == storage_key) | (AssetModel.file_url == storage_key))
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
@@ -43,6 +43,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
output_height=getattr(model, "output_height", 720) or 720,
|
||||
cover_url=getattr(model, "cover_url", "") or "",
|
||||
title_config=dict(getattr(model, "title_config", {}) or {}),
|
||||
extra_meta=dict(getattr(model, "extra_meta", {}) or {}),
|
||||
logs=model.logs or "[]",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
@@ -88,6 +89,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
output_height=task.output_height,
|
||||
cover_url=task.cover_url or "",
|
||||
title_config=dict(task.title_config) if task.title_config else {},
|
||||
extra_meta=dict(task.extra_meta) if task.extra_meta else {},
|
||||
logs=task.logs,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
@@ -322,6 +324,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
model.output_height = task.output_height
|
||||
model.cover_url = task.cover_url or ""
|
||||
model.title_config = dict(task.title_config) if task.title_config else {}
|
||||
model.extra_meta = dict(task.extra_meta) if task.extra_meta else {}
|
||||
model.logs = task.logs
|
||||
self.session.commit()
|
||||
return task
|
||||
|
||||
@@ -46,6 +46,43 @@ class SQLAlchemyIngestJobRepository:
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
"""返回 asset 最近一条未失败的 ingest job(PENDING/PROCESSING/COMPLETED 均算存在,用于幂等判断)。"""
|
||||
if not asset_id:
|
||||
return None
|
||||
# 优先返回仍在跑的 (PENDING/PROCESSING),否则返回最新一条 COMPLETED
|
||||
model = (
|
||||
self.session.query(IngestJobModel)
|
||||
.filter(IngestJobModel.asset_id == asset_id)
|
||||
.filter(IngestJobModel.status.in_([IngestJobStatus.PENDING.value, IngestJobStatus.PROCESSING.value]))
|
||||
.order_by(IngestJobModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
model = (
|
||||
self.session.query(IngestJobModel)
|
||||
.filter(IngestJobModel.asset_id == asset_id)
|
||||
.filter(IngestJobModel.status == IngestJobStatus.COMPLETED.value)
|
||||
.order_by(IngestJobModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return IngestJob(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
library_id=model.library_id,
|
||||
storage_key=model.storage_key,
|
||||
status=IngestJobStatus(model.status),
|
||||
error_message=model.error_message,
|
||||
result_asset_id=model.result_asset_id,
|
||||
file_hash=model.file_hash or "",
|
||||
asset_id=getattr(model, "asset_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[IngestJob]:
|
||||
models = self.session.query(IngestJobModel).filter(IngestJobModel.project_id == project_id).all()
|
||||
return [self.get(model.id) for model in models if self.get(model.id) is not None]
|
||||
|
||||
@@ -56,7 +56,7 @@ class UserModel(Base):
|
||||
is_member = Column(Boolean, nullable=False, default=False)
|
||||
member_type = Column(String(20), nullable=True)
|
||||
member_expires_at = Column(DateTime, nullable=True)
|
||||
points_balance = Column(Integer, nullable=False, default=0)
|
||||
points_balance = Column(Float, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
@@ -775,9 +775,9 @@ class PointsAccountModel(Base):
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, unique=True, index=True)
|
||||
balance = Column(Integer, nullable=False, default=0)
|
||||
total_earned = Column(Integer, nullable=False, default=0)
|
||||
total_spent = Column(Integer, nullable=False, default=0)
|
||||
balance = Column(Float, nullable=False, default=0)
|
||||
total_earned = Column(Float, nullable=False, default=0)
|
||||
total_spent = Column(Float, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -792,8 +792,8 @@ class PointsTransactionModel(Base):
|
||||
account_id = Column(String(36), nullable=False, index=True)
|
||||
type = Column(String(20), nullable=False, index=True) # earn / spend / refund
|
||||
source = Column(String(50), nullable=False, index=True)
|
||||
amount = Column(Integer, nullable=False)
|
||||
balance_after = Column(Integer, nullable=False)
|
||||
amount = Column(Float, nullable=False)
|
||||
balance_after = Column(Float, nullable=False)
|
||||
description = Column(String(255), nullable=False, default="")
|
||||
ref_id = Column(String(100), nullable=False, default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
@@ -918,3 +918,88 @@ 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)
|
||||
# v1.5 音频/视频参数
|
||||
voice_id = Column(String(200), nullable=False, default="")
|
||||
voice_source = Column(String(20), nullable=False, default="")
|
||||
video_ratio = Column(String(10), nullable=False, default="9:16")
|
||||
video_model = Column(String(100), nullable=False, default="")
|
||||
# 结果与状态
|
||||
status = Column(String(30), nullable=False, default="pending", index=True)
|
||||
current_stage = Column(String(200), nullable=False, default="") # 细粒度阶段 snake_case
|
||||
phase_message = Column(String(500), nullable=False, default="") # 阶段中文提示文案
|
||||
heartbeat_at = Column(DateTime, nullable=True, index=True) # worker 心跳,用于僵尸任务超时回收
|
||||
intent_result = Column(JSON, nullable=True)
|
||||
image_analysis = Column(JSON, nullable=True)
|
||||
storyboard = Column(JSON, nullable=True)
|
||||
generated_copy_text = Column(Text, nullable=False, default="")
|
||||
copy_result = Column(
|
||||
JSON, nullable=True
|
||||
) # v1.6: 编导脚本结构{overview,scene_and_lighting,shots,hard_constraints,negative_prompts,voiceover_script}
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Float, nullable=False, default=0)
|
||||
video_resolution = Column(String(20), nullable=False, default="720p")
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0)
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="")
|
||||
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))
|
||||
|
||||
@@ -81,6 +81,41 @@ def ensure_database_exists(database_url: str) -> None:
|
||||
admin_engine.dispose()
|
||||
|
||||
|
||||
_VIRAL_VIDEO_BACKFILL_COLS = [
|
||||
("storyboard", "JSON"),
|
||||
("generated_copy_text", "TEXT NOT NULL DEFAULT ''"),
|
||||
("voice_id", "VARCHAR(200) NOT NULL DEFAULT ''"),
|
||||
("voice_source", "VARCHAR(20) NOT NULL DEFAULT ''"),
|
||||
("video_ratio", "VARCHAR(10) NOT NULL DEFAULT '9:16'"),
|
||||
("video_model", "VARCHAR(100) NOT NULL DEFAULT ''"),
|
||||
("copy_result", "JSON"),
|
||||
]
|
||||
|
||||
|
||||
def _ensure_viral_video_columns(connection) -> None:
|
||||
"""Idempotently add new columns to viral_video_jobs; create_all will not ALTER existing tables."""
|
||||
from sqlalchemy import inspect as _inspect
|
||||
|
||||
try:
|
||||
insp = _inspect(connection)
|
||||
if not insp.has_table("viral_video_jobs"):
|
||||
return
|
||||
existing = {c["name"] for c in insp.get_columns("viral_video_jobs")}
|
||||
except Exception:
|
||||
return
|
||||
import logging as _logging
|
||||
|
||||
_log = _logging.getLogger(__name__)
|
||||
for col, ddl in _VIRAL_VIDEO_BACKFILL_COLS:
|
||||
if col in existing:
|
||||
continue
|
||||
try:
|
||||
connection.execute(text(f"ALTER TABLE viral_video_jobs ADD COLUMN {col} {ddl}"))
|
||||
_log.info("added column viral_video_jobs.%s", col)
|
||||
except Exception as e:
|
||||
_log.warning("add column %s failed: %s", col, e)
|
||||
|
||||
|
||||
def initialize_database(engine) -> None:
|
||||
"""初始化数据库 schema。
|
||||
|
||||
@@ -100,4 +135,5 @@ def initialize_database(engine) -> None:
|
||||
text("SELECT pg_advisory_unlock(:lock_id)"),
|
||||
{"lock_id": SCHEMA_INIT_LOCK_ID},
|
||||
)
|
||||
_ensure_viral_video_columns(connection)
|
||||
connection.commit()
|
||||
|
||||
+257
@@ -0,0 +1,257 @@
|
||||
"""爆款视频任务 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 15,
|
||||
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 "",
|
||||
voice_id=getattr(model, "voice_id", "") or "",
|
||||
voice_source=getattr(model, "voice_source", "") or "",
|
||||
video_ratio=getattr(model, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(model, "video_model", "") or "",
|
||||
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
|
||||
current_stage=getattr(model, "current_stage", "") or "",
|
||||
phase_message=getattr(model, "phase_message", "") or "",
|
||||
heartbeat_at=getattr(model, "heartbeat_at", None),
|
||||
intent_result=dict(model.intent_result) if model.intent_result else None,
|
||||
image_analysis=dict(model.image_analysis) if getattr(model, "image_analysis", None) else None,
|
||||
storyboard=list(model.storyboard) if getattr(model, "storyboard", None) else None,
|
||||
generated_copy_text=getattr(model, "generated_copy_text", "") or "",
|
||||
copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
video_resolution=getattr(model, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(model, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(model, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(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,
|
||||
voice_id=job.voice_id,
|
||||
voice_source=job.voice_source,
|
||||
video_ratio=job.video_ratio,
|
||||
video_model=job.video_model,
|
||||
status=job.status,
|
||||
current_stage=job.current_stage or "",
|
||||
phase_message=job.phase_message or "",
|
||||
heartbeat_at=job.heartbeat_at,
|
||||
intent_result=job.intent_result,
|
||||
image_analysis=job.image_analysis,
|
||||
storyboard=job.storyboard,
|
||||
generated_copy_text=job.generated_copy_text,
|
||||
copy_result=job.copy_result,
|
||||
result_video_url=job.result_video_url,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(job, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
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.current_stage = job.current_stage or ""
|
||||
model.phase_message = job.phase_message or ""
|
||||
model.heartbeat_at = job.heartbeat_at
|
||||
model.intent_result = job.intent_result
|
||||
model.image_analysis = job.image_analysis
|
||||
model.storyboard = job.storyboard
|
||||
model.generated_copy_text = job.generated_copy_text or ""
|
||||
model.copy_result = job.copy_result
|
||||
model.result_video_url = job.result_video_url
|
||||
model.video_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
model.credits_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
model.credits_transaction_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
model.credits_cost = float(getattr(job, "credits_cost", 0) or 0)
|
||||
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
|
||||
# v1.5 three-stage: persist user-editable params so resume uses latest values
|
||||
model.user_copy_text = job.user_copy_text
|
||||
model.industry = job.industry
|
||||
model.target_customer = job.target_customer
|
||||
model.persona_id = job.persona_id
|
||||
model.viral_structure = job.viral_structure
|
||||
model.marketing_purpose = job.marketing_purpose
|
||||
model.bgm_preference = job.bgm_preference
|
||||
model.duration = job.duration
|
||||
model.fusion_level = job.fusion_level
|
||||
model.reference_audio_path = job.reference_audio_path
|
||||
model.reference_video_url = job.reference_video_url
|
||||
model.style_strength = job.style_strength
|
||||
model.style_template_id = job.style_template_id
|
||||
model.voice_id = job.voice_id or ""
|
||||
model.voice_source = job.voice_source or ""
|
||||
model.video_ratio = job.video_ratio or "9:16"
|
||||
model.video_model = job.video_model or ""
|
||||
model.updated_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
|
||||
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", "image_analyzed", "copy_generated"]
|
||||
),
|
||||
)
|
||||
.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,
|
||||
}
|
||||
@@ -52,6 +52,18 @@ class RenderedOutput:
|
||||
def from_dict(cls, data: dict[str, Any]) -> "RenderedOutput":
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("rendered_output must be a dict")
|
||||
# fingerprint_chunks 历史上有两种位置:
|
||||
# 1) 顶层 ``fingerprint_chunks``(由 compute_render_fingerprint_and_dedup 直接返回)
|
||||
# 2) 嵌套在 ``fingerprint_dict["chunks"]``(VideoFingerprint.to_dict() 序列化的结构)
|
||||
# 顶层优先;顶层为空时回退到嵌套位置,兼容旧数据。
|
||||
fp_dict = data.get("fingerprint_dict") or {}
|
||||
chunks_raw = data.get("fingerprint_chunks")
|
||||
if not chunks_raw and isinstance(fp_dict, dict):
|
||||
chunks_raw = fp_dict.get("chunks")
|
||||
# md5 同样可能在顶层或嵌套在 fingerprint_dict 内(历史数据兼容)
|
||||
md5_value = data.get("video_fingerprint_md5")
|
||||
if not md5_value and isinstance(fp_dict, dict):
|
||||
md5_value = fp_dict.get("md5")
|
||||
return cls(
|
||||
file_url=str(data.get("file_url") or ""),
|
||||
file_size=int(data.get("file_size") or 0),
|
||||
@@ -65,14 +77,14 @@ class RenderedOutput:
|
||||
batch_id=str(data.get("batch_id") or ""),
|
||||
project_id=str(data.get("project_id") or ""),
|
||||
user_id=str(data.get("user_id") or ""),
|
||||
fingerprint_dict=data.get("fingerprint_dict"),
|
||||
fingerprint_chunks=data.get("fingerprint_chunks"),
|
||||
fingerprint_dict=fp_dict or None,
|
||||
fingerprint_chunks=chunks_raw if isinstance(chunks_raw, list) else None,
|
||||
is_duplicate=bool(data.get("is_duplicate", False)),
|
||||
duplicate_of=data.get("duplicate_of"),
|
||||
duplicate_rate=_safe_float(data.get("duplicate_rate")),
|
||||
match_count=_safe_int(data.get("match_count")),
|
||||
visual_similarity=_safe_float(data.get("visual_similarity")),
|
||||
video_fingerprint_md5=str(data.get("video_fingerprint_md5") or ""),
|
||||
video_fingerprint_md5=str(md5_value or ""),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -51,7 +51,7 @@ class APISettings(SharedSettings):
|
||||
def validate_jwt_secret_key(cls, v):
|
||||
if v is None or v == "":
|
||||
raise ValueError(
|
||||
"JWT_SECRET_KEY must be set via environment variable. " "Do not use default value in production!"
|
||||
"JWT_SECRET_KEY must be set via environment variable. Do not use default value in production!"
|
||||
)
|
||||
# Block known insecure default values
|
||||
insecure_defaults = [
|
||||
@@ -63,7 +63,7 @@ class APISettings(SharedSettings):
|
||||
]
|
||||
if v.lower() in [d.lower() for d in insecure_defaults]:
|
||||
raise ValueError(
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. " "Please set a strong random secret via environment variable."
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. Please set a strong random secret via environment variable."
|
||||
)
|
||||
return v
|
||||
|
||||
@@ -248,6 +248,10 @@ class APISettings(SharedSettings):
|
||||
def OSS_ENDPOINT(self) -> str:
|
||||
return self.oss_endpoint
|
||||
|
||||
@property
|
||||
def OSS_INTERNAL_ENDPOINT(self) -> str:
|
||||
return self.effective_oss_internal_endpoint
|
||||
|
||||
@property
|
||||
def OSS_ACCESS_KEY_ID(self) -> str:
|
||||
return self.oss_access_key_id
|
||||
|
||||
+34
-2
@@ -47,12 +47,37 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── OSS 阿里云 ──────────────────────────────────────────────────────
|
||||
oss_endpoint: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
# 内网 endpoint:ECS VPC 内访问 OSS 用(千兆带宽、免公网流量费)。
|
||||
# 为空时自动从 oss_endpoint 推导:若 oss_endpoint 是阿里云公网域名(形如
|
||||
# oss-cn-<region>.aliyuncs.com),自动加 -internal 得到内网域名;其他情况
|
||||
# (自定义域名/本地 MinIO/非阿里云)回退使用 oss_endpoint。
|
||||
# 显式填同值可以覆盖自动推导、强制所有流量都走公网。
|
||||
oss_internal_endpoint: str = ""
|
||||
oss_access_key_id: str = ""
|
||||
oss_access_key_secret: str = ""
|
||||
oss_bucket_name: str = "xiaoxia-autocut"
|
||||
oss_direct_upload_max_mb: int = 2000
|
||||
oss_direct_upload_expire_seconds: int = 900
|
||||
|
||||
@property
|
||||
def effective_oss_internal_endpoint(self) -> str:
|
||||
"""实际用于 SDK 内网访问的 endpoint(带 -internal 自动推导)。"""
|
||||
if self.oss_internal_endpoint:
|
||||
return self.oss_internal_endpoint
|
||||
ep = self.oss_endpoint.strip()
|
||||
scheme = ""
|
||||
host = ep
|
||||
if ep.startswith("https://"):
|
||||
scheme = "https://"
|
||||
host = ep[len("https://") :]
|
||||
elif ep.startswith("http://"):
|
||||
scheme = "http://"
|
||||
host = ep[len("http://") :]
|
||||
# 阿里云公网域名自动推导:oss-cn-<region>.aliyuncs.com → oss-cn-<region>-internal.aliyuncs.com
|
||||
if host.endswith(".aliyuncs.com") and "-internal" not in host and host.startswith("oss-cn-"):
|
||||
host = host[: -len(".aliyuncs.com")] + "-internal.aliyuncs.com"
|
||||
return f"{scheme}{host}" if scheme else host
|
||||
|
||||
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
|
||||
cosyvoice_api_key: str = ""
|
||||
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
@@ -65,17 +90,24 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-1-6-250615"
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
|
||||
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
|
||||
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
|
||||
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 = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
mediakit_timeout: int = 60
|
||||
mediakit_cover_enabled: bool = False # 封面抽帧是否走MediaKit(默认false走本地ffmpeg+cv2,<2s完成)
|
||||
|
||||
# ── 积分/会员系统 (#1895) ────────────────────────────────────────────
|
||||
# 积分系统总开关(产品要求 #1895:暂停积分系统但保留全部代码/表/接口)。
|
||||
|
||||
@@ -9,9 +9,9 @@ from uuid import uuid4
|
||||
class PointsAccount:
|
||||
id: str
|
||||
user_id: str
|
||||
balance: int = 0
|
||||
total_earned: int = 0
|
||||
total_spent: int = 0
|
||||
balance: float = 0.0
|
||||
total_earned: float = 0.0
|
||||
total_spent: float = 0.0
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
+207
-54
@@ -1,32 +1,198 @@
|
||||
"""积分消耗规则配置 (#1895)"""
|
||||
"""积分消耗规则配置 (#1895)
|
||||
|
||||
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
|
||||
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
|
||||
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
# ============ 爆款视频动态定价 (#2151) ============
|
||||
# key = (model_id, resolution, has_video_input),单位:元/百万token
|
||||
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.5", "480p", False): 70.0,
|
||||
("seedance-2.5", "720p", False): 70.0,
|
||||
("seedance-2.5", "1080p", False): 77.0,
|
||||
("seedance-2.5", "480p", True): 42.0,
|
||||
("seedance-2.5", "720p", True): 42.0,
|
||||
("seedance-2.5", "1080p", True): 46.0,
|
||||
("seedance-2.0", "480p", False): 46.0,
|
||||
("seedance-2.0", "720p", False): 46.0,
|
||||
("seedance-2.0", "1080p", False): 51.0,
|
||||
}
|
||||
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
|
||||
VIRAL_VIDEO_FIXED_COST = 0.15
|
||||
# 利润系数
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
|
||||
# Seedance 输出帧率
|
||||
VIRAL_VIDEO_FPS = 24
|
||||
|
||||
# 分辨率别名映射 -> 标准 key
|
||||
_RESOLUTION_ALIASES: dict[str, str] = {
|
||||
"480p": "480p",
|
||||
"普清": "480p",
|
||||
"default": "480p",
|
||||
"low": "480p",
|
||||
"sd": "480p",
|
||||
"720p": "720p",
|
||||
"高清": "720p",
|
||||
"medium": "720p",
|
||||
"hd": "720p",
|
||||
"1080p": "1080p",
|
||||
"超清": "1080p",
|
||||
"high": "1080p",
|
||||
"ultra": "1080p",
|
||||
"全能": "1080p",
|
||||
"fhd": "1080p",
|
||||
}
|
||||
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
|
||||
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080}
|
||||
|
||||
|
||||
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
|
||||
"""把 (resolution, ratio) 解析为 (width, height)。
|
||||
|
||||
resolution 数字代表短边像素数(480p/720p/1080p 等):
|
||||
- 横屏 16:9:短边是 height,width = short * 16/9
|
||||
- 竖屏 9:16:短边是 width,height = short * 16/9
|
||||
- 方屏 1:1:width = height = short
|
||||
"""
|
||||
key = str(resolution or "").strip()
|
||||
key_l = key.lower()
|
||||
res_key = _RESOLUTION_ALIASES.get(key_l) or _RESOLUTION_ALIASES.get(key) or "720p"
|
||||
short = _RESOLUTION_SHORT_SIDE.get(res_key, 720)
|
||||
r = str(ratio or "").strip().lower()
|
||||
if r == "16:9":
|
||||
# 横屏:短边是 height,width 向上取整并对齐偶数
|
||||
w = math.ceil(short * 16 / 9)
|
||||
h = short
|
||||
elif r == "1:1":
|
||||
w, h = short, short
|
||||
else:
|
||||
# 9:16 竖屏(默认):短边是 width,height 向上取整并对齐偶数
|
||||
w = short
|
||||
h = math.ceil(short * 16 / 9)
|
||||
# 对齐到偶数(视频编码要求)
|
||||
w = w + (w % 2)
|
||||
h = h + (h % 2)
|
||||
return int(w), int(h)
|
||||
|
||||
|
||||
def _match_model_prefix(model: str) -> str:
|
||||
"""匹配 model 前缀。"""
|
||||
m = (model or "").strip().lower()
|
||||
for prefix in ("seedance-2.5", "seedance-2.0"):
|
||||
if m.startswith(prefix):
|
||||
return prefix
|
||||
return "seedance-2.5"
|
||||
|
||||
|
||||
def _infer_resolution_key(width: int, height: int) -> str:
|
||||
"""从实际 (width, height) 用短边推断 resolution key。"""
|
||||
short = min(int(width or 720), int(height or 720))
|
||||
if short >= 1000:
|
||||
return "1080p"
|
||||
if short >= 650:
|
||||
return "720p"
|
||||
return "480p"
|
||||
|
||||
|
||||
def calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> tuple[float, dict]:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
|
||||
Returns:
|
||||
(credits, breakdown) 二元组:
|
||||
- credits: 四舍五入保留两位小数的最终积分
|
||||
- breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier /
|
||||
model_price / width / height / fps 字段,便于前端展示计费明细。
|
||||
"""
|
||||
w = max(1, int(width or 1))
|
||||
h = max(1, int(height or 1))
|
||||
effective_fps = int(fps or VIRAL_VIDEO_FPS)
|
||||
|
||||
prefix = _match_model_prefix(model)
|
||||
res_key = _infer_resolution_key(w, h)
|
||||
key = (prefix, res_key, bool(has_video_input))
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
|
||||
if price is None:
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
|
||||
if actual_tokens is not None and actual_tokens > 0:
|
||||
tokens = float(actual_tokens)
|
||||
else:
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
tokens = dur * w * h * effective_fps / 1024.0
|
||||
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
|
||||
credits = round(float(total), 2)
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
|
||||
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
|
||||
"model_price": float(price),
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
|
||||
def calculate_viral_video_credits(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> float:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),仅返回积分值(向后兼容包装器)。
|
||||
|
||||
内部调用 calculate_viral_video_credits_with_breakdown,仅返回 credits 部分,
|
||||
保持旧调用方签名与返回值类型不变。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
"""
|
||||
credits, _ = calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds=duration_seconds,
|
||||
width=width,
|
||||
height=height,
|
||||
model=model,
|
||||
has_video_input=has_video_input,
|
||||
actual_tokens=actual_tokens,
|
||||
fps=fps,
|
||||
)
|
||||
return credits
|
||||
|
||||
|
||||
# ============ 场景定义 ============
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称)
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称), dynamic(是否动态定价)
|
||||
# 说明:爆款视频(viral_video)走动态定价(预扣→结算多退少补),因此不使用 @points_gate
|
||||
# 装饰器,base_points=0,dynamic=True;前端展示场景列表时仍可看到。
|
||||
|
||||
POINTS_SCENES: dict[str, dict] = {
|
||||
"ai_voice": {
|
||||
"base_points": 1,
|
||||
"unit": "分钟",
|
||||
"name": "AI 配音",
|
||||
"description": "AI 配音每分钟消耗 1 积分(免费用户上浮 15%,会员 8~9 折)",
|
||||
},
|
||||
"ai_video": {
|
||||
"base_points": 3,
|
||||
"unit": "条",
|
||||
"name": "智能混剪",
|
||||
"extra_per_30s": 1,
|
||||
"description": "智能混剪每条 3 积分起,视频超过 30 秒后每 30 秒加 1 积分;免费用户每日 2 条免费额度",
|
||||
},
|
||||
"ai_digital_human": {
|
||||
"base_points": 15,
|
||||
"unit": "分钟",
|
||||
"name": "AI 数字人",
|
||||
"description": "AI 数字人每分钟消耗 15 积分",
|
||||
},
|
||||
"voice_clone_train": {
|
||||
"base_points": 0,
|
||||
"unit": "次",
|
||||
@@ -39,23 +205,16 @@ POINTS_SCENES: dict[str, dict] = {
|
||||
"name": "声音克隆合成",
|
||||
"description": "克隆音色合成每分钟消耗 1 积分",
|
||||
},
|
||||
"douyin_extract": {
|
||||
"base_points": 1,
|
||||
"viral_video": {
|
||||
"base_points": 0,
|
||||
"unit": "次",
|
||||
"name": "抖音链接提取",
|
||||
"description": "抖音文案提取每次 1 积分",
|
||||
"name": "爆款视频",
|
||||
"dynamic": True,
|
||||
"description": "爆款视频动态定价(按视频时长/分辨率/模型计算,预扣→结算多退少补)",
|
||||
},
|
||||
"ai_rewrite": {"base_points": 1, "unit": "次", "name": "AI 改写文案", "description": "AI 改写文案每次 1 积分"},
|
||||
"ai_title": {
|
||||
"base_points": 1,
|
||||
"unit": "次",
|
||||
"name": "AI 标题生成",
|
||||
"description": "AI 生成标题每次 1 积分(免费用户实际上浮后 2 积分/次)",
|
||||
},
|
||||
"ai_cover": {"base_points": 1, "unit": "张", "name": "AI 封面生成", "description": "AI 封面生成每张 1 积分"},
|
||||
}
|
||||
|
||||
# 免费用户积分消耗上浮系数
|
||||
# 免费用户积分消耗上浮系数(仅对 voice_clone_synth 生效)
|
||||
FREE_USER_MULTIPLIER = 1.15
|
||||
|
||||
# ============ 积分包定义 ============
|
||||
@@ -81,9 +240,6 @@ MEMBER_DISCOUNT: dict[str, float] = {
|
||||
"yearly": 0.8,
|
||||
}
|
||||
|
||||
# 每日免费混剪次数(免费用户)
|
||||
DAILY_FREE_CLIP_LIMIT = 2
|
||||
|
||||
|
||||
def calculate_points_cost(
|
||||
scene_key: str,
|
||||
@@ -91,48 +247,45 @@ def calculate_points_cost(
|
||||
quantity: int = 1,
|
||||
duration_minutes: float = 0,
|
||||
member_type: str | None = None,
|
||||
) -> int:
|
||||
) -> float:
|
||||
"""计算指定场景的积分消耗。
|
||||
|
||||
Args:
|
||||
scene_key: 场景标识,如 "ai_voice"、"ai_video"
|
||||
scene_key: 场景标识(当前支持 voice_clone_train/voice_clone_synth/viral_video;
|
||||
viral_video 为动态定价场景,此处返回 0,由业务侧调用
|
||||
calculate_viral_video_credits 手动计算)
|
||||
is_member: 是否付费会员
|
||||
quantity: 数量(按次计费场景)
|
||||
duration_minutes: 时长分钟数(按时长计费场景)
|
||||
member_type: 会员类型 (monthly/quarterly/yearly),用于折扣
|
||||
|
||||
Returns:
|
||||
实际消耗积分(已含免费用户 ×1.15 上浮或会员折扣)
|
||||
|
||||
Raises:
|
||||
ValueError: 未知场景标识
|
||||
实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/动态/已下线场景统一返回 0。
|
||||
"""
|
||||
scene = POINTS_SCENES.get(scene_key)
|
||||
if not scene:
|
||||
raise ValueError(f"Unknown points scene: {scene_key}")
|
||||
# 已下线/未注册的场景统一返回 0(免费),保持向后兼容
|
||||
return 0.0
|
||||
|
||||
# 动态定价场景(如 viral_video)由业务侧手动计算,这里统一返回 0
|
||||
if scene.get("dynamic"):
|
||||
return 0.0
|
||||
|
||||
base = scene["base_points"]
|
||||
if base == 0:
|
||||
return 0
|
||||
return 0.0
|
||||
|
||||
# —— 计算基础消耗 ——
|
||||
unit = scene["unit"]
|
||||
if unit == "分钟":
|
||||
total_base = base * max(1, math.ceil(duration_minutes))
|
||||
elif unit in ("条", "次", "张"):
|
||||
elif unit in ("次", "张"):
|
||||
total_base = base * quantity
|
||||
# 混剪特殊逻辑:视频超过 30s 后每 +30s 额外加 1 积分
|
||||
if scene_key == "ai_video" and duration_minutes > 0.5:
|
||||
extra_segments = math.ceil((duration_minutes * 60 - 30) / 30)
|
||||
if extra_segments > 0:
|
||||
total_base += scene.get("extra_per_30s", 1) * extra_segments
|
||||
else:
|
||||
total_base = base
|
||||
|
||||
# —— 会员折扣 / 免费用户上浮 ——
|
||||
if is_member and member_type and member_type in MEMBER_DISCOUNT:
|
||||
total_base = max(1, math.floor(total_base * MEMBER_DISCOUNT[member_type]))
|
||||
elif not is_member:
|
||||
total_base = math.ceil(total_base * FREE_USER_MULTIPLIER)
|
||||
|
||||
return total_base
|
||||
return float(total_base)
|
||||
|
||||
@@ -13,7 +13,6 @@ from typing import Any
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import (
|
||||
DAILY_FREE_CLIP_LIMIT,
|
||||
POINTS_PACKAGES,
|
||||
)
|
||||
|
||||
@@ -84,7 +83,7 @@ class PointsService:
|
||||
|
||||
# ──────────────── 余额检查 ────────────────
|
||||
|
||||
def check_balance(self, user_id: str, amount: int, db: Session) -> dict[str, Any]:
|
||||
def check_balance(self, user_id: str, amount: float, db: Session) -> dict[str, Any]:
|
||||
"""检查余额是否足够。"""
|
||||
account_data = self.get_or_create_account(user_id, db)
|
||||
balance = account_data["balance"]
|
||||
@@ -100,7 +99,7 @@ class PointsService:
|
||||
def deduct_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: int,
|
||||
amount: float,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -109,7 +108,7 @@ class PointsService:
|
||||
"""扣减积分(事务性:SELECT FOR UPDATE → 检查余额 → 扣减 → 流水 → 同步用户表)。
|
||||
|
||||
Returns:
|
||||
{"success": True/False, "balance": int, "transaction_id": str|None}
|
||||
{"success": True/False, "balance": float, "transaction_id": str|None}
|
||||
"""
|
||||
PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models()
|
||||
|
||||
@@ -173,7 +172,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分扣减失败: user_id=%s, amount=%d, source=%s",
|
||||
"积分扣减失败: user_id=%s, amount=%.2f, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -185,7 +184,7 @@ class PointsService:
|
||||
def add_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: int,
|
||||
amount: float,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -242,7 +241,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分增加失败: user_id=%s, amount=%d, source=%s",
|
||||
"积分增加失败: user_id=%s, amount=%.2f, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -254,7 +253,7 @@ class PointsService:
|
||||
def refund_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: int,
|
||||
amount: float,
|
||||
source: str,
|
||||
db: Session,
|
||||
ref_id: str = "",
|
||||
@@ -270,6 +269,92 @@ class PointsService:
|
||||
ref_id=ref_id,
|
||||
)
|
||||
|
||||
# ──────────────── 爆款视频(viral_video)动态定价 ────────────────
|
||||
|
||||
def deduct_viral_video(self, user_id: str, credits: float, job_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频预扣积分(confirm-copy 阶段)。"""
|
||||
return self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=float(credits or 0),
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频生成",
|
||||
ref_id=job_id,
|
||||
)
|
||||
|
||||
def settle_viral_video(
|
||||
self,
|
||||
user_id: str,
|
||||
estimated: float,
|
||||
actual: float,
|
||||
txn_id: str,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""爆款视频完成后按实际 tokens 结算(多退少补)。
|
||||
|
||||
- actual < estimated: 退差额
|
||||
- actual > estimated: 补扣差额(余额不足时记 warning,不阻塞完成)
|
||||
- |diff| < 0.01: 不动
|
||||
"""
|
||||
diff = round(float(actual or 0) - float(estimated or 0), 2)
|
||||
if abs(diff) < 0.01:
|
||||
return {"success": True, "action": "none", "diff": 0.0}
|
||||
if diff < 0:
|
||||
refund = round(-diff, 2)
|
||||
try:
|
||||
res = self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频结算退费",
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "refund", "diff": -refund, "amount": refund}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算退费异常 user_id=%s refund=%.2f", user_id, refund)
|
||||
return {"success": False, "action": "refund", "diff": -refund}
|
||||
else:
|
||||
extra = round(diff, 2)
|
||||
try:
|
||||
res = self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=extra,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频结算补扣",
|
||||
ref_id=txn_id,
|
||||
)
|
||||
if not res.get("success"):
|
||||
logger.warning(
|
||||
"[viral_video] 结算补扣余额不足 user_id=%s extra=%.2f balance=%s (不阻塞任务完成)",
|
||||
user_id,
|
||||
extra,
|
||||
res.get("balance"),
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "deduct", "diff": extra, "amount": extra}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算补扣异常 user_id=%s extra=%.2f", user_id, extra)
|
||||
return {"success": False, "action": "deduct", "diff": extra}
|
||||
|
||||
def refund_viral_video(self, user_id: str, credits: float, txn_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频失败全额退款。"""
|
||||
amount = float(credits or 0)
|
||||
if amount <= 0:
|
||||
return {"success": True, "action": "none", "amount": 0.0}
|
||||
try:
|
||||
return self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频失败退款",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 失败退款异常 user_id=%s amount=%.2f", user_id, amount)
|
||||
return {"success": False, "action": "refund", "amount": amount}
|
||||
|
||||
# ──────────────── 流水查询 ────────────────
|
||||
|
||||
def get_transactions(
|
||||
@@ -324,132 +409,16 @@ class PointsService:
|
||||
"page_size": page_size,
|
||||
}
|
||||
|
||||
# ──────────────── 每日免费混剪额度 ────────────────
|
||||
|
||||
def _daily_key(self, user_id: str) -> str:
|
||||
"""生成 Redis 每日额度 key。格式: daily_usage:{user_id}:{YYYYMMDD}:free_clip"""
|
||||
today = datetime.now(UTC).strftime("%Y%m%d")
|
||||
return f"daily_usage:{user_id}:{today}:free_clip"
|
||||
|
||||
def check_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""检查今日是否还有免费混剪额度。
|
||||
|
||||
优先查 Redis,Redis 不可用时降级到 DB。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
current = redis_client.get(key)
|
||||
if current is None:
|
||||
return True
|
||||
return int(current) < DAILY_FREE_CLIP_LIMIT
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 查询每日额度")
|
||||
|
||||
# 降级到 DB
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if record is None:
|
||||
return True
|
||||
return record.count < DAILY_FREE_CLIP_LIMIT
|
||||
|
||||
def record_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""记录使用一次免费混剪。
|
||||
|
||||
先 INCR Redis;如果超限回退 Redis。DB 使用 upsert 语义(唯一约束)。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
new_count = redis_client.incr(key)
|
||||
if new_count == 1:
|
||||
redis_client.expire(key, 48 * 3600) # TTL 48h
|
||||
if new_count <= DAILY_FREE_CLIP_LIMIT:
|
||||
return True
|
||||
# 超限,回退 Redis
|
||||
redis_client.decr(key)
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 记录每日额度")
|
||||
|
||||
# 降级/兜底到 DB(upsert 语义)
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if record is None:
|
||||
if DAILY_FREE_CLIP_LIMIT <= 0:
|
||||
return False
|
||||
record = DailyUsageRecordModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
usage_type="free_clip",
|
||||
usage_date=datetime.now(UTC),
|
||||
count=1,
|
||||
)
|
||||
db.add(record)
|
||||
else:
|
||||
if record.count >= DAILY_FREE_CLIP_LIMIT:
|
||||
return False
|
||||
record.count += 1
|
||||
|
||||
db.commit()
|
||||
return True
|
||||
# ──────────────── 每日免费混剪额度(已下线:智能混剪全免费) ────────────────
|
||||
|
||||
def get_daily_usage(self, user_id: str, db: Session) -> dict[str, Any]:
|
||||
"""查询今日免费额度使用情况。"""
|
||||
redis_client = _get_redis_client()
|
||||
used = 0
|
||||
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
val = redis_client.get(key)
|
||||
used = int(val) if val else 0
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if used == 0:
|
||||
# 从 DB 查
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
used = record.count if record else 0
|
||||
|
||||
"""查询今日免费额度使用情况(智能混剪已全免费,返回 unlimited)。"""
|
||||
now = datetime.now(UTC)
|
||||
tomorrow = (now + timedelta(days=1)).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
return {
|
||||
"free_clips_used": used,
|
||||
"free_clips_limit": DAILY_FREE_CLIP_LIMIT,
|
||||
"free_clips_remaining": max(0, DAILY_FREE_CLIP_LIMIT - used),
|
||||
"free_clips_used": 0,
|
||||
"free_clips_limit": -1, # -1 表示 unlimited
|
||||
"free_clips_remaining": -1,
|
||||
"reset_at": tomorrow.isoformat(),
|
||||
}
|
||||
|
||||
|
||||
Executable
+251
@@ -0,0 +1,251 @@
|
||||
"""ViralVideoJob 领域模型 — 爆款视频任务.
|
||||
|
||||
v1.6 重大简化:Seedance 2.5 单次最长30秒,单次调用直接出片,不再分段/拼接/ffmpeg concat。
|
||||
状态机(三步分步):
|
||||
pending -> running -> image_analyzed -> running -> copy_generated -> running -> completed
|
||||
wait_user_confirm -> running -> completed (旧路径兼容)
|
||||
任意阶段 fail; 任意非终态 cancel.
|
||||
failed -> pending (retry 重置后重跑)。
|
||||
"""
|
||||
|
||||
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"
|
||||
IMAGE_ANALYZED = "image_analyzed"
|
||||
COPY_GENERATED = "copy_generated"
|
||||
WAIT_USER_CONFIRM = "wait_user_confirm"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class ViralVideoStage(StrEnum):
|
||||
IMAGE_ANALYSIS = "image_analysis"
|
||||
VIDEO_ANALYSIS = "video_analysis"
|
||||
INTENT_PARSING = "intent_parsing"
|
||||
SCRIPT_GENERATION = "script_generation" # v1.6: 编导分镜脚本(融合原 copy_fusion+storyboard+review)
|
||||
REVIEW = "review"
|
||||
TTS = "tts"
|
||||
RENDERING = "rendering" # v1.6: 单次 Seedance 生成(BGM/音效/画面一次出片)
|
||||
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):
|
||||
IMAGE_ANALYSIS = "image_analysis"
|
||||
INTENT_PARSING = "intent_parsing"
|
||||
SCRIPT_GENERATION = "script_generation"
|
||||
REVIEW = "review"
|
||||
VIDEO_STYLE_INTEGRATION = "video_style_integration"
|
||||
STYLE_CONSTRAINT = "style_constraint"
|
||||
|
||||
|
||||
STAGE_LABELS = {
|
||||
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
|
||||
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
|
||||
ViralVideoStage.INTENT_PARSING: "意图解析",
|
||||
ViralVideoStage.SCRIPT_GENERATION: "编导脚本生成",
|
||||
ViralVideoStage.REVIEW: "合规审核",
|
||||
ViralVideoStage.TTS: "AI 配音",
|
||||
ViralVideoStage.RENDERING: "视频生成",
|
||||
ViralVideoStage.UPLOADING: "上传发布",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ViralVideoJob:
|
||||
"""爆款视频任务领域实体(v1.6 单次 Seedance 出片版)。"""
|
||||
|
||||
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 = 15 # v1.6: 默认15秒,上限30秒(Seedance 2.5 单次最大30s)
|
||||
user_copy_text: str = ""
|
||||
fusion_level: str = FusionLevel.AI_POLISH
|
||||
reference_audio_path: str = ""
|
||||
reference_video_url: str = ""
|
||||
style_strength: str = StyleStrength.MEDIUM
|
||||
style_guide: dict | None = None
|
||||
style_template_id: str = ""
|
||||
# v1.5.1 音频/视频参数
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
# v1.4+ 产物
|
||||
image_analysis: dict | None = None
|
||||
intent_result: dict | None = None
|
||||
generated_copy_text: str = "" # v1.6: 存 voiceover_script(纯口播对白),字段名兼容
|
||||
storyboard: list | None = None # v1.6: 存 copy_result.shots,字段名兼容
|
||||
copy_result: dict | None = None # v1.6: 完整编导脚本结构
|
||||
# 状态
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
status: ViralVideoStatus = ViralVideoStatus.PENDING
|
||||
current_stage: str = "" # 细粒度阶段(ViralVideoStage.value,snake_case)
|
||||
phase_message: str = "" # 阶段中文提示文案,前端轮询直接展示
|
||||
heartbeat_at: datetime | None = None # worker 心跳时间,用于超时僵尸任务检测
|
||||
result_video_url: str = ""
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_transaction_id: str = ""
|
||||
credits_cost: float = 0.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.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.WAIT_USER_CONFIRM,
|
||||
ViralVideoStatus.RUNNING,
|
||||
):
|
||||
raise ValueError(f"Cannot transition from {self.status} to running")
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
now = datetime.now(timezone.utc)
|
||||
if self.started_at is None:
|
||||
self.started_at = now
|
||||
self.heartbeat_at = now
|
||||
self.updated_at = now
|
||||
|
||||
def touch_heartbeat(self) -> None:
|
||||
"""更新心跳时间(worker 在长任务中周期性调用,用于超时检测)。"""
|
||||
now = datetime.now(timezone.utc)
|
||||
if self.started_at is None:
|
||||
self.started_at = now
|
||||
self.heartbeat_at = now
|
||||
self.updated_at = now
|
||||
|
||||
def mark_image_analyzed(self) -> None:
|
||||
if self.status not in (ViralVideoStatus.PENDING, ViralVideoStatus.RUNNING):
|
||||
raise ValueError(f"Cannot transition from {self.status} to image_analyzed")
|
||||
self.status = ViralVideoStatus.IMAGE_ANALYZED
|
||||
if self.started_at is None:
|
||||
self.started_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_copy_generated(self, copy_result: dict) -> None:
|
||||
"""v1.6 阶段2完成:编导脚本(含 voiceover_script/shots/硬约束/负面词)已生成。"""
|
||||
if self.status not in (
|
||||
ViralVideoStatus.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.RUNNING,
|
||||
ViralVideoStatus.PENDING,
|
||||
):
|
||||
raise ValueError(f"Cannot transition from {self.status} to copy_generated")
|
||||
self.status = ViralVideoStatus.COPY_GENERATED
|
||||
self.copy_result = copy_result or {}
|
||||
if isinstance(copy_result, dict):
|
||||
self.generated_copy_text = copy_result.get("voiceover_script", "") or ""
|
||||
shots = copy_result.get("shots") or []
|
||||
self.storyboard = list(shots) if isinstance(shots, list) else []
|
||||
self.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_image_analyzed(self, **kwargs) -> None:
|
||||
if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING):
|
||||
raise ValueError(f"Cannot resume from {self.status} to copy-gen")
|
||||
for k, v in kwargs.items():
|
||||
if hasattr(self, k) and v not in (None, "", []):
|
||||
setattr(self, k, v)
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_copy_generated(self, edited_copy: str | None = None) -> None:
|
||||
"""阶段2->阶段3:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
|
||||
if self.status != ViralVideoStatus.COPY_GENERATED:
|
||||
raise ValueError(f"Cannot resume from {self.status} to render")
|
||||
if edited_copy and isinstance(self.copy_result, dict):
|
||||
self.copy_result = {**self.copy_result, "voiceover_script": edited_copy}
|
||||
self.generated_copy_text = edited_copy
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_confirm(self) -> None:
|
||||
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise ValueError(f"Cannot resume from {self.status}")
|
||||
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,
|
||||
)
|
||||
|
||||
@property
|
||||
def effective_copy_text(self) -> str:
|
||||
"""TTS 用的最终口播文案:优先 copy_result.voiceover_script,兼容老字段。"""
|
||||
if isinstance(self.copy_result, dict) and self.copy_result.get("voiceover_script"):
|
||||
return self.copy_result["voiceover_script"]
|
||||
return self.generated_copy_text or self.user_copy_text or "你好,给大家推荐一款好物"
|
||||
|
||||
@property
|
||||
def voiceover_script(self) -> str:
|
||||
return self.effective_copy_text
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user