Compare commits
206 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 9344314eac | |||
| 3f49867384 | |||
| dd420c556f | |||
| 2d823a9255 | |||
| ed24c7cd68 | |||
| 69f88434bd | |||
| 599388d9e0 | |||
| 9ffe909dc0 | |||
| 6243196408 | |||
| 3d8f2c2ce2 | |||
| 3eef497dfe | |||
| 61c15eb987 | |||
| 0012ecad30 | |||
| e1994ada0a | |||
| f9f6c53ef4 | |||
| 9699a1fcde | |||
| 305e2bd9d5 | |||
| 30cc58441f | |||
| b3b5dbd459 | |||
| 9f5948dd2b | |||
| fbc1df36e0 | |||
| d36cc09be5 | |||
| fa4dbc6761 | |||
| 77133eb9a0 | |||
| 35207ae040 | |||
| 5af16f1aea | |||
| c232721fff | |||
| fc27d4e81e | |||
| 3c782f89d1 | |||
| 2b4f11036f | |||
| 2881da65cf | |||
| 1037e218bb | |||
| f0514d7487 | |||
| bf8f62ec5b | |||
| cb911f5eba | |||
| f6fa3ea653 | |||
| 6a7709ba43 | |||
| 2c70dd4c29 | |||
| 5548e78eee | |||
| 44bd96b148 | |||
| dda67cd10f | |||
| 233d0272a9 | |||
| 631dd643c0 | |||
| 50489b05d6 | |||
| 252ea71d50 | |||
| d5f6e9499d | |||
| afd6b9c0b8 | |||
| 8a6ebb49be | |||
| 9e024caff5 | |||
| c7e9878f0d | |||
| b6243c8ab8 | |||
| 67e4b3fc6f | |||
| df654bce19 | |||
| bd99c161dd | |||
| 7fee203693 | |||
| 562ffc53cc | |||
| 9a67f727b3 | |||
| 5bb714f25b | |||
| 0f592add29 | |||
| 42dd5fabc6 | |||
| c7aed2c152 | |||
| 3428f4ef73 | |||
| 13c57771b9 | |||
| 5eedf190c1 | |||
| 7a9a97b31a | |||
| a6f89067e7 | |||
| 718deb32b9 | |||
| 054e81c3e7 | |||
| 07cf055a97 | |||
| 0e78f175fe | |||
| 371d1b9daf | |||
| 1a459d1ad2 | |||
| f097e16bc1 | |||
| 69eea54cd2 | |||
| b5cbc3b482 | |||
| 9106b4de2e | |||
| 7d639d5f9a | |||
| db6c237d68 | |||
| 7b68b94df6 | |||
| 8e19f24984 | |||
| 866d71a431 | |||
| 042512a527 | |||
| 3bb9c5dd4e | |||
| c87810a4a6 | |||
| 75ec9db439 | |||
| 0d70074182 | |||
| 071a3707f4 | |||
| 4104401759 | |||
| f0d46b44d7 | |||
| d63f7b4650 | |||
| b28ac8bd1b | |||
| 1b02df4d4a | |||
| 76e11cb2f9 | |||
| 1baf29c76a | |||
| 83465397fe | |||
| 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 | |||
| 0a004db1bd |
+15
-1
@@ -211,10 +211,24 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
|
||||
# 用于 AI 文案生成、智能剪辑等需要大模型能力的场景
|
||||
|
||||
DOUBAO_API_KEY=your-doubao-api-key
|
||||
DOUBAO_MODEL=doubao-seed-1-6-250615
|
||||
DOUBAO_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_FAST_MODEL=doubao-seed-2-1-lite-260915
|
||||
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-seed-2-1-pro-260915
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
# Embedding 向量化模型
|
||||
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
|
||||
# 视频模型(Seedance 2.5,统一走方舟;真人参考图通过信任链自动 AI 化)
|
||||
DOUBAO_VIDEO_MODEL=doubao-seedance-2-5-260628
|
||||
DOUBAO_VIDEO_TIMEOUT=480
|
||||
DOUBAO_VIDEO_POLL_INTERVAL=10
|
||||
# 图片模型(Seedream 5.0 Pro,用于信任链真人 AI 化 + 文生图)
|
||||
DOUBAO_IMAGE_MODEL=doubao-seedream-5-0-pro-260628
|
||||
DOUBAO_IMAGE_TIMEOUT=120
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -1187,6 +1187,14 @@ jobs:
|
||||
DOUBAO_MODEL: "${{ secrets.DOUBAO_MODEL }}"
|
||||
DOUBAO_BASE_URL: "${{ secrets.DOUBAO_BASE_URL }}"
|
||||
DOUBAO_VISION_MODEL: "${{ secrets.DOUBAO_VISION_MODEL }}"
|
||||
DOUBAO_VISION_LITE_MODEL: "${{ secrets.DOUBAO_VISION_LITE_MODEL }}"
|
||||
DOUBAO_VISION_USE_LITE: "${{ secrets.DOUBAO_VISION_USE_LITE }}"
|
||||
DOUBAO_IMAGE_MODEL: "${{ secrets.DOUBAO_IMAGE_MODEL }}"
|
||||
DOUBAO_IMAGE_SIZE: "${{ secrets.DOUBAO_IMAGE_SIZE }}"
|
||||
DOUBAO_IMAGE_TIMEOUT: "${{ secrets.DOUBAO_IMAGE_TIMEOUT }}"
|
||||
DOUBAO_FAST_MODEL: "${{ secrets.DOUBAO_FAST_MODEL }}"
|
||||
DOUBAO_TIMEOUT: "${{ secrets.DOUBAO_TIMEOUT }}"
|
||||
DOUBAO_MAX_RETRIES: "${{ secrets.DOUBAO_MAX_RETRIES }}"
|
||||
WECHAT_APP_ID: "${{ secrets.WECHAT_APP_ID }}"
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
Mon Oct 5 04:09:11 PM CST 2026
|
||||
2198 lite/pro并行竞速 (commit 9699a1f) — CI rebuild trigger Mon Oct 5 08:09:11 AM UTC 2026
|
||||
+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
|
||||
@@ -0,0 +1,31 @@
|
||||
"""viral_video_jobs 增加 pre_trusted_images 列(信任链Seedream预热结果)
|
||||
|
||||
Revision ID: 094_viral_video_pre_trusted
|
||||
Revises: 093_viral_video_pricing_points_float
|
||||
Create Date: 2026-10-04
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "094_viral_video_pre_trusted"
|
||||
down_revision = "093"
|
||||
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 "pre_trusted_images" not in cols:
|
||||
op.add_column("viral_video_jobs", sa.Column("pre_trusted_images", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "pre_trusted_images" in cols:
|
||||
op.drop_column("viral_video_jobs", "pre_trusted_images")
|
||||
@@ -0,0 +1,102 @@
|
||||
"""爆款视频 Prompt 模板配置表(#2040)。
|
||||
|
||||
086 曾预留同名旧表(id varchar / content / variables json),从未被业务使用;
|
||||
本迁移将其替换为 #2040 新结构。
|
||||
|
||||
Revision ID: 095_viral_video_prompt_templates
|
||||
Revises: 094_viral_video_pre_trusted
|
||||
Create Date: 2026-10-04
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "095_viral_video_prompt_templates"
|
||||
down_revision = "094_viral_video_pre_trusted"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _table_exists(conn, name: str) -> bool:
|
||||
return name in sa.inspect(conn).get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
# 086 预留的旧结构表:先删除(无业务数据、无任何引用)
|
||||
if _table_exists(conn, "viral_video_prompt_templates"):
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.Integer, primary_key=True, autoincrement=True),
|
||||
sa.Column("name", sa.String(128), nullable=False),
|
||||
sa.Column("prompt_type", sa.String(32), nullable=False),
|
||||
sa.Column("version", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column("system_prompt", sa.Text, nullable=False),
|
||||
sa.Column("user_prompt_template", sa.Text, nullable=False),
|
||||
sa.Column("example_output", sa.Text, nullable=True),
|
||||
sa.Column("is_active", sa.Boolean, nullable=False, server_default=sa.text("true")),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.func.now(),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.func.now(),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_vvpt_type_active",
|
||||
"viral_video_prompt_templates",
|
||||
["prompt_type", "is_active"],
|
||||
)
|
||||
op.create_index(
|
||||
"uq_vvpt_type_version",
|
||||
"viral_video_prompt_templates",
|
||||
["prompt_type", "version"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if _table_exists(conn, "viral_video_prompt_templates"):
|
||||
op.drop_index("uq_vvpt_type_version", table_name="viral_video_prompt_templates")
|
||||
op.drop_index("ix_vvpt_type_active", table_name="viral_video_prompt_templates")
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
|
||||
# 恢复 086 的旧预留结构
|
||||
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(),
|
||||
),
|
||||
)
|
||||
@@ -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"),
|
||||
|
||||
@@ -41,7 +41,6 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -270,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),
|
||||
|
||||
@@ -163,7 +163,6 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
return matched or None
|
||||
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -465,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),
|
||||
|
||||
@@ -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}",
|
||||
|
||||
@@ -145,19 +145,22 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表。"""
|
||||
packages = []
|
||||
for code, pkg in POINTS_PACKAGES.items():
|
||||
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
|
||||
packages.append(
|
||||
PointsPackageItem(
|
||||
code=code,
|
||||
name=pkg["name"],
|
||||
points=pkg["points"],
|
||||
price_cents=pkg["price_cents"],
|
||||
unit_price=unit_price,
|
||||
)
|
||||
"""查询可购买的积分包列表(读管理后台 credit_packages 表真实数据)。
|
||||
|
||||
仅返回 is_active=true;后台改价/启停后最多 30 秒生效。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_points_packages
|
||||
|
||||
packages = [
|
||||
PointsPackageItem(
|
||||
code=row["code"],
|
||||
name=row["name"],
|
||||
points=row["points"],
|
||||
price_cents=row["price_cents"],
|
||||
unit_price=row["unit_price"],
|
||||
)
|
||||
for row in get_points_packages()
|
||||
]
|
||||
mt = _member_type(current_user)
|
||||
discount = MEMBER_DISCOUNT.get(mt) if mt else None
|
||||
return PointsPackagesResponse(packages=packages, user_discount=discount)
|
||||
@@ -169,17 +172,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 +188,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 +201,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),
|
||||
|
||||
@@ -86,33 +86,13 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
"""
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for plan_id, info in MEMBERSHIP_PRICES.items():
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"price_cents": info["price_cents"],
|
||||
"monthly_price_cents": monthly_cents,
|
||||
"duration_days": days,
|
||||
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
|
||||
"features": features,
|
||||
})
|
||||
return {"plans": plans}
|
||||
return {"plans": get_membership_plans()}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
|
||||
@@ -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,
|
||||
@@ -390,6 +407,7 @@ async def prepare_direct_upload(
|
||||
duplicated=True,
|
||||
skip_transfer=True,
|
||||
asset_id=existing.id,
|
||||
url=_get_existing_asset_url(existing, storage_service),
|
||||
)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
@@ -443,6 +461,7 @@ async def prepare_direct_upload(
|
||||
duplicated=False,
|
||||
skip_transfer=False,
|
||||
asset_id=pending_asset_id,
|
||||
url="",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,981 @@
|
||||
"""爆款视频 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.points_rules import list_viral_video_models
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
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,
|
||||
pre_trusted_images=getattr(job, "pre_trusted_images", None) or None,
|
||||
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("/models")
|
||||
def list_available_models() -> dict:
|
||||
"""返回爆款视频可用模型列表(供前端模型选择器使用)。"""
|
||||
dashscope_available = get_dashscope_client() is not None
|
||||
models = list_viral_video_models(
|
||||
include_placeholder=False,
|
||||
dashscope_available=dashscope_available,
|
||||
)
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@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, "任务准备中")
|
||||
+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
+320
@@ -0,0 +1,320 @@
|
||||
"""爆款视频 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 = ""
|
||||
pre_trusted_images: list[str] | None = None
|
||||
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,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
import apiClient from "@/api/client"
|
||||
import type {
|
||||
GenerateViralVideoRequest,
|
||||
HistoryResponse,
|
||||
StyleTemplate,
|
||||
ViralVideoJob,
|
||||
ImageAnalysisResult,
|
||||
CopyResult,
|
||||
AnalyzeImagesRequest,
|
||||
GenerateCopyRequest,
|
||||
ConfirmCopyRequest,
|
||||
ViralVideoModel,
|
||||
ViralVideoModelsResponse,
|
||||
} 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)
|
||||
}
|
||||
|
||||
/** 获取支持的视频模型列表(GET /viral-video/models)。后端返回 {models: [...]} 包装 */
|
||||
export function getViralVideoModels() {
|
||||
return apiClient.get<ViralVideoModelsResponse>("/viral-video/models").then((r) => {
|
||||
const data = r.data as ViralVideoModelsResponse | ViralVideoModel[] | null | undefined
|
||||
if (Array.isArray(data)) return data
|
||||
if (data && Array.isArray((data as ViralVideoModelsResponse).models)) {
|
||||
return (data as ViralVideoModelsResponse).models
|
||||
}
|
||||
return []
|
||||
})
|
||||
}
|
||||
/** ── 三步拆分:前端 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,317 @@
|
||||
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" | "my_voice"
|
||||
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" | "my_voice"
|
||||
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" | "my_voice"
|
||||
/** 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" | "my_voice"
|
||||
/** Seedance 视频比例(9:16/16:9/1:1 等) */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 视频模型描述(GET /viral-video/models) */
|
||||
export interface ViralVideoModel {
|
||||
key: string
|
||||
display_name: string
|
||||
supports_audio: boolean
|
||||
supported_resolutions: string[]
|
||||
max_duration: number
|
||||
/** 计费模式(可选):per_second / per_video / token 等 */
|
||||
billing_mode?: string
|
||||
is_default?: boolean
|
||||
}
|
||||
|
||||
/** GET /viral-video/models 响应包装 */
|
||||
export interface ViralVideoModelsResponse {
|
||||
models: ViralVideoModel[]
|
||||
}
|
||||
|
||||
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
|
||||
export interface ConfirmCopyRequest {
|
||||
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
|
||||
edited_copy?: string
|
||||
/** 视频模型 key,覆盖默认 */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
|
||||
export interface StoryboardSegment {
|
||||
order: number
|
||||
type: string
|
||||
description: string
|
||||
text: string
|
||||
duration: number
|
||||
ken_burns?: string
|
||||
transition?: string
|
||||
}
|
||||
@@ -18,6 +18,8 @@ export interface VoiceClone {
|
||||
language: string
|
||||
gender: string
|
||||
error_message: string | null
|
||||
/** CosyVoice 实际使用的音色 ID(status=ready 时由后端填充,用于 TTS 调用) */
|
||||
voice_id?: string | null
|
||||
created_at: string
|
||||
updated_at: string
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ export const toVoiceClone = (profile: VoiceCloneProfile): VoiceClone => ({
|
||||
language: profile.language || "",
|
||||
gender: profile.gender || "",
|
||||
error_message: profile.error_message || null,
|
||||
voice_id: profile.voice_id,
|
||||
created_at: profile.created_at,
|
||||
updated_at: profile.updated_at,
|
||||
})
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
/* DurationWheelPicker —— 弹层式滚轮选择器(样式与表单一致) */
|
||||
|
||||
/* 触发按钮:外观复用 .vv-select 风格 */
|
||||
.dw-trigger {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
width: 100%;
|
||||
height: 36px;
|
||||
padding: 0 12px;
|
||||
background: #fff;
|
||||
border: 1px solid #e0e0e8;
|
||||
border-radius: 8px;
|
||||
font-size: 13px;
|
||||
color: #1f2937;
|
||||
cursor: pointer;
|
||||
box-sizing: border-box;
|
||||
transition: all 0.15s;
|
||||
user-select: none;
|
||||
}
|
||||
.dw-trigger:hover {
|
||||
border-color: #c0c0d0;
|
||||
}
|
||||
.dw-trigger-open,
|
||||
.dw-trigger:focus-within {
|
||||
border-color: #7c3aed !important;
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.12);
|
||||
}
|
||||
.dw-trigger-disabled {
|
||||
opacity: 0.5;
|
||||
pointer-events: none;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.dw-trigger-val {
|
||||
flex: 1;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.dw-trigger-placeholder {
|
||||
color: #9ca3af;
|
||||
}
|
||||
.dw-trigger-arrow {
|
||||
font-size: 10px;
|
||||
color: #9ca3af;
|
||||
margin-left: 8px;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
.dw-trigger-arrow-up {
|
||||
transform: rotate(180deg);
|
||||
}
|
||||
|
||||
/* 弹层容器 */
|
||||
.dw-popup {
|
||||
padding: 8px;
|
||||
min-width: 140px;
|
||||
}
|
||||
|
||||
/* 滚轮 */
|
||||
.dw-picker {
|
||||
position: relative;
|
||||
width: 100%;
|
||||
overflow: hidden;
|
||||
border-radius: 8px;
|
||||
background: #fafafe;
|
||||
border: 1px solid #e5e7eb;
|
||||
}
|
||||
.dw-picker-list {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
list-style: none;
|
||||
height: 100%;
|
||||
overflow-y: scroll;
|
||||
scroll-snap-type: y mandatory;
|
||||
-webkit-overflow-scrolling: touch;
|
||||
scrollbar-width: none;
|
||||
}
|
||||
.dw-picker-list::-webkit-scrollbar {
|
||||
display: none;
|
||||
}
|
||||
.dw-picker-item {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
justify-content: center;
|
||||
gap: 3px;
|
||||
scroll-snap-align: center;
|
||||
cursor: pointer;
|
||||
font-size: 15px;
|
||||
color: #9ca3af;
|
||||
font-weight: 400;
|
||||
transition:
|
||||
color 0.15s,
|
||||
transform 0.15s,
|
||||
font-weight 0.15s;
|
||||
}
|
||||
.dw-picker-item-val {
|
||||
font-variant-numeric: tabular-nums;
|
||||
}
|
||||
.dw-picker-item-unit {
|
||||
font-size: 13px;
|
||||
color: inherit;
|
||||
}
|
||||
.dw-picker-item-active {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-val {
|
||||
font-size: 18px;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-unit {
|
||||
font-size: 14px;
|
||||
}
|
||||
|
||||
/* 中心选中条 */
|
||||
.dw-picker-mask {
|
||||
position: absolute;
|
||||
left: 6px;
|
||||
right: 6px;
|
||||
pointer-events: none;
|
||||
background: #f5f0ff;
|
||||
border-radius: 6px;
|
||||
z-index: 1;
|
||||
}
|
||||
.dw-picker-mask::before,
|
||||
.dw-picker-mask::after {
|
||||
content: "";
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 1px;
|
||||
background: #d8c4ff;
|
||||
}
|
||||
.dw-picker-mask::before {
|
||||
top: 0;
|
||||
}
|
||||
.dw-picker-mask::after {
|
||||
bottom: 0;
|
||||
}
|
||||
|
||||
/* 上下渐变 */
|
||||
.dw-picker-fade {
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 40%;
|
||||
pointer-events: none;
|
||||
z-index: 2;
|
||||
}
|
||||
.dw-picker-fade-top {
|
||||
top: 0;
|
||||
background: linear-gradient(to bottom, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
.dw-picker-fade-bottom {
|
||||
bottom: 0;
|
||||
background: linear-gradient(to top, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
|
||||
/* 弹层按钮区 */
|
||||
.dw-popup-actions {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
justify-content: flex-end;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn {
|
||||
border-radius: 6px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary {
|
||||
background: #7c3aed;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary:hover {
|
||||
background: #6d28d9 !important;
|
||||
}
|
||||
|
||||
/* 覆盖 antd Popover 默认内边距 */
|
||||
.dw-popover .ant-popover-inner {
|
||||
padding: 0 !important;
|
||||
overflow: hidden;
|
||||
}
|
||||
.dw-popover .ant-popover-arrow {
|
||||
display: none;
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
/**
|
||||
* DurationWheelPicker —— 竖屏滚轮式时长选择器(弹层版)
|
||||
*
|
||||
* 设计:
|
||||
* - 外观是和其他表单 Select 一致的输入框(白色底+1px灰边+紫色focus ring)
|
||||
* - 点击输入框弹出 Popover,内部是滚轮 picker(原生 scroll-snap,零依赖)
|
||||
* - 滚轮样式:白底容器,选中行 #7c3aed 紫字加粗+浅紫背景条
|
||||
* - 支持触摸/鼠标滚轮/点击;松手吸附;底部"确认/取消"按钮
|
||||
* - 默认范围 15–30 秒,步长 1 秒
|
||||
*/
|
||||
import React, { useEffect, useMemo, useRef, useState, useCallback } from "react"
|
||||
import { Popover, Button } from "antd"
|
||||
import { DownOutlined } from "@ant-design/icons"
|
||||
import "./DurationWheelPicker.css"
|
||||
|
||||
export interface DurationWheelPickerProps {
|
||||
value?: number
|
||||
min?: number
|
||||
max?: number
|
||||
step?: number
|
||||
unit?: string
|
||||
onChange?: (value: number) => void
|
||||
placeholder?: string
|
||||
disabled?: boolean
|
||||
/** 弹层宽度,默认 160px */
|
||||
popupWidth?: number
|
||||
/** 弹层内滚轮高度,默认 180px */
|
||||
wheelHeight?: number
|
||||
}
|
||||
|
||||
const ITEM_HEIGHT = 36
|
||||
|
||||
const DurationWheelPicker: React.FC<DurationWheelPickerProps> = ({
|
||||
value = 20,
|
||||
min = 15,
|
||||
max = 30,
|
||||
step = 1,
|
||||
unit = "秒",
|
||||
onChange,
|
||||
placeholder = "请选择时长",
|
||||
disabled = false,
|
||||
popupWidth = 160,
|
||||
wheelHeight = 180,
|
||||
}) => {
|
||||
const options = useMemo(() => {
|
||||
const arr: number[] = []
|
||||
for (let v = min; v <= max; v += step) arr.push(v)
|
||||
return arr
|
||||
}, [min, max, step])
|
||||
|
||||
const [open, setOpen] = useState(false)
|
||||
// 弹层内暂存值,点确认才提交
|
||||
const [draft, setDraft] = useState<number>(value)
|
||||
const listRef = useRef<HTMLUListElement>(null)
|
||||
const scrollTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setDraft(value)
|
||||
// 下一帧滚到当前值
|
||||
requestAnimationFrame(() => scrollToValue(value, false))
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [open])
|
||||
|
||||
const scrollToValue = useCallback(
|
||||
(v: number, smooth = true) => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = options.indexOf(v)
|
||||
if (idx < 0) return
|
||||
list.scrollTo({ top: idx * ITEM_HEIGHT, behavior: smooth ? "smooth" : "auto" })
|
||||
},
|
||||
[options],
|
||||
)
|
||||
|
||||
const handleScroll = () => {
|
||||
if (scrollTimerRef.current) clearTimeout(scrollTimerRef.current)
|
||||
scrollTimerRef.current = setTimeout(() => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = Math.round(list.scrollTop / ITEM_HEIGHT)
|
||||
const clamped = Math.max(0, Math.min(options.length - 1, idx))
|
||||
const targetTop = clamped * ITEM_HEIGHT
|
||||
if (Math.abs(list.scrollTop - targetTop) > 1) {
|
||||
list.scrollTo({ top: targetTop, behavior: "smooth" })
|
||||
}
|
||||
setDraft(options[clamped])
|
||||
}, 100)
|
||||
}
|
||||
|
||||
const handleConfirm = () => {
|
||||
onChange?.(draft)
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleCancel = () => {
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleItemClick = (v: number) => {
|
||||
setDraft(v)
|
||||
scrollToValue(v, true)
|
||||
}
|
||||
|
||||
const maskTop = wheelHeight / 2 - ITEM_HEIGHT / 2
|
||||
|
||||
const wheel = (
|
||||
<div className="dw-popup">
|
||||
<div
|
||||
className="dw-picker"
|
||||
style={{ height: wheelHeight, width: popupWidth - 24 /* padding */ }}
|
||||
>
|
||||
<div className="dw-picker-mask" style={{ top: maskTop, height: ITEM_HEIGHT }} aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-top" aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-bottom" aria-hidden />
|
||||
<ul
|
||||
ref={listRef}
|
||||
className="dw-picker-list"
|
||||
onScroll={handleScroll}
|
||||
style={{
|
||||
paddingTop: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
paddingBottom: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
}}
|
||||
>
|
||||
{options.map((v) => {
|
||||
const isActive = v === draft
|
||||
return (
|
||||
<li
|
||||
key={v}
|
||||
className={`dw-picker-item${isActive ? " dw-picker-item-active" : ""}`}
|
||||
style={{ height: ITEM_HEIGHT, lineHeight: `${ITEM_HEIGHT}px` }}
|
||||
onClick={() => handleItemClick(v)}
|
||||
aria-selected={isActive}
|
||||
role="option"
|
||||
>
|
||||
<span className="dw-picker-item-val">{v}</span>
|
||||
<span className="dw-picker-item-unit">{unit}</span>
|
||||
</li>
|
||||
)
|
||||
})}
|
||||
</ul>
|
||||
</div>
|
||||
<div className="dw-popup-actions">
|
||||
<Button size="small" onClick={handleCancel}>
|
||||
取消
|
||||
</Button>
|
||||
<Button size="small" type="primary" onClick={handleConfirm}>
|
||||
确认
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
return (
|
||||
<Popover
|
||||
open={!disabled && open}
|
||||
onOpenChange={(v) => setOpen(v)}
|
||||
content={wheel}
|
||||
trigger="click"
|
||||
placement="bottomLeft"
|
||||
overlayClassName="dw-popover"
|
||||
overlayStyle={{ padding: 0 }}
|
||||
overlayInnerStyle={{ padding: 0, borderRadius: 10 }}
|
||||
destroyTooltipOnHide
|
||||
>
|
||||
<div
|
||||
className={`dw-trigger${disabled ? " dw-trigger-disabled" : ""}${open ? " dw-trigger-open" : ""}`}
|
||||
style={{ height: 36 }}
|
||||
>
|
||||
<span className={`dw-trigger-val${value != null ? "" : " dw-trigger-placeholder"}`}>
|
||||
{value != null ? `${value}${unit}` : placeholder}
|
||||
</span>
|
||||
<DownOutlined className={`dw-trigger-arrow${open ? " dw-trigger-arrow-up" : ""}`} />
|
||||
</div>
|
||||
</Popover>
|
||||
)
|
||||
}
|
||||
|
||||
export default DurationWheelPicker
|
||||
@@ -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": "成片库",
|
||||
|
||||
@@ -10,4 +10,4 @@
|
||||
* 功能流程不做积分预校验,直接走生成。
|
||||
* - true:展示完整积分系统 UI。
|
||||
*/
|
||||
export const ENABLE_CREDIT_SYSTEM = false
|
||||
export const ENABLE_CREDIT_SYSTEM = true
|
||||
|
||||
@@ -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),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
||||
@@ -80,6 +80,8 @@ const AiAvatarPage: React.FC = () => {
|
||||
const [finalizeLoading, setFinalizeLoading] = useState(false)
|
||||
|
||||
/* ── 对口型轮询 ── */
|
||||
/** 对口型轮询总时长上限(10分钟):超过后停止轮询并提示去历史记录查看 */
|
||||
const LIPSYNC_POLL_MAX_MS = 10 * 60 * 1000
|
||||
const lipsyncTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
/* ── 渲染进度轮询 ── */
|
||||
const renderTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
@@ -270,7 +272,24 @@ const AiAvatarPage: React.FC = () => {
|
||||
// 如果是预合成模式,后端会同步把状态置为 submitted(甚至可能已返回 running),
|
||||
// 但仍需轮询等 completed
|
||||
if (lipsyncTimerRef.current) clearInterval(lipsyncTimerRef.current)
|
||||
// 轮询间隔 5 秒;单请求超时 5 分钟(见 api/aiAvatar.ts);总轮询上限 10 分钟
|
||||
// 单次请求失败/超时不中断轮询,继续下一轮;超过总上限后停止并提示用户去历史记录查看
|
||||
lipsyncTimerRef.current = setInterval(async () => {
|
||||
// 总时长保护:超过 10 分钟停止轮询
|
||||
if (Date.now() - lipsyncStartAtRef.current > LIPSYNC_POLL_MAX_MS) {
|
||||
if (lipsyncTimerRef.current) {
|
||||
clearInterval(lipsyncTimerRef.current)
|
||||
lipsyncTimerRef.current = null
|
||||
}
|
||||
if (lipsyncTickRef.current) {
|
||||
clearInterval(lipsyncTickRef.current)
|
||||
lipsyncTickRef.current = null
|
||||
}
|
||||
setLipsyncStatus("failed")
|
||||
setLipsyncErrorMessage("渲染时间较长,请稍后在历史记录中查看")
|
||||
message.warning("对口型渲染时间较长,已停止自动刷新,请稍后在历史记录中查看")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const updated = await getLipsyncJob(job.id)
|
||||
state.setLipsyncJob(updated)
|
||||
@@ -296,9 +315,10 @@ const AiAvatarPage: React.FC = () => {
|
||||
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
|
||||
}
|
||||
} catch (err) {
|
||||
console.error("[对口型] 轮询错误:", err)
|
||||
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
|
||||
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
|
||||
}
|
||||
}, 3000)
|
||||
}, 5000)
|
||||
} catch (err) {
|
||||
console.error("[对口型] 创建失败:", {
|
||||
status: (err as { response?: { status?: number } })?.response?.status,
|
||||
@@ -580,20 +600,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 && (
|
||||
|
||||
@@ -72,7 +72,8 @@ export const previewTts = async (data: {
|
||||
}
|
||||
|
||||
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
|
||||
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -91,7 +92,10 @@ export const submitRender = async (data: {
|
||||
}
|
||||
|
||||
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
|
||||
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
|
||||
timeout: 300_000,
|
||||
})
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
@@ -523,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()
|
||||
|
||||
@@ -10,7 +10,7 @@ export interface BatchTaskState {
|
||||
taskId: string
|
||||
/** 变体序号(0-based,与标题/封面数组对齐) */
|
||||
variantIndex: number
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed"
|
||||
status: "running" | "completed" | "awaiting_cover" | "failed" | "queued"
|
||||
progress: number
|
||||
error: string | null
|
||||
/** 完成后的成片视频 */
|
||||
@@ -374,5 +374,42 @@ export function useGenerationPolling({
|
||||
}
|
||||
}, [])
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
/**
|
||||
* 批量队列模式:逐任务追加到轮询队列(支持串行提交、429 排队重试场景)。
|
||||
* 与 startPollingBatch 不同的是:
|
||||
* - 不会 reset batchContextRef;多次调用会累积
|
||||
* - 不触发整体 onComplete / onFailed(完成判定交给外层 useEffect 按状态聚合)
|
||||
* - 仍通过 onBatchTaskUpdate 回传单任务状态
|
||||
*/
|
||||
const pollBatchTaskQueued = useCallback(
|
||||
(taskId: string, variantIndex: number) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.set(taskId, variantIndex)
|
||||
onBatchTaskUpdate?.(taskId, {
|
||||
taskId,
|
||||
variantIndex,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
})
|
||||
pollSingleTask(taskId, Date.now(), {
|
||||
onTaskProgress: (pct) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: pct })
|
||||
},
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
const finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
onBatchTaskUpdate?.(taskId, { status: finalStatus, progress: 100, videos })
|
||||
},
|
||||
onTaskFailed: (msg) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "failed", error: msg })
|
||||
},
|
||||
}).catch(() => {
|
||||
/* onTaskFailed 已处理 */
|
||||
})
|
||||
},
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, pollBatchTaskQueued, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
import { useCallback, useState } from "react"
|
||||
import { message } from "antd"
|
||||
import { generateCover } from "@/api/generation"
|
||||
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
/** onCoversChange 支持直接传值或函数式 updater(函数式用于串行回写避免闭包覆盖) */
|
||||
@@ -182,15 +182,9 @@ export function useBatchCovers({
|
||||
async (index: number, file: File) => {
|
||||
addUploading(index)
|
||||
try {
|
||||
const libs = await getAssetLibraries()
|
||||
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
|
||||
if (!imageLib) {
|
||||
message.error("未找到素材库,请先创建")
|
||||
return
|
||||
}
|
||||
const result = await uploadAssetDirect({
|
||||
file,
|
||||
library_id: imageLib.id,
|
||||
kind: "image",
|
||||
})
|
||||
const url = result?.url || ""
|
||||
if (url) {
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
* 视频生成 Hook
|
||||
* 封装视频生成的核心逻辑、状态管理、轮询等
|
||||
*/
|
||||
import { useState, useCallback, useEffect } from "react"
|
||||
import { useState, useCallback, useEffect, useRef } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor"
|
||||
import { createGenerationTask } from "@/api/tasks/tasks"
|
||||
import type { CreateGenerationTaskRequest } from "@/api/tasks/types"
|
||||
import type { UseGenerateVideoProps } from "./generate-video/types"
|
||||
import { getGenerationPhase } from "./generate-video/phase"
|
||||
import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling"
|
||||
@@ -15,6 +17,26 @@ import { extractBackendError, translateError } from "./generate-video/errorUtils
|
||||
|
||||
export type GenerationCompleteStatus = "completed" | "awaiting_cover" | null
|
||||
|
||||
/** 判断是否是用户队列已满 429(需要排队重试而非直接报错) */
|
||||
function isUserQueueFullError(err: unknown): { waitMs: number } | null {
|
||||
if (!axios.isAxiosError(err)) return null
|
||||
if (err.response?.status !== 429 && err.response?.status !== 503) return null
|
||||
const detail = (err.response?.data as { detail?: unknown })?.detail
|
||||
const code =
|
||||
typeof detail === "object" && detail !== null ? (detail as { code?: string }).code : undefined
|
||||
if (code === "USER_QUEUE_FULL" || code === "SYSTEM_QUEUE_FULL") {
|
||||
const waitSec =
|
||||
typeof detail === "object" && detail !== null
|
||||
? Number((detail as { estimated_wait_seconds?: number }).estimated_wait_seconds) || 0
|
||||
: 0
|
||||
return { waitMs: Math.max(15_000, waitSec * 1000 || 30_000) }
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/** sleep */
|
||||
const sleep = (ms: number) => new Promise<void>((r) => setTimeout(r, ms))
|
||||
|
||||
export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { selectedTemplate, onGenerationSuccess } = props
|
||||
|
||||
@@ -31,6 +53,22 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
|
||||
|
||||
/** 排队中重试的定时器,unmount / 新提交时清理 */
|
||||
const queueTimersRef = useRef<number[]>([])
|
||||
const cancelledRef = useRef(false)
|
||||
|
||||
const clearQueueTimers = useCallback(() => {
|
||||
queueTimersRef.current.forEach((id) => clearTimeout(id))
|
||||
queueTimersRef.current = []
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
cancelledRef.current = true
|
||||
clearQueueTimers()
|
||||
}
|
||||
}, [clearQueueTimers])
|
||||
|
||||
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
|
||||
setBatchTasks((prev) => {
|
||||
const list = prev || []
|
||||
@@ -58,7 +96,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
(videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => {
|
||||
setGenerating(false)
|
||||
setGenerated(true)
|
||||
const finalStatus: GenerationCompleteStatus = taskStatus ?? "completed"
|
||||
setCompletionStatus(finalStatus)
|
||||
@@ -81,21 +118,30 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
[onGenerationSuccess],
|
||||
)
|
||||
const handleFailed = useCallback((errorMsg: string) => {
|
||||
setGenerating(false)
|
||||
setGenerateError(errorMsg)
|
||||
}, [])
|
||||
|
||||
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频。
|
||||
当全部任务都已结束(completed/awaiting_cover/failed)且无排队/渲染中任务时,关闭 generating。 */
|
||||
useEffect(() => {
|
||||
if (batchTasks.length === 0) return
|
||||
const byVariant = new Map<number, GeneratedVideo>()
|
||||
let hasQueued = false
|
||||
let hasRunning = false
|
||||
let hasSuccess = false
|
||||
let allDone = true
|
||||
batchTasks.forEach((t) => {
|
||||
if (
|
||||
t.status === "completed" ||
|
||||
(t.status === "awaiting_cover" && t.videos && t.videos.length > 0)
|
||||
) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
if (t.status === "queued") hasQueued = true
|
||||
else if (t.status === "running") hasRunning = true
|
||||
if (t.status === "completed" || t.status === "awaiting_cover") {
|
||||
hasSuccess = true
|
||||
if (t.videos && t.videos.length > 0) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
}
|
||||
}
|
||||
if (t.status !== "completed" && t.status !== "awaiting_cover" && t.status !== "failed") {
|
||||
allDone = false
|
||||
}
|
||||
})
|
||||
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
|
||||
@@ -105,17 +151,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
return ordered
|
||||
})
|
||||
if (allDone && !hasQueued && !hasRunning) {
|
||||
setGenerating(false)
|
||||
if (hasSuccess) {
|
||||
setGenerated(true)
|
||||
setCompletionStatus("awaiting_cover")
|
||||
}
|
||||
}
|
||||
}, [batchTasks])
|
||||
|
||||
const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({
|
||||
const { startPolling, pollBatchTaskQueued, retryTask, clearTimer } = useGenerationPolling({
|
||||
onProgress: handleProgress,
|
||||
onComplete: handleComplete,
|
||||
onFailed: handleFailed,
|
||||
onBatchTaskUpdate: handleBatchTaskUpdate,
|
||||
})
|
||||
|
||||
/** 根据 props 构造基础 payload(批量/单任务共用的字段) */
|
||||
const buildBasePayload = useCallback((): Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> => {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
const bgmConfig = {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
}
|
||||
|
||||
const titleConfig = props.titleSettings?.title
|
||||
? {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color: props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
}
|
||||
: undefined
|
||||
|
||||
const payload: Omit<
|
||||
CreateGenerationTaskRequest,
|
||||
"count" | "titles" | "voice_library_ids" | "cover_urls" | "variant_plan_ids"
|
||||
> = {
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: bgmConfig as CreateGenerationTaskRequest["bgm_config"],
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(titleConfig ? ({ title_config: titleConfig } as Record<string, unknown>) : {}),
|
||||
}
|
||||
|
||||
return payload
|
||||
}, [props, selectedTemplate])
|
||||
|
||||
/* ── 生成视频 ──
|
||||
返回 true 表示任务创建成功并已开始轮询;false 表示校验未通过或创建失败 */
|
||||
返回 true 表示任务创建成功并已开始轮询(含排队中);false 表示校验未通过或创建失败 */
|
||||
const generate = useCallback(async (): Promise<boolean> => {
|
||||
const errorMsg = validateGenerateInputs(props)
|
||||
if (errorMsg) {
|
||||
@@ -123,6 +337,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
return false
|
||||
}
|
||||
|
||||
cancelledRef.current = false
|
||||
clearQueueTimers()
|
||||
setGenerating(true)
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
@@ -133,25 +349,16 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setCurrentTaskId("")
|
||||
clearTimer()
|
||||
|
||||
const basePayload = buildBasePayload()
|
||||
const assetIds = basePayload.asset_ids
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
|
||||
try {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
|
||||
// from-assets 已由 useStep2Materials 在用户选素材时(debounce 800ms)调用,
|
||||
// 后端已改为异步秒级返回,这里做一次轻量兜底:
|
||||
// 单次查 clips,已有则直接放行;没有则再调一次 from-assets。
|
||||
// from-assets 兜底:片段不存在则补一次
|
||||
if (assetIds.length > 0 && selectedTemplate) {
|
||||
try {
|
||||
const clipList = await getEditPlanClips(selectedTemplate, { limit: 500 })
|
||||
if (clipList.items.length === 0) {
|
||||
// 片段不存在(极端情况:useStep2Materials 的 debounce 还没触发)
|
||||
// 手动补一次 from-assets(后端秒级返回)
|
||||
await createClipsFromAssets(selectedTemplate, assetIds, "main")
|
||||
}
|
||||
} catch {
|
||||
@@ -159,221 +366,185 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
}
|
||||
}
|
||||
|
||||
const isBatch = (props.previewCount || 1) > 1
|
||||
const hide = message.loading(
|
||||
isBatch ? `正在生成 ${props.previewCount} 个视频...` : "正在生成预览视频...",
|
||||
0,
|
||||
)
|
||||
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
|
||||
const indexes =
|
||||
isBatch && props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = isBatch ? indexes.length : 1
|
||||
|
||||
// 标题文字数组:批量时按勾选顺序
|
||||
const titlesArr =
|
||||
isBatch && (props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: []
|
||||
// 配音数组:独立配音模式按勾选顺序;否则不传(回退共用 voice_library_id)
|
||||
const voiceArr =
|
||||
isBatch && props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map((i) => props.variantVoiceLibraryIds![i] || voiceLibraryId)
|
||||
: []
|
||||
// 封面数组:批量时按勾选顺序(未设置封面的变体传空串,后端回退智能封面)
|
||||
const coversArr =
|
||||
isBatch && props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: []
|
||||
// #1744 变体 plan 数组:预览阶段后端独立选片产出的 plan id,按勾选顺序回传,
|
||||
// 后端直接关联这些 plan 渲染(不再重新选片)→ 预览所见即成片。
|
||||
// 全部为空(降级本地模拟/后端端点未上线)时不传,后端走自身独立选片。
|
||||
const variantPlansArr =
|
||||
isBatch && props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: []
|
||||
const hasVariantPlans = variantPlansArr.some((id) => !!id)
|
||||
|
||||
try {
|
||||
const taskResp = await createGenerationTask({
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: assetIds,
|
||||
output_width: outputWidth,
|
||||
output_height: outputHeight,
|
||||
cover_url: coverUrl,
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
tts_style: props.ttsStyle || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: {
|
||||
enabled: props.bgm !== false,
|
||||
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
|
||||
},
|
||||
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
|
||||
...(isBatch ? { count: batchCount } : {}),
|
||||
...(titlesArr.length ? { titles: titlesArr } : {}),
|
||||
...(voiceArr.length ? { voice_library_ids: voiceArr } : {}),
|
||||
...(coversArr.length ? { cover_urls: coversArr } : {}),
|
||||
...(hasVariantPlans ? { variant_plan_ids: variantPlansArr } : {}),
|
||||
...(props.titleSettings?.title
|
||||
? {
|
||||
title_config: {
|
||||
text: props.titleSettings.title,
|
||||
font: props.titleSettings.font,
|
||||
font_size: props.titleSettings.size,
|
||||
font_color: props.titleSettings.color,
|
||||
position: props.titleSettings.position,
|
||||
...(props.titleSettings.position === "custom" &&
|
||||
props.titleSettings.posX != null &&
|
||||
props.titleSettings.posY != null
|
||||
? {
|
||||
pos_x: Math.round(props.titleSettings.posX),
|
||||
pos_y: Math.round(props.titleSettings.posY),
|
||||
}
|
||||
: {}),
|
||||
bold: props.titleSettings.bold,
|
||||
italic: props.titleSettings.italic,
|
||||
stroke: props.titleSettings.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.strokeWidth ?? 4,
|
||||
color: props.titleSettings.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.shadowBlur ?? 4,
|
||||
color: props.titleSettings.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
line_height: props.titleSettings.lineHeight ?? 1.2,
|
||||
margin_top: props.titleSettings.marginTop ?? 24,
|
||||
max_chars_per_line: props.titleSettings.maxCharsPerLine ?? 0,
|
||||
...(props.titleSettings.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.bgColor,
|
||||
padding: props.titleSettings.bgPadding,
|
||||
radius: props.titleSettings.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
line_overrides: (props.titleSettings.lineOverrides ?? []).map((lo) => ({
|
||||
line_index: lo.line_index,
|
||||
text: lo.text,
|
||||
size: lo.size,
|
||||
color: lo.color,
|
||||
bold: lo.bold,
|
||||
italic: lo.italic,
|
||||
stroke: lo.stroke,
|
||||
highlights: lo.highlights?.map((h) => ({
|
||||
word: h.word,
|
||||
color: h.color,
|
||||
bold: h.bold,
|
||||
scale: h.scale,
|
||||
})),
|
||||
})),
|
||||
...(props.titleSettings.coverTitle
|
||||
? {
|
||||
cover_title_config: {
|
||||
title: props.titleSettings.coverTitle.title,
|
||||
font: props.titleSettings.coverTitle.font,
|
||||
font_size: props.titleSettings.coverTitle.size,
|
||||
font_color: props.titleSettings.coverTitle.color,
|
||||
bold: props.titleSettings.coverTitle.bold,
|
||||
italic: props.titleSettings.coverTitle.italic,
|
||||
position: props.titleSettings.coverTitle.position,
|
||||
stroke: props.titleSettings.coverTitle.stroke
|
||||
? {
|
||||
enabled: true,
|
||||
width: props.titleSettings.coverTitle.strokeWidth ?? 4,
|
||||
color: props.titleSettings.coverTitle.strokeColor ?? "#000000",
|
||||
}
|
||||
: { enabled: false },
|
||||
shadow: props.titleSettings.coverTitle.shadow
|
||||
? {
|
||||
enabled: true,
|
||||
offset_x: props.titleSettings.coverTitle.shadowOffsetX ?? 2,
|
||||
offset_y: props.titleSettings.coverTitle.shadowOffsetY ?? 2,
|
||||
blur: props.titleSettings.coverTitle.shadowBlur ?? 4,
|
||||
color:
|
||||
props.titleSettings.coverTitle.shadowColor ?? "rgba(0,0,0,0.8)",
|
||||
}
|
||||
: { enabled: false },
|
||||
...(props.titleSettings.coverTitle.bgEnabled
|
||||
? {
|
||||
background: {
|
||||
enabled: true,
|
||||
color: props.titleSettings.coverTitle.bgColor,
|
||||
padding: props.titleSettings.coverTitle.bgPadding,
|
||||
radius: props.titleSettings.coverTitle.bgRadius,
|
||||
},
|
||||
}
|
||||
: { background: { enabled: false } }),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
})
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
if (taskIds.length > 1) {
|
||||
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
|
||||
setCurrentTaskId("")
|
||||
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
|
||||
} else {
|
||||
if (!isBatch) {
|
||||
/* ── 单视频:原逻辑(一次提交 count=1) ── */
|
||||
const hide = message.loading("正在生成预览视频...", 0)
|
||||
try {
|
||||
const taskResp = await createGenerationTask({ ...basePayload, count: 1 })
|
||||
hide()
|
||||
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
|
||||
if (taskIds.length === 0) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
setCurrentTaskId(taskIds[0])
|
||||
startPolling(taskIds[0])
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
}
|
||||
} catch (err) {
|
||||
hide()
|
||||
throw err
|
||||
return true
|
||||
}
|
||||
|
||||
/* ── 批量:支持任意数量视频,按队列容量串行提交,429 自动排队重试 ── */
|
||||
const indexes = props.selectedVariantIndexes?.length
|
||||
? props.selectedVariantIndexes
|
||||
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
|
||||
const batchCount = indexes.length
|
||||
|
||||
const titlesAll =
|
||||
(props.variantTitles?.length || 0) >= batchCount
|
||||
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
|
||||
: indexes.map(() => props.titleSettings?.title || "")
|
||||
const voiceArrAll =
|
||||
props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
|
||||
? indexes.map(
|
||||
(i) => props.variantVoiceLibraryIds![i] || basePayload.voice_library_id || "",
|
||||
)
|
||||
: []
|
||||
const coversAll = props.variantCoverUrls?.length
|
||||
? indexes.map((i) => props.variantCoverUrls![i] || "")
|
||||
: indexes.map(() => "")
|
||||
const plansAll = props.variantPlanIds?.length
|
||||
? indexes.map((i) => props.variantPlanIds![i] || "")
|
||||
: indexes.map(() => "")
|
||||
|
||||
const hasAnyVoice = voiceArrAll.some((v) => !!v)
|
||||
const hasAnyCover = coversAll.some((u) => !!u)
|
||||
const hasAnyPlan = plansAll.some((id) => !!id)
|
||||
|
||||
// 先用占位 ID 把所有变体卡片置为 queued,UI 可见
|
||||
const placeholderIds = indexes.map((_, i) => `__queued_${Date.now()}_${i}`)
|
||||
const initialTasks: BatchTaskState[] = indexes.map((variantIndex, i) => ({
|
||||
taskId: placeholderIds[i],
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
}))
|
||||
setBatchTasks(initialTasks)
|
||||
|
||||
message.loading({
|
||||
content: `已提交 ${batchCount} 个视频任务,系统按队列容量依次渲染…`,
|
||||
key: "batch-gen",
|
||||
duration: 3,
|
||||
})
|
||||
|
||||
/** 将占位 taskId 更新为真实 taskId(卡片引用同一对象) */
|
||||
const replacePlaceholder = (placeholderId: string, realTaskId: string) => {
|
||||
setBatchTasks((prev) => {
|
||||
const idx = prev.findIndex((t) => t.taskId === placeholderId)
|
||||
if (idx === -1) return prev
|
||||
const next = [...prev]
|
||||
next[idx] = { ...next[idx], taskId: realTaskId }
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
/** 提交某一索引的单任务(count=1),成功后返回真实 taskId;429/503 则返回 waitMs */
|
||||
const submitOne = async (
|
||||
i: number,
|
||||
): Promise<{ queued: true; waitMs: number } | { queued: false; taskId: string }> => {
|
||||
const body: CreateGenerationTaskRequest = {
|
||||
...basePayload,
|
||||
count: 1,
|
||||
titles: [titlesAll[i] || ""],
|
||||
...(hasAnyVoice
|
||||
? { voice_library_ids: [voiceArrAll[i] || basePayload.voice_library_id || ""] }
|
||||
: {}),
|
||||
...(hasAnyCover ? { cover_urls: [coversAll[i] || ""] } : {}),
|
||||
...(hasAnyPlan && plansAll[i] ? { variant_plan_ids: [plansAll[i]] } : {}),
|
||||
}
|
||||
try {
|
||||
const resp = await createGenerationTask(body)
|
||||
const item = resp.items?.[0]
|
||||
const tid = item?.id
|
||||
if (!tid) throw new Error("创建任务成功但未返回任务 ID")
|
||||
return { queued: false, taskId: tid }
|
||||
} catch (err) {
|
||||
const q = isUserQueueFullError(err)
|
||||
if (q) return { queued: true, waitMs: q.waitMs }
|
||||
throw err
|
||||
}
|
||||
}
|
||||
|
||||
// 串行提交:每次提交一个;429/503 则等待后重试;其它错误立即标记该任务失败
|
||||
let fatalErr: unknown = null
|
||||
for (let i = 0; i < batchCount; i++) {
|
||||
if (cancelledRef.current) return false
|
||||
const variantIndex = indexes[i]
|
||||
const placeholderId = placeholderIds[i]
|
||||
let attempt = 0
|
||||
let submitted = false
|
||||
while (!submitted) {
|
||||
if (cancelledRef.current) return false
|
||||
attempt++
|
||||
try {
|
||||
const result = await submitOne(i)
|
||||
if (!result.queued) {
|
||||
replacePlaceholder(placeholderId, result.taskId)
|
||||
// 先更新到 running,再启动单任务增量轮询(不触发整体 onComplete)
|
||||
pollBatchTaskQueued(result.taskId, variantIndex)
|
||||
submitted = true
|
||||
} else {
|
||||
// 排队:保持 queued 状态,等待后重试
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "queued",
|
||||
progress: 0,
|
||||
error: null,
|
||||
})
|
||||
if (attempt === 1) {
|
||||
message.info({
|
||||
content: `队列繁忙,${Math.round(result.waitMs / 1000)} 秒后自动继续提交后续视频…`,
|
||||
key: "batch-gen",
|
||||
duration: 4,
|
||||
})
|
||||
}
|
||||
await sleep(Math.min(result.waitMs, 60_000))
|
||||
}
|
||||
} catch (err) {
|
||||
// 非限流错误:该任务标记失败,继续后续任务(不阻断整个批量)
|
||||
console.error("[batch generate] 任务提交失败:", err)
|
||||
const msg = translateError(extractBackendError(err))
|
||||
handleBatchTaskUpdate(placeholderId, {
|
||||
taskId: placeholderId,
|
||||
variantIndex,
|
||||
status: "failed",
|
||||
error: msg,
|
||||
progress: 0,
|
||||
})
|
||||
submitted = true
|
||||
if (!fatalErr) fatalErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (fatalErr) {
|
||||
// 有任务失败但其余已成功,整体不 throw;由 UI 展示单个失败卡片
|
||||
}
|
||||
return true
|
||||
} catch (err: unknown) {
|
||||
console.error("[handleGenerate] 生成失败:", err)
|
||||
setGenerating(false)
|
||||
const backendMsg = extractBackendError(err)
|
||||
console.error("[handleGenerate] 错误信息:", backendMsg, "完整错误:", err)
|
||||
const finalMsg = translateError(backendMsg)
|
||||
setGenerateError(finalMsg)
|
||||
setGenerating(false)
|
||||
message.error(finalMsg)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
|
||||
}, [
|
||||
props,
|
||||
clearTimer,
|
||||
startPolling,
|
||||
selectedTemplate,
|
||||
buildBasePayload,
|
||||
handleBatchTaskUpdate,
|
||||
clearQueueTimers,
|
||||
pollBatchTaskQueued,
|
||||
])
|
||||
|
||||
const retry = useCallback(() => {
|
||||
setGenerateError(null)
|
||||
@@ -383,9 +554,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
/** 第5步:单独重试某个失败任务 */
|
||||
const retryBatchTask = useCallback(
|
||||
(taskId: string) => {
|
||||
handleBatchTaskUpdate(taskId, { status: "running", progress: 0, error: null, videos: [] })
|
||||
retryTask(taskId)
|
||||
},
|
||||
[retryTask],
|
||||
[retryTask, handleBatchTaskUpdate],
|
||||
)
|
||||
|
||||
const dismissError = useCallback(() => {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,263 @@
|
||||
/**
|
||||
* 爆款视频素材选择弹窗(通用版,支持 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" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary"
|
||||
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,
|
||||
},
|
||||
|
||||
@@ -553,7 +553,20 @@ def concat_video_files(
|
||||
if work_dir is None:
|
||||
work_dir = output_path.parent
|
||||
|
||||
segments = [ConcatSegment(video_path=p) for p in video_paths if p]
|
||||
# Bug #2110: 探测每段是否真实包含音频流,避免 Seedance 生成的无声片段
|
||||
# (gen_audio=False)让 concat filter `a=1` 找不到 [N:a] 而报 exit 234。
|
||||
from video_processing.ffmpeg_utils import probe_has_audio as _probe_has_audio
|
||||
|
||||
segments: list[ConcatSegment] = []
|
||||
for p in video_paths:
|
||||
if not p:
|
||||
continue
|
||||
try:
|
||||
has_audio = _probe_has_audio(p)
|
||||
except Exception:
|
||||
has_audio = True # 探测失败保守认为有音频
|
||||
segments.append(ConcatSegment(video_path=p, has_audio=has_audio))
|
||||
|
||||
config = ConcatConfig(segments=segments, force_reencode=force_reencode)
|
||||
|
||||
engine = ConcatEngine(work_dir)
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""爆款视频 Worker 侧模块(#2039/#2040/#2051)。
|
||||
|
||||
video_analyzer(#2051):参考视频风格分析 6 步管线,输出 style_guide + clips 渲染参数映射。
|
||||
#2040 的 prompt 系统(prompts/prompt_store/llm_runner)由 #2040 分支提供,本文件不依赖它。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from apps.worker.viral_video.video_analyzer import (
|
||||
DEFAULT_ANALYSIS_TIMEOUT,
|
||||
MAX_REFERENCE_DURATION_SEC,
|
||||
MAX_REFERENCE_SIZE_MB,
|
||||
STYLE_GUIDE_SCHEMA,
|
||||
analyze_video_style,
|
||||
build_render_params_for_clip,
|
||||
map_bgm_bpm,
|
||||
map_camera_to_ken_burns,
|
||||
map_color_to_video_filter,
|
||||
map_transition_to_xfade,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_ANALYSIS_TIMEOUT",
|
||||
"MAX_REFERENCE_DURATION_SEC",
|
||||
"MAX_REFERENCE_SIZE_MB",
|
||||
"STYLE_GUIDE_SCHEMA",
|
||||
"analyze_video_style",
|
||||
"build_render_params_for_clip",
|
||||
"map_bgm_bpm",
|
||||
"map_camera_to_ken_burns",
|
||||
"map_color_to_video_filter",
|
||||
"map_transition_to_xfade",
|
||||
]
|
||||
@@ -0,0 +1,961 @@
|
||||
"""参考爆款视频风格分析模块(#2051,v1.3)。
|
||||
|
||||
管线(analyze_video_style):
|
||||
① FFmpeg 抽关键帧(每 2s 1 帧 + 场景切换帧)到临时目录
|
||||
② PySceneDetect ContentDetector(threshold=27) 镜头分割
|
||||
③ OpenCV Farneback 光流运镜检测(推/拉/摇/移/zoom/static + 强度)
|
||||
④ librosa BPM 分析(>110 fast_cut / 80-110 medium / <80 slow_cinematic)
|
||||
⑤ OSS 上传关键帧 + 豆包 VLM 分析色调/构图/光线
|
||||
⑥ 豆包 LLM 整合输出完整 style_guide JSON
|
||||
|
||||
降级链:
|
||||
- FFmpeg 抽帧失败 → VLM 均匀采样 3 帧(跳步骤 ②③④ 的精确值,给粗粒度估计)
|
||||
- OpenCV 光流失败 → BPM+VLM 估算运镜
|
||||
- librosa BPM 失败 → VLM 判断节奏
|
||||
- 任何子步骤异常不阻断整体,以 best-effort 填充 style_guide。
|
||||
|
||||
资源约束:
|
||||
- 参考视频 ≤60s 且 ≤100MB;分析总超时 ≤60s;临时帧 try/finally 清理。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
import tempfile
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 资源约束 ─────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_REFERENCE_DURATION_SEC = 60
|
||||
MAX_REFERENCE_SIZE_MB = 100
|
||||
DEFAULT_ANALYSIS_TIMEOUT = 60 # 秒
|
||||
KEYFRAME_INTERVAL_SEC = 2
|
||||
SCENEDETECT_THRESHOLD = 27
|
||||
VLM_SAMPLE_FRAMES = 5 # 上传给 VLM 的关键帧上限
|
||||
FARNEBACK_PARAMS = dict(pyr_scale=0.5, levels=3, winsize=15, iterations=3, poly_n=5, poly_sigma=1.2, flags=0)
|
||||
|
||||
# ── style_guide 输出 schema(最小校验参考,不强制 jsonschema 依赖) ───────
|
||||
|
||||
STYLE_GUIDE_SCHEMA: dict[str, Any] = {
|
||||
"style_name": str,
|
||||
"avg_shot_duration": float,
|
||||
"shot_count": int,
|
||||
"pace": str, # fast_cut | medium | slow_cinematic
|
||||
"bpm": int,
|
||||
"camera_movements": list,
|
||||
"transitions": list,
|
||||
"color_palette": list,
|
||||
"color_tone": str, # warm | cool | high_sat | low_sat | vintage | fresh | dramatic | bright
|
||||
"color_filter": str, # none | warm_vintage | cool_fresh | high_contrast | soft_pastel | dramatic_cinematic
|
||||
"composition": dict,
|
||||
"lighting": str,
|
||||
"mood": str,
|
||||
"visual_keywords": list,
|
||||
"ken_burns_params": dict,
|
||||
"transition_map": dict,
|
||||
"video_filter_eq_params": dict,
|
||||
"ken_burns_direction_hint": str,
|
||||
}
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ShotBoundary:
|
||||
"""一段镜头(帧号区间)。"""
|
||||
|
||||
index: int
|
||||
start_sec: float
|
||||
end_sec: float
|
||||
movement: str = (
|
||||
"static" # push_in | pull_out | pan_left | pan_right | tilt_up | tilt_down | static | zoom_in | zoom_out
|
||||
)
|
||||
intensity: str = "low" # low | medium | high
|
||||
transition: str = "hard_cut" # 到下一个镜头的转场
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnalysisArtifacts:
|
||||
"""中间产物(降级路径用)。"""
|
||||
|
||||
frames_dir: Path
|
||||
frame_paths: list[Path] = field(default_factory=list)
|
||||
shots: list[ShotBoundary] = field(default_factory=list)
|
||||
bpm: int = 0
|
||||
vlm_descriptions: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _strip_code_fence(text: str) -> str:
|
||||
"""移除 markdown 代码块围栏,返回纯文本。"""
|
||||
t = text.strip()
|
||||
for fence in ("```json", "```JSON", "```"):
|
||||
if t.startswith(fence):
|
||||
t = t[len(fence) :].lstrip()
|
||||
if t.endswith("```"):
|
||||
t = t[:-3].rstrip()
|
||||
return t
|
||||
|
||||
|
||||
# ── FFmpeg / ffprobe ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _ffmpeg_bin() -> str:
|
||||
return shutil.which("ffmpeg") or "ffmpeg"
|
||||
|
||||
|
||||
def _ffprobe_bin() -> str:
|
||||
return shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
def _probe_duration(video_path: str | Path) -> float:
|
||||
"""用 ffprobe 取视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
_ffprobe_bin(),
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(video_path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
text=True,
|
||||
) # nosec B603
|
||||
return float(out.strip() or 0)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("ffprobe 时长探测失败 %s: %s", video_path, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _extract_keyframes(video_path: Path, out_dir: Path, interval: int = KEYFRAME_INTERVAL_SEC) -> list[Path]:
|
||||
"""按固定间隔抽帧;同时检测场景切换帧(select='gt(scene,...)')。"""
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
# 固定间隔
|
||||
fixed_tpl = str(out_dir / "f_%04d.jpg")
|
||||
cmd_fixed = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
f"fps=1/{interval}",
|
||||
"-q:v",
|
||||
"3",
|
||||
fixed_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_fixed, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
# 场景切换帧(独立命名,scene_ 前缀)
|
||||
scene_tpl = str(out_dir / "scene_%04d.jpg")
|
||||
cmd_scene = [
|
||||
_ffmpeg_bin(),
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vf",
|
||||
"select='gt(scene,0.35)',showinfo",
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
"3",
|
||||
scene_tpl,
|
||||
]
|
||||
subprocess.run(
|
||||
cmd_scene, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DEFAULT_ANALYSIS_TIMEOUT, check=False
|
||||
) # nosec B603
|
||||
frames = sorted(out_dir.glob("f_*.jpg")) + sorted(out_dir.glob("scene_*.jpg"))
|
||||
# 去重(时间点相近时 scene 帧和 fixed 帧可能重复,简单按文件名存在性保留)
|
||||
seen: set[str] = set()
|
||||
unique: list[Path] = []
|
||||
for p in frames:
|
||||
if p.name not in seen:
|
||||
seen.add(p.name)
|
||||
unique.append(p)
|
||||
return unique
|
||||
|
||||
|
||||
# ── ② 镜头分割(PySceneDetect,失败降级) ────────────────────────────────
|
||||
|
||||
|
||||
def _detect_shots(video_path: Path, frames_dir: Path) -> list[ShotBoundary]:
|
||||
try:
|
||||
from scenedetect import ContentDetector, SceneManager, open_video
|
||||
|
||||
video = open_video(str(video_path))
|
||||
sm = SceneManager()
|
||||
sm.add_detector(ContentDetector(threshold=SCENEDETECT_THRESHOLD))
|
||||
sm.detect_scenes(video)
|
||||
scenes = sm.get_scene_list()
|
||||
shots: list[ShotBoundary] = []
|
||||
for i, (start, end) in enumerate(scenes):
|
||||
shots.append(
|
||||
ShotBoundary(
|
||||
index=i,
|
||||
start_sec=start.get_seconds(),
|
||||
end_sec=end.get_seconds(),
|
||||
)
|
||||
)
|
||||
if shots:
|
||||
return shots
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("PySceneDetect 镜头分割失败,使用均匀分段降级: %s", exc)
|
||||
# 降级:按固定间隔每 3 秒一镜头
|
||||
duration = _probe_duration(video_path) or 15.0
|
||||
dur = max(3.0, min(duration, float(MAX_REFERENCE_DURATION_SEC)))
|
||||
shots = []
|
||||
seg = 3.0
|
||||
i = 0
|
||||
t = 0.0
|
||||
while t < dur - 0.1:
|
||||
shots.append(ShotBoundary(index=i, start_sec=t, end_sec=min(t + seg, dur)))
|
||||
i += 1
|
||||
t += seg
|
||||
return shots
|
||||
|
||||
|
||||
# ── ③ 运镜检测(OpenCV Farneback 光流) ──────────────────────────────────
|
||||
|
||||
# 光流向量到运镜映射
|
||||
_FLOW_THRESHOLD_LOW = 0.3
|
||||
_FLOW_THRESHOLD_HIGH = 1.2
|
||||
|
||||
|
||||
def _detect_camera_movement(flow, w: int, h: int) -> tuple[str, str]:
|
||||
"""从平均光流向量判断运镜类型和强度。"""
|
||||
import numpy as np # noqa: PLC0415 - numpy 已在 requirements 中
|
||||
|
||||
fx = float(np.median(flow[..., 0]))
|
||||
fy = float(np.median(flow[..., 1]))
|
||||
trans_mag = math.hypot(fx, fy)
|
||||
# 发散/收敛判断 zoom:比较边缘流沿径向外指的平均分量(稳健版)
|
||||
cx, cy = w / 2.0, h / 2.0
|
||||
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
||||
rx, ry = (xs - cx) / max(cx, 1.0), (ys - cy) / max(cy, 1.0)
|
||||
rmag = np.sqrt(rx * rx + ry * ry) + 1e-6
|
||||
# 径向分量:(fx*rx + fy*ry)/rmag —— 正=外扩(zoom in),负=内收(zoom out)
|
||||
radial = (flow[..., 0] * rx + flow[..., 1] * ry) / rmag
|
||||
# 只看边缘带(|r|>0.5),且减去平移贡献:径向减去平均平移投影
|
||||
edge_mask = (rmag > 0.5).astype(np.float32)
|
||||
if edge_mask.sum() > 10:
|
||||
trans_radial = (fx * rx + fy * ry) / rmag
|
||||
zoom_signal = float(np.mean((radial - trans_radial)[edge_mask > 0]))
|
||||
else:
|
||||
zoom_signal = 0.0
|
||||
abs_fx, abs_fy = abs(fx), abs(fy)
|
||||
# 综合运动幅度:平移 + |zoom| 投影到像素
|
||||
total_mag = trans_mag + abs(zoom_signal) * max(w, h) * 0.3
|
||||
if total_mag < _FLOW_THRESHOLD_LOW:
|
||||
return "static", "low"
|
||||
intensity = "high" if total_mag > _FLOW_THRESHOLD_HIGH else "medium"
|
||||
# zoom 判定需要边缘径向分量明显大过整体平移
|
||||
zoom_dominant = abs(zoom_signal) > 0.6 and abs(zoom_signal) * max(w, h) * 0.3 > trans_mag * 1.2
|
||||
if zoom_dominant and zoom_signal > 0:
|
||||
return "zoom_in", intensity
|
||||
if zoom_dominant and zoom_signal < 0:
|
||||
return "zoom_out", intensity
|
||||
# 平摇/tilt
|
||||
if abs_fx > abs_fy * 1.5:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
if abs_fy > abs_fx * 1.5:
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
# 轨道/跟拍:以主轴为主
|
||||
if abs_fx >= abs_fy:
|
||||
return "pan_right" if fx > 0 else "pan_left", intensity
|
||||
return "tilt_down" if fy > 0 else "tilt_up", intensity
|
||||
|
||||
|
||||
def _analyze_movements(video_path: Path, shots: list[ShotBoundary]) -> None:
|
||||
"""对每个 shot 的首尾帧算光流,填充 movement/intensity。失败时静默降级为 static/low。"""
|
||||
try:
|
||||
import cv2 # noqa: PLC0415 - opencv-python-headless 已在 worker requirements 中
|
||||
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OpenCV 不可用,运镜检测降级为 static/low: %s", exc)
|
||||
return
|
||||
try:
|
||||
cap = cv2.VideoCapture(str(video_path))
|
||||
for shot in shots:
|
||||
mid_t = (shot.start_sec + shot.end_sec) / 2.0
|
||||
dt = max(0.2, min(0.5, (shot.end_sec - shot.start_sec) / 4.0))
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, max(0.0, (mid_t - dt)) * 1000)
|
||||
ok1, f1 = cap.read()
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, min(mid_t + dt, shot.end_sec - 0.05) * 1000)
|
||||
ok2, f2 = cap.read()
|
||||
if not (ok1 and ok2):
|
||||
continue
|
||||
g1 = cv2.cvtColor(f1, cv2.COLOR_BGR2GRAY)
|
||||
g2 = cv2.cvtColor(f2, cv2.COLOR_BGR2GRAY)
|
||||
h, w = g1.shape
|
||||
# 降采样加速
|
||||
scale = 360.0 / h if h > 360 else 1.0
|
||||
if scale < 1.0:
|
||||
g1 = cv2.resize(g1, (int(w * scale), int(h * scale)))
|
||||
g2 = cv2.resize(g2, (int(w * scale), int(h * scale)))
|
||||
flow = cv2.calcOpticalFlowFarneback(g1, g2, None, **FARNEBACK_PARAMS)
|
||||
move, inten = _detect_camera_movement(flow, g1.shape[1], g1.shape[0])
|
||||
shot.movement = move
|
||||
shot.intensity = inten
|
||||
cap.release()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("运镜检测异常,已降级: %s", exc)
|
||||
|
||||
|
||||
# ── ④ librosa BPM ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _detect_bpm(video_path: Path) -> int:
|
||||
"""提取音轨并估算 BPM;失败返回 0。"""
|
||||
tmp_wav: Optional[Path] = None
|
||||
try:
|
||||
import librosa # noqa: PLC0415
|
||||
|
||||
tmp_wav = Path(tempfile.mkstemp(suffix=".wav")[1])
|
||||
# ffmpeg 抽 22050Hz 单声道 wav
|
||||
subprocess.run(
|
||||
[_ffmpeg_bin(), "-y", "-i", str(video_path), "-vn", "-ac", "1", "-ar", "22050", "-f", "wav", str(tmp_wav)],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=20,
|
||||
check=False,
|
||||
) # nosec B603
|
||||
if not tmp_wav.exists() or tmp_wav.stat().st_size < 1024:
|
||||
return 0
|
||||
y, sr = librosa.load(str(tmp_wav), sr=22050, mono=True)
|
||||
if len(y) < sr * 2:
|
||||
return 0
|
||||
tempo, _ = librosa.beat.beat_track(y=y, sr=sr)
|
||||
try:
|
||||
bpm = int(round(float(tempo)))
|
||||
except Exception: # noqa: BLE001
|
||||
bpm = int(round(float(tempo[0]))) if len(tempo) else 0
|
||||
return max(40, min(bpm, 220))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("librosa BPM 分析失败: %s", exc)
|
||||
return 0
|
||||
finally:
|
||||
if tmp_wav and tmp_wav.exists():
|
||||
try:
|
||||
tmp_wav.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _pace_from_bpm(bpm: int) -> str:
|
||||
if bpm >= 110:
|
||||
return "fast_cut"
|
||||
if bpm >= 80:
|
||||
return "medium"
|
||||
if bpm > 0:
|
||||
return "slow_cinematic"
|
||||
return "medium"
|
||||
|
||||
|
||||
# ── ⑤ VLM 帧分析 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sample_frames(frame_paths: list[Path], shots: list[ShotBoundary], k: int = VLM_SAMPLE_FRAMES) -> list[Path]:
|
||||
"""从全量帧中均匀选 k 张代表性帧(优先场景帧)。"""
|
||||
if not frame_paths:
|
||||
return []
|
||||
scene_frames = sorted(p for p in frame_paths if p.name.startswith("scene_"))
|
||||
fixed_frames = sorted(p for p in frame_paths if p.name.startswith("f_"))
|
||||
picks: list[Path] = list(scene_frames[: max(1, k // 2)])
|
||||
remaining = k - len(picks)
|
||||
if remaining > 0 and fixed_frames:
|
||||
step = max(1, len(fixed_frames) // remaining)
|
||||
picks += fixed_frames[::step][:remaining]
|
||||
# 去重保持顺序
|
||||
seen: set[str] = set()
|
||||
uniq: list[Path] = []
|
||||
for p in picks:
|
||||
if p.name not in seen and p.exists():
|
||||
seen.add(p.name)
|
||||
uniq.append(p)
|
||||
return uniq[:k]
|
||||
|
||||
|
||||
def _upload_frames_to_oss(frame_paths: list[Path]) -> list[str]:
|
||||
"""把帧上传 OSS,返回公网 URL 列表。失败时降级为 data URI。"""
|
||||
urls: list[str] = []
|
||||
try:
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
for p in frame_paths:
|
||||
try:
|
||||
key = f"viral-video/analysis/{uuid.uuid4().hex}/{p.name}"
|
||||
url = upload_to_oss(p, key)
|
||||
if url:
|
||||
urls.append(url)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("单帧 OSS 上传失败 %s: %s", p.name, exc)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("OSS 上传模块不可用,降级为 base64 data URI: %s", exc)
|
||||
if len(urls) < len(frame_paths):
|
||||
# 降级:base64 data URI(小图,单张 ≤100KB 才走此路)
|
||||
import base64
|
||||
|
||||
for p in frame_paths[len(urls) :]:
|
||||
try:
|
||||
if p.stat().st_size > 120_000:
|
||||
continue
|
||||
b64 = base64.b64encode(p.read_bytes()).decode("ascii")
|
||||
urls.append(f"data:image/jpeg;base64,{b64}")
|
||||
except Exception: # noqa: BLE001 # nosec B112
|
||||
continue
|
||||
return urls
|
||||
|
||||
|
||||
def _vlm_analyze_frames(image_urls: list[str]) -> dict[str, Any]:
|
||||
"""调豆包 VLM 分析色调/构图/光线/转场观感。"""
|
||||
if not image_urls:
|
||||
return {}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演和调色师。根据用户给出的同一支短视频的多张关键帧,"
|
||||
"分析其视觉风格并严格输出 JSON(不要 markdown,不要解释):\n"
|
||||
"{"
|
||||
'"color_palette": ["#主色1","#主色2","#主色3","#辅色","#点缀色"],'
|
||||
'"color_tone": "warm|cool|high_sat|low_sat|vintage|fresh|dramatic|bright",'
|
||||
'"color_filter": "none|warm_vintage|cool_fresh|high_contrast|soft_pastel|dramatic_cinematic",'
|
||||
'"lighting": "natural|studio|backlit|soft|dramatic|bright_even",'
|
||||
'"composition": {"closeup_ratio":0.0,"medium_ratio":0.0,"wide_ratio":0.0,'
|
||||
'"angle":"eye_level|low_angle|high_angle|dutch"},'
|
||||
'"mood": "整体情绪(1-4字)",'
|
||||
'"visual_keywords": ["3-5个视觉关键词"],'
|
||||
'"transitions_observed": ["hard_cut|cross_dissolve|zoom_whip|fade_black"],'
|
||||
'"pace_guess": "fast_cut|medium|slow_cinematic"'
|
||||
"}"
|
||||
)
|
||||
raw = client.vision_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": "请分析这支参考视频的风格。"},
|
||||
],
|
||||
images=image_urls,
|
||||
temperature=0.2,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return {}
|
||||
raw = _strip_code_fence(raw)
|
||||
# 容忍模型可能前后加文本
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
return json.loads(raw[i : j + 1])
|
||||
return {}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("VLM 帧分析失败: %s", exc)
|
||||
return {}
|
||||
|
||||
|
||||
# ── ⑥ LLM 整合 style_guide ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _llm_synthesize(
|
||||
shots: list[ShotBoundary],
|
||||
bpm: int,
|
||||
vlm: dict[str, Any],
|
||||
style_strength: str,
|
||||
) -> dict[str, Any]:
|
||||
"""把结构化信号整合成 style_guide;LLM 不可用时走规则合成。"""
|
||||
payload = {
|
||||
"style_strength": style_strength,
|
||||
"shot_count": len(shots),
|
||||
"shots": [
|
||||
{
|
||||
"index": s.index,
|
||||
"start_sec": round(s.start_sec, 2),
|
||||
"end_sec": round(s.end_sec, 2),
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"transition": s.transition,
|
||||
}
|
||||
for s in shots
|
||||
],
|
||||
"bpm": bpm,
|
||||
"pace_guess": _pace_from_bpm(bpm),
|
||||
"vlm": vlm,
|
||||
}
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise RuntimeError("豆包客户端未配置")
|
||||
sys_prompt = (
|
||||
"你是资深短视频导演。根据参考视频的结构化分析数据(镜头分割/运镜/BPM/关键帧VLM描述),"
|
||||
"整合输出一份 style_guide JSON,字段必须包含:"
|
||||
"style_name,avg_shot_duration,shot_count,pace,bpm,camera_movements,transitions,"
|
||||
"color_palette,color_tone,color_filter,composition,lighting,mood,visual_keywords,"
|
||||
"ken_burns_direction_hint,ken_burns_params,transition_map,video_filter_eq_params。"
|
||||
"严格输出一个合法 JSON 对象,不要 markdown/解释。"
|
||||
)
|
||||
user_text = "分析数据:\n" + json.dumps(payload, ensure_ascii=False)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": sys_prompt}, {"role": "user", "content": user_text}],
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
if raw:
|
||||
raw = _strip_code_fence(raw)
|
||||
i, j = raw.find("{"), raw.rfind("}")
|
||||
if i >= 0 and j > i:
|
||||
result = json.loads(raw[i : j + 1])
|
||||
if isinstance(result, dict) and result.get("style_name"):
|
||||
return result
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("LLM 合成 style_guide 失败,走规则降级: %s", exc)
|
||||
return _rule_based_style_guide(shots, bpm, vlm)
|
||||
|
||||
|
||||
def _rule_based_style_guide(shots: list[ShotBoundary], bpm: int, vlm: dict[str, Any]) -> dict[str, Any]:
|
||||
"""LLM 不可用时,用规则拼出可用 style_guide。"""
|
||||
durations = [s.end_sec - s.start_sec for s in shots] or [3.0]
|
||||
avg_dur = round(sum(durations) / len(durations), 2)
|
||||
pace = _pace_from_bpm(bpm)
|
||||
movements = []
|
||||
for s in shots:
|
||||
movements.append(
|
||||
{
|
||||
"shot_index": s.index + 1,
|
||||
"movement": s.movement,
|
||||
"intensity": s.intensity,
|
||||
"duration": round(s.end_sec - s.start_sec, 2),
|
||||
"subject_hint": _default_subject_hint(s.movement),
|
||||
}
|
||||
)
|
||||
transitions = []
|
||||
for i in range(len(shots) - 1):
|
||||
transitions.append({"between_shot": [i + 1, i + 2], "type": shots[i].transition})
|
||||
color_palette = vlm.get("color_palette") or ["#E0E0E0", "#333333", "#F5F5F5", "#888888", "#FF6B35"]
|
||||
color_tone = vlm.get("color_tone") or "bright"
|
||||
color_filter = vlm.get("color_filter") or "none"
|
||||
lighting = vlm.get("lighting") or "bright_even"
|
||||
composition = vlm.get("composition") or {
|
||||
"closeup_ratio": 0.4,
|
||||
"medium_ratio": 0.4,
|
||||
"wide_ratio": 0.2,
|
||||
"angle": "eye_level",
|
||||
}
|
||||
mood = vlm.get("mood") or "明快"
|
||||
vk = vlm.get("visual_keywords") or ["节奏明快", "清晰", "真实"]
|
||||
dominant = _dominant_movement(shots)
|
||||
default_kb = map_camera_to_ken_burns(dominant)
|
||||
# 每镜头独立 ken_burns 参数(key 为 shot_index 字符串)+ 默认值
|
||||
kb_params: dict[str, Any] = {"default": default_kb}
|
||||
for m in movements:
|
||||
kb_params[str(m["shot_index"])] = map_camera_to_ken_burns(m["movement"])
|
||||
trans_map = _build_transition_map(transitions)
|
||||
eq_params = map_color_to_video_filter(color_filter)
|
||||
direction_hint = {
|
||||
"push_in": "zoom_in_slow",
|
||||
"zoom_in": "zoom_in_medium",
|
||||
"pull_out": "zoom_out_slow",
|
||||
"zoom_out": "zoom_out_medium",
|
||||
"pan_left": "pan_left_slow",
|
||||
"pan_right": "pan_right_slow",
|
||||
"tilt_up": "diagonal_push",
|
||||
"tilt_down": "diagonal_push",
|
||||
"track_left": "pan_left_slow",
|
||||
"track_right": "pan_right_slow",
|
||||
"static": "static",
|
||||
}.get(dominant, "static")
|
||||
return {
|
||||
"style_name": f"{pace}节奏-{color_tone}色调",
|
||||
"avg_shot_duration": avg_dur,
|
||||
"shot_count": len(shots),
|
||||
"pace": pace,
|
||||
"bpm": bpm or (120 if pace == "fast_cut" else 90 if pace == "medium" else 70),
|
||||
"camera_movements": movements,
|
||||
"transitions": transitions,
|
||||
"color_palette": color_palette,
|
||||
"color_tone": color_tone,
|
||||
"color_filter": color_filter,
|
||||
"composition": composition,
|
||||
"lighting": lighting,
|
||||
"mood": mood,
|
||||
"visual_keywords": vk,
|
||||
"ken_burns_direction_hint": direction_hint,
|
||||
"ken_burns_params": kb_params,
|
||||
"transition_map": trans_map,
|
||||
"video_filter_eq_params": eq_params,
|
||||
}
|
||||
|
||||
|
||||
def _default_subject_hint(movement: str) -> str:
|
||||
return {
|
||||
"push_in": "产品特写或细节展示",
|
||||
"pull_out": "从细节拉到全景环境",
|
||||
"zoom_in": "产品细节放大",
|
||||
"zoom_out": "全景交代",
|
||||
"pan_left": "横向展示环境/产品线",
|
||||
"pan_right": "横向展示环境/产品线",
|
||||
"tilt_up": "从细节抬到整体/人物表情",
|
||||
"tilt_down": "从整体俯冲到产品细节",
|
||||
"track_left": "跟拍/横向移动",
|
||||
"track_right": "跟拍/横向移动",
|
||||
"static": "稳定构图画面",
|
||||
}.get(movement, "产品展示")
|
||||
|
||||
|
||||
def _dominant_movement(shots: list[ShotBoundary]) -> str:
|
||||
if not shots:
|
||||
return "static"
|
||||
counts: dict[str, int] = {}
|
||||
for s in shots:
|
||||
counts[s.movement] = counts.get(s.movement, 0) + 1
|
||||
return max(counts, key=counts.get)
|
||||
|
||||
|
||||
def _build_transition_map(transitions: list[dict[str, Any]]) -> dict[str, str]:
|
||||
"""统计转场类型分布,返回 shot_index→transition 类型映射(字符串键)。"""
|
||||
m: dict[str, str] = {}
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or [0, 0]
|
||||
if len(pair) >= 2:
|
||||
m[f"{pair[0]}-{pair[1]}"] = t.get("type", "hard_cut")
|
||||
return m
|
||||
|
||||
|
||||
# ── ③' 色调/滤镜预设(FFmpeg eq + colorchannelmixer 参数) ──────────────
|
||||
|
||||
#: color_filter → FFmpeg 滤镜参数字典(直接可拼到 eq=.../colorchannelmixer=...)
|
||||
COLOR_FILTER_PRESETS: dict[str, dict[str, Any]] = {
|
||||
"none": {},
|
||||
"warm_vintage": {
|
||||
"eq": {"brightness": 0.02, "contrast": 1.05, "saturation": 0.9, "gamma": 1.05},
|
||||
"colorchannelmixer": {"rr": 1.1, "gg": 0.98, "bb": 0.82, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"cool_fresh": {
|
||||
"eq": {"brightness": 0.03, "contrast": 1.08, "saturation": 1.05},
|
||||
"colorchannelmixer": {"rr": 0.9, "gg": 1.0, "bb": 1.12, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"high_contrast": {
|
||||
"eq": {"brightness": 0.0, "contrast": 1.3, "saturation": 1.2},
|
||||
"colorchannelmixer": {},
|
||||
},
|
||||
"soft_pastel": {
|
||||
"eq": {"brightness": 0.05, "contrast": 0.92, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 1.03, "bb": 1.05, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
"dramatic_cinematic": {
|
||||
"eq": {"brightness": -0.03, "contrast": 1.2, "saturation": 0.85},
|
||||
"colorchannelmixer": {"rr": 1.05, "gg": 0.98, "bb": 0.9, "ra": 0, "ga": 0, "ba": 0, "aa": 1},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_color_to_video_filter(color_filter: str) -> dict[str, Any]:
|
||||
"""color_filter 枚举 → FFmpeg eq/colorchannelmixer 参数字典(渲染端直接使用)。"""
|
||||
preset = COLOR_FILTER_PRESETS.get(color_filter) or COLOR_FILTER_PRESETS["none"]
|
||||
# 返回深拷贝防污染
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── ③'' 运镜 → ken_burns 参数映射 ───────────────────────────────────────
|
||||
|
||||
#: 运镜类型 → URS 可直接消费的 ken_burns 参数字典
|
||||
CAMERA_TO_KEN_BURNS: dict[str, dict[str, Any]] = {
|
||||
"static": {
|
||||
"type": "static",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"push_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.12,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_in": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.0,
|
||||
"zoom_end": 1.18,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pull_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.12,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"zoom_out": {
|
||||
"type": "zoom",
|
||||
"zoom_start": 1.18,
|
||||
"zoom_end": 1.0,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"pan_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.08,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_up": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.08,
|
||||
"zoom_end": 1.14,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": -0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"tilt_down": {
|
||||
"type": "pan+zoom",
|
||||
"zoom_start": 1.14,
|
||||
"zoom_end": 1.08,
|
||||
"pan_x": 0.0,
|
||||
"pan_y": 0.05,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_left": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": -0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
"track_right": {
|
||||
"type": "pan",
|
||||
"zoom_start": 1.05,
|
||||
"zoom_end": 1.05,
|
||||
"pan_x": 0.10,
|
||||
"pan_y": 0.0,
|
||||
"duration_factor": 1.0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def map_camera_to_ken_burns(movement: str) -> dict[str, Any]:
|
||||
"""运镜类型 → URS ken_burns 参数字典。未知类型回退 static。"""
|
||||
preset = CAMERA_TO_KEN_BURNS.get(movement) or CAMERA_TO_KEN_BURNS["static"]
|
||||
return json.loads(json.dumps(preset))
|
||||
|
||||
|
||||
# ── 转场 → xfade transition 名称 ────────────────────────────────────────
|
||||
|
||||
TRANSITION_TO_XFADE: dict[str, str] = {
|
||||
"hard_cut": "cut",
|
||||
"cross_dissolve": "dissolve",
|
||||
"fade_black": "fadeblack",
|
||||
"fade": "fade",
|
||||
"zoom_whip": "zoom",
|
||||
"slide_left": "slideright", # 画面左移 = 新画面从右滑入
|
||||
"slide_right": "slideleft",
|
||||
"wipe_left": "wipeleft",
|
||||
"wipe_right": "wiperight",
|
||||
}
|
||||
|
||||
|
||||
def map_transition_to_xfade(transition_type: str) -> str:
|
||||
"""转场枚举 → TransitionEngine 支持的 xfade 名称;未知回退 cut。"""
|
||||
return TRANSITION_TO_XFADE.get(transition_type, "cut")
|
||||
|
||||
|
||||
# ── BPM → BGM 推荐 BPM ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def map_bgm_bpm(bpm: int) -> int:
|
||||
"""BGM 选曲 BPM:参考视频 BPM ±5。bpm=0 返回 90(默认 medium)。"""
|
||||
if bpm <= 0:
|
||||
return 90
|
||||
return max(60, min(bpm, 180))
|
||||
|
||||
|
||||
# ── 单 clip 渲染参数聚合(给 URS build_render_plan 使用) ───────────────
|
||||
|
||||
|
||||
def build_render_params_for_clip(
|
||||
clip_index: int,
|
||||
style_guide: dict[str, Any],
|
||||
*,
|
||||
duration_sec: Optional[float] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""根据 style_guide 为第 clip_index 个 clip 生成可直接喂给 URS 的渲染参数。"""
|
||||
shot_idx = clip_index + 1
|
||||
movements = style_guide.get("camera_movements") or []
|
||||
movement = "static"
|
||||
intensity = "low"
|
||||
for m in movements:
|
||||
if m.get("shot_index") == shot_idx:
|
||||
movement = m.get("movement", "static")
|
||||
intensity = m.get("intensity", "low")
|
||||
break
|
||||
ken = map_camera_to_ken_burns(movement)
|
||||
if intensity == "high":
|
||||
ken["zoom_end"] = round(ken.get("zoom_end", 1.0) * 1.08, 3)
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 1.3, 3)
|
||||
elif intensity == "low":
|
||||
for k in ("pan_x", "pan_y"):
|
||||
ken[k] = round(ken.get(k, 0.0) * 0.6, 3)
|
||||
transitions = style_guide.get("transitions") or []
|
||||
trans_type = "hard_cut"
|
||||
for t in transitions:
|
||||
pair = t.get("between_shot") or []
|
||||
if len(pair) >= 2 and pair[0] == shot_idx:
|
||||
trans_type = t.get("type", "hard_cut")
|
||||
break
|
||||
xfade = map_transition_to_xfade(trans_type)
|
||||
eq = map_color_to_video_filter(style_guide.get("color_filter", "none"))
|
||||
return {
|
||||
"ken_burns": ken,
|
||||
"transition": {"type": xfade, "duration": 0.3 if xfade != "cut" else 0.0},
|
||||
"video_filter": eq,
|
||||
"bgm_bpm_hint": map_bgm_bpm(int(style_guide.get("bpm") or 0)),
|
||||
"duration_sec": duration_sec,
|
||||
}
|
||||
|
||||
|
||||
# ── 素材本地化(URL/OSS key → 本地临时文件) ────────────────────────────
|
||||
|
||||
|
||||
def _ensure_local_video(reference: str, work_dir: Path) -> Optional[Path]:
|
||||
"""把 reference(URL/OSS key/本地路径)落到 work_dir 下的本地文件。"""
|
||||
p = Path(reference)
|
||||
if p.exists() and p.is_file():
|
||||
return p
|
||||
try:
|
||||
from video_processing.oss_helpers import download_asset
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
ok = download_asset(reference, target)
|
||||
if ok and target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("download_asset 失败,尝试 http 直连: %s", exc)
|
||||
if reference.startswith(("http://", "https://")):
|
||||
try:
|
||||
import httpx # noqa: PLC0415 - 项目依赖,延迟导入
|
||||
|
||||
target = work_dir / f"ref_{uuid.uuid4().hex}.mp4"
|
||||
with httpx.Client(timeout=20.0, follow_redirects=True) as client:
|
||||
with client.stream("GET", reference) as resp:
|
||||
resp.raise_for_status()
|
||||
with open(target, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=64 * 1024):
|
||||
f.write(chunk)
|
||||
if target.exists() and target.stat().st_size > 0:
|
||||
return target
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("HTTP 下载参考视频失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── 入口 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def analyze_video_style(
|
||||
reference_video_path: str | Path,
|
||||
style_strength: str = "medium",
|
||||
*,
|
||||
timeout_sec: int = DEFAULT_ANALYSIS_TIMEOUT,
|
||||
) -> dict[str, Any]:
|
||||
"""分析参考视频风格,返回 style_guide dict。
|
||||
|
||||
Args:
|
||||
reference_video_path: 本地路径、HTTP(S) URL 或 OSS storage key。
|
||||
style_strength: light | medium | strict。
|
||||
timeout_sec: 单步超时(秒),默认 60。
|
||||
|
||||
Returns:
|
||||
style_guide dict,详见 STYLE_GUIDE_SCHEMA。任何子步骤失败都会降级,不抛异常。
|
||||
"""
|
||||
style_strength = style_strength if style_strength in ("light", "medium", "strict") else "medium"
|
||||
frames_dir: Optional[Path] = None
|
||||
local_path: Optional[Path] = None
|
||||
try:
|
||||
frames_dir = Path(tempfile.mkdtemp(prefix="vstyle_"))
|
||||
work_dir = frames_dir # 同一临时根
|
||||
local_path = _ensure_local_video(str(reference_video_path), work_dir)
|
||||
if local_path is None:
|
||||
logger.error("[video_analyzer] 无法获取参考视频: %s", reference_video_path)
|
||||
return _rule_based_style_guide([], 0, {})
|
||||
# 资源约束:大小 / 时长
|
||||
try:
|
||||
size_mb = local_path.stat().st_size / (1024 * 1024)
|
||||
if size_mb > MAX_REFERENCE_SIZE_MB:
|
||||
logger.warning(
|
||||
"[video_analyzer] 参考视频 %.1fMB 超上限,按前 %ds 分析", size_mb, MAX_REFERENCE_DURATION_SEC
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
duration = _probe_duration(local_path)
|
||||
if duration > MAX_REFERENCE_DURATION_SEC:
|
||||
duration = MAX_REFERENCE_DURATION_SEC
|
||||
# ① 抽帧
|
||||
try:
|
||||
frame_paths = _extract_keyframes(local_path, frames_dir / "frames")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("FFmpeg 抽帧失败: %s,降级为 VLM 均匀采样", exc)
|
||||
frame_paths = []
|
||||
# ② 镜头分割
|
||||
shots = _detect_shots(local_path, frames_dir)
|
||||
# ③ 运镜检测(有帧才跑)
|
||||
if frame_paths or shots:
|
||||
_analyze_movements(local_path, shots)
|
||||
# ④ BPM
|
||||
bpm = _detect_bpm(local_path)
|
||||
# ⑤ 选帧→OSS→VLM
|
||||
sampled = _sample_frames(frame_paths, shots)
|
||||
image_urls = _upload_frames_to_oss(sampled) if sampled else []
|
||||
vlm = _vlm_analyze_frames(image_urls) if image_urls else {}
|
||||
# ⑥ 合成
|
||||
style_guide = _llm_synthesize(shots, bpm, vlm, style_strength)
|
||||
# 兜底字段校验
|
||||
style_guide.setdefault("style_strength", style_strength)
|
||||
style_guide.setdefault("pace", _pace_from_bpm(bpm))
|
||||
style_guide.setdefault("bpm", bpm)
|
||||
style_guide.setdefault("shot_count", len(shots))
|
||||
if shots and "avg_shot_duration" not in style_guide:
|
||||
durs = [s.end_sec - s.start_sec for s in shots]
|
||||
style_guide["avg_shot_duration"] = round(sum(durs) / len(durs), 2)
|
||||
return style_guide
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("[video_analyzer] 整体分析异常,返回最小占位 style_guide: %s", exc)
|
||||
return _rule_based_style_guide([], 0, {"mood": "未知"})
|
||||
finally:
|
||||
# 临时帧清理
|
||||
if frames_dir and frames_dir.exists():
|
||||
shutil.rmtree(frames_dir, ignore_errors=True)
|
||||
@@ -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%)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析:火山OCR专用API + doubao-lite强约束JSON并行,单次pro VLM兜底。"""
|
||||
|
||||
from .fast_path import analyze_image_v2, analyze_images_v2 # noqa: F401
|
||||
@@ -0,0 +1,307 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""把 fast_json VLM 输出 + OCR 文本组装为与旧 _normalize() 完全一致的 dict。
|
||||
|
||||
目标:下游(信任链t2i/intent_parsing/script_generation)零改动。
|
||||
必出字段:name, brand, category, appearance, packaging, text_on_package,
|
||||
key_features, scene, mood, portrait_prompt, summary, _source
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# ---------- portrait_prompt 模板 ----------
|
||||
# 目标:60-100 字的人物穿搭描述,用于 Seedream 纯文生图。要求具体、风格化、视觉细节丰富。
|
||||
# 旧 VLM 输出格式参考:"一位25岁左右的亚洲女性,身穿白色V领短袖T恤,黑色高腰阔腿裤,
|
||||
# 搭配银色项链,长发披肩,表情自信,街拍风格,阳光明媚的城市街头"
|
||||
|
||||
|
||||
def _join_parts(*parts: str | None) -> str:
|
||||
return "".join(p for p in parts if p)
|
||||
|
||||
|
||||
_AGE_PREFIX = {
|
||||
"青年": "年轻",
|
||||
"中年": "中年",
|
||||
"老年": "老年",
|
||||
}
|
||||
# gender 后缀
|
||||
_GENDER_WORD = {"男": "男性", "女": "女性"}
|
||||
|
||||
|
||||
def _person_subject(fj: dict[str, Any]) -> str:
|
||||
"""人物主语:年轻女性 / 中年男性 / 少女 / 小男孩 / 人物 等。"""
|
||||
gender = fj.get("gender") or ""
|
||||
age = fj.get("age_range") or ""
|
||||
gw = _GENDER_WORD.get(gender, "")
|
||||
if age == "儿童":
|
||||
if gender == "女":
|
||||
return "小女孩"
|
||||
if gender == "男":
|
||||
return "小男孩"
|
||||
return "儿童"
|
||||
if age == "青少年":
|
||||
if gender == "女":
|
||||
return "少女"
|
||||
if gender == "男":
|
||||
return "少年"
|
||||
return "青少年"
|
||||
prefix = _AGE_PREFIX.get(age, "")
|
||||
if gw:
|
||||
return f"{prefix}{gw}" if prefix else gw
|
||||
return f"{prefix}人物" if prefix else "人物"
|
||||
|
||||
|
||||
def _build_wear_sentence(fj: dict[str, Any]) -> str:
|
||||
"""穿搭段:上装+下装/连衣裙,带颜色+材质+图案。"""
|
||||
upper = fj.get("upper_wear") or ""
|
||||
upper_color = fj.get("upper_color") or ""
|
||||
lower = fj.get("lower_wear") or ""
|
||||
lower_color = fj.get("lower_color") or ""
|
||||
dress_color = fj.get("dress_color") or ""
|
||||
material = fj.get("material") or ""
|
||||
pattern = fj.get("pattern") or ""
|
||||
|
||||
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
|
||||
if is_dress:
|
||||
c = dress_color or upper_color
|
||||
wear = f"{c}{upper}" if c else upper
|
||||
if material and material not in wear:
|
||||
wear = f"{material}{wear}"
|
||||
if pattern and pattern not in wear and pattern != "纯色":
|
||||
wear += f",{pattern}图案"
|
||||
return f"身穿{wear}"
|
||||
|
||||
parts: list[str] = []
|
||||
if upper:
|
||||
up = f"{upper_color}{upper}" if upper_color else upper
|
||||
if material and material not in up:
|
||||
up = f"{material}{up}"
|
||||
if pattern and pattern != "纯色" and pattern not in up:
|
||||
up += f"({pattern})"
|
||||
parts.append(f"上身{up}" if up else "")
|
||||
if lower:
|
||||
lo = f"{lower_color}{lower}" if lower_color else lower
|
||||
parts.append(f"下身{lo}" if lo else "")
|
||||
return ",".join(p for p in parts if p)
|
||||
|
||||
|
||||
def _build_portrait_prompt(fj: dict[str, Any]) -> str:
|
||||
"""组装最终 portrait_prompt(目标 60-100 字,用于 Seedream 纯文生图)。"""
|
||||
if not fj.get("has_person"):
|
||||
# 非人像:用商品+场景+mood 拼一段
|
||||
name = fj.get("product_name") or "商品"
|
||||
brand = fj.get("brand") or ""
|
||||
colors = fj.get("colors") or []
|
||||
style = fj.get("style") or ""
|
||||
scene = fj.get("scene") or ""
|
||||
mood = fj.get("mood") or ""
|
||||
pieces = []
|
||||
if brand:
|
||||
pieces.append(brand)
|
||||
pieces.append(name)
|
||||
if colors:
|
||||
pieces.append("、".join(colors[:3]) + "配色")
|
||||
if style:
|
||||
pieces.append(style + "风格")
|
||||
if mood:
|
||||
pieces.append(mood + "氛围")
|
||||
if scene and scene not in ("通用",):
|
||||
pieces.append(scene + "场景")
|
||||
pieces.append("产品特写")
|
||||
prompt = ",".join(p for p in pieces if p)
|
||||
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
|
||||
|
||||
subject = _person_subject(fj)
|
||||
wear = _build_wear_sentence(fj)
|
||||
|
||||
accessories = fj.get("accessories") or []
|
||||
if isinstance(accessories, str):
|
||||
accessories = [accessories]
|
||||
acc_str = ""
|
||||
if accessories:
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a)
|
||||
|
||||
hairstyle = fj.get("hairstyle") or ""
|
||||
expression = fj.get("expression") or ""
|
||||
pose = fj.get("pose") or ""
|
||||
style = fj.get("style") or ""
|
||||
scene = fj.get("scene") or ""
|
||||
mood = fj.get("mood") or ""
|
||||
|
||||
detail_parts: list[str] = []
|
||||
if hairstyle:
|
||||
detail_parts.append(hairstyle)
|
||||
if expression and expression not in ("自然", "平静"):
|
||||
detail_parts.append(f"神情{expression}")
|
||||
if pose and pose not in ("站立",):
|
||||
detail_parts.append(pose)
|
||||
|
||||
style_parts: list[str] = []
|
||||
if style:
|
||||
style_parts.append(style)
|
||||
if mood:
|
||||
style_parts.append(mood)
|
||||
if scene and scene not in ("通用",):
|
||||
style_parts.append(scene)
|
||||
|
||||
pieces = [f"一位{subject}"]
|
||||
if wear:
|
||||
pieces.append(wear)
|
||||
if acc_str:
|
||||
pieces.append(acc_str.lstrip(","))
|
||||
if detail_parts:
|
||||
pieces.append(",".join(detail_parts))
|
||||
if style_parts:
|
||||
# 风格词之间不用逗号,用空格紧凑
|
||||
pieces.append("".join(style_parts) + "风格")
|
||||
else:
|
||||
pieces.append("人像写真")
|
||||
|
||||
full = ",".join(p for p in pieces if p)
|
||||
# 过短补充镜头词
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
# 过长截断
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
|
||||
|
||||
# ---------- 商品字段 ----------
|
||||
|
||||
|
||||
def _infer_name(fj: dict[str, Any], ocr_texts: list[str]) -> str:
|
||||
pname = fj.get("product_name")
|
||||
if pname and pname != "未识别":
|
||||
return str(pname)
|
||||
# 人物图 → name 用穿搭主件
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or ""
|
||||
if "连衣裙" in up:
|
||||
return up
|
||||
return up or "人物穿搭"
|
||||
if ocr_texts:
|
||||
# 商品名可能是 OCR 最长的一行(品牌/产品名)
|
||||
return max(ocr_texts, key=len)
|
||||
return "未识别"
|
||||
|
||||
|
||||
def _infer_brand(fj: dict[str, Any], ocr_texts: list[str]) -> str:
|
||||
brand = fj.get("brand")
|
||||
if brand:
|
||||
return str(brand)
|
||||
# OCR 里短的、纯字母/汉字短串可能是 brand
|
||||
for t in ocr_texts:
|
||||
if 1 < len(t) <= 12:
|
||||
return t
|
||||
return "无法判断"
|
||||
|
||||
|
||||
def _infer_category(fj: dict[str, Any]) -> str:
|
||||
cat = fj.get("category")
|
||||
if cat:
|
||||
return str(cat)
|
||||
if fj.get("has_person"):
|
||||
return "服饰"
|
||||
return "非产品图"
|
||||
|
||||
|
||||
def _build_appearance(fj: dict[str, Any]) -> str:
|
||||
"""外观描述:颜色+款式+材质+图案 拼成一段。"""
|
||||
parts: list[str] = []
|
||||
for key, _label in [
|
||||
("upper_color", "主色"),
|
||||
("upper_wear", "款式"),
|
||||
("material", "材质"),
|
||||
("pattern", "图案"),
|
||||
]:
|
||||
v = fj.get(key)
|
||||
if v and v not in ("无法判断", "未知", "纯色"):
|
||||
parts.append(str(v))
|
||||
if not parts:
|
||||
if fj.get("has_person"):
|
||||
return "人像穿搭整体造型"
|
||||
return "无法判断"
|
||||
return "、".join(parts)
|
||||
|
||||
|
||||
def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
|
||||
feats: list[str] = []
|
||||
for key in (
|
||||
"upper_wear",
|
||||
"lower_wear",
|
||||
"upper_color",
|
||||
"lower_color",
|
||||
"dress_color",
|
||||
"material",
|
||||
"pattern",
|
||||
"style",
|
||||
"accessories",
|
||||
):
|
||||
v = fj.get(key)
|
||||
if not v:
|
||||
continue
|
||||
if isinstance(v, list):
|
||||
feats.extend(str(x) for x in v if x)
|
||||
elif isinstance(v, str) and v not in ("无法判断", "未知", "纯色"):
|
||||
feats.append(v)
|
||||
if ocr_texts:
|
||||
feats.append(f"画面文字: {'/'.join(ocr_texts[:3])}")
|
||||
# 去重
|
||||
out: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for f in feats:
|
||||
f = f.strip()
|
||||
if f and f not in seen and len(f) <= 30:
|
||||
seen.add(f)
|
||||
out.append(f)
|
||||
return out[:6] if out else ["无法判断"]
|
||||
|
||||
|
||||
def assemble_result(
|
||||
idx: int,
|
||||
fast_json: dict[str, Any] | None,
|
||||
ocr_texts: list[str],
|
||||
) -> dict[str, Any]:
|
||||
"""把 fast_json 结果 + OCR 文本组装成下游兼容的 product dict。"""
|
||||
fj = fast_json or {}
|
||||
ocr_texts = ocr_texts or []
|
||||
|
||||
portrait_prompt = _build_portrait_prompt(fj)
|
||||
name = _infer_name(fj, ocr_texts)
|
||||
brand = _infer_brand(fj, ocr_texts)
|
||||
category = _infer_category(fj)
|
||||
appearance = _build_appearance(fj)
|
||||
key_features = _build_key_features(fj, ocr_texts)
|
||||
scene = fj.get("scene") or "通用"
|
||||
mood = fj.get("mood") or ""
|
||||
packaging = "无法判断" # 包装细节专用API无,保留占位
|
||||
text_on_package = ocr_texts[:8]
|
||||
summary = _build_summary(fj, name, brand, category)
|
||||
|
||||
return {
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance,
|
||||
"packaging": packaging,
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": key_features,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary,
|
||||
"_source": "v2_fast_json",
|
||||
}
|
||||
|
||||
|
||||
def _build_summary(fj: dict, name: str, brand: str, category: str) -> str:
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or "穿搭"
|
||||
style = fj.get("style") or ""
|
||||
base = f"{style}{up}" if style and style not in up else up
|
||||
return base
|
||||
if brand != "无法判断" and name != brand:
|
||||
return f"{brand} {name}"
|
||||
return name
|
||||
@@ -0,0 +1,144 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析主路径:每图并行 OCR(火山专用API)+ lite JSON VLM,失败时单次 pro VLM 兜底。
|
||||
|
||||
设计原则(灵应10-05要求):
|
||||
- 主力路径简洁:单图2路并行,外层N图全并发
|
||||
- 兜底简单:单次 pro VLM 调用,无竞速/重试/复杂超时
|
||||
- 输出 dict 格式与旧版完全一致,下游零改动
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Any
|
||||
|
||||
from . import assembler, ocr_volc, vlm_fallback, vlm_fast_json
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 可通过环境变量调参(有默认值,无需配置即可跑)
|
||||
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "6"))
|
||||
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "5"))
|
||||
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "5"))
|
||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
|
||||
|
||||
_FALLBACK_RESULT = {
|
||||
"name": "未识别",
|
||||
"brand": "无法判断",
|
||||
"category": "非产品图",
|
||||
"appearance": "无法判断",
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": ["无法判断"],
|
||||
"scene": "通用",
|
||||
"mood": "",
|
||||
"portrait_prompt": "无法判断",
|
||||
"summary": "未识别",
|
||||
}
|
||||
|
||||
|
||||
def _is_usable(r: dict[str, Any]) -> bool:
|
||||
"""结果可用判定:portrait_prompt 是核心,有效就算 usable。"""
|
||||
pp = (r.get("portrait_prompt") or "").strip()
|
||||
if pp and pp not in ("无人像", "无法判断", "未识别"):
|
||||
return True
|
||||
name = (r.get("name") or "").strip()
|
||||
if name and name not in ("未识别", "无法判断", "未知"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
||||
"""单张图片 V2 分析。"""
|
||||
t0 = time.time()
|
||||
|
||||
# 第1层:OCR + lite JSON VLM 并行
|
||||
fj_result: dict[str, Any] | None = None
|
||||
ocr_result: list[str] = []
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
||||
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||
try:
|
||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||
try:
|
||||
res = fut.result(timeout=1)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||
continue
|
||||
if fut is f_fj and isinstance(res, dict):
|
||||
fj_result = res
|
||||
elif fut is f_ocr and isinstance(res, list):
|
||||
ocr_result = res
|
||||
except TimeoutError:
|
||||
# fast 整体超时,取消还没跑完的子任务,继续走 pro 兜底
|
||||
for f in (f_fj, f_ocr):
|
||||
if not f.done():
|
||||
f.cancel()
|
||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||
|
||||
fast_elapsed = time.time() - t0
|
||||
|
||||
# 组装 fast 结果
|
||||
if fj_result:
|
||||
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
||||
if _is_usable(assembled):
|
||||
assembled["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
logger.info(
|
||||
"[vision.v2] 图片 #%d fast命中 elapsed=%.2fs pp=%s",
|
||||
idx,
|
||||
fast_elapsed,
|
||||
(assembled.get("portrait_prompt") or "")[:40],
|
||||
)
|
||||
return assembled
|
||||
|
||||
# 第2层:pro VLM 单次兜底
|
||||
pro_t0 = time.time()
|
||||
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, timeout=_PRO_TIMEOUT)
|
||||
if pro_result and _is_usable(pro_result):
|
||||
pro_result["_fallback_used"] = True
|
||||
pro_result["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
pro_result["_pro_elapsed"] = round(time.time() - pro_t0, 2)
|
||||
if ocr_result and not pro_result.get("text_on_package"):
|
||||
pro_result["text_on_package"] = ocr_result[:8]
|
||||
logger.info("[vision.v2] 图片 #%d pro兜底命中 total=%.2fs", idx, time.time() - t0)
|
||||
return pro_result
|
||||
|
||||
# 最终:返回最小可用结果
|
||||
logger.warning("[vision.v2] 图片 #%d 全路径失败 elapsed=%.2fs", idx, time.time() - t0)
|
||||
out = dict(_FALLBACK_RESULT)
|
||||
out["_source"] = "v2_all_failed"
|
||||
out["text_on_package"] = ocr_result[:8]
|
||||
out["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
return out
|
||||
|
||||
|
||||
def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
|
||||
"""批量图片 V2 分析,外层全并发。"""
|
||||
if not img_urls:
|
||||
return []
|
||||
workers = min(_IMG_WORKERS, len(img_urls), 16)
|
||||
results: list[dict[str, Any] | None] = [None] * len(img_urls)
|
||||
|
||||
logger.info("[vision.v2] 开始图片分析 n=%d workers=%d fast_timeout=%.0fs", len(img_urls), workers, _FAST_TIMEOUT)
|
||||
t0 = time.time()
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
future_to_idx = {pool.submit(analyze_image_v2, idx, url): idx for idx, url in enumerate(img_urls)}
|
||||
for fut in as_completed(future_to_idx):
|
||||
idx = future_to_idx[fut]
|
||||
try:
|
||||
results[idx] = fut.result()
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
|
||||
r = dict(_FALLBACK_RESULT)
|
||||
r["_source"] = "v2_future_exception"
|
||||
results[idx] = r
|
||||
|
||||
elapsed = time.time() - t0
|
||||
succ = sum(1 for r in results if r and _is_usable(r))
|
||||
fb = sum(1 for r in results if r and r.get("_fallback_used"))
|
||||
logger.info("[vision.v2] 完成 n=%d usable=%d pro_fallback=%d elapsed=%.2fs", len(img_urls), succ, fb, elapsed)
|
||||
return [r for r in results if r is not None]
|
||||
@@ -0,0 +1,109 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""火山引擎 AI MediaKit OCR(同步)调用封装。
|
||||
|
||||
接口:POST {mediakit_base_url}/tools-sync/ocr
|
||||
鉴权:Bearer {mediakit_api_key}
|
||||
请求体:{"image_url": "<公网可访问URL>"} (部分版本也支持 image_base64)
|
||||
响应:{"code":0,"data":{"texts":[{"text":"...","bbox":[x,y,w,h],...},...],...}}
|
||||
|
||||
目标:识别商品包装/Logo/水印上的文字,作为 fast_json VLM 的补充。
|
||||
返回值:识别到的文本字符串列表(失败返回 [])。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT = 8 # OCR 秒级返回,8s 绰绰有余
|
||||
|
||||
|
||||
def call_ocr(img_url: str, *, timeout: int = DEFAULT_TIMEOUT) -> list[str]:
|
||||
"""调用 MediaKit 同步 OCR,返回去重后的纯文本列表。
|
||||
|
||||
不做重试(外层降级逻辑负责)。失败/未配置返回空列表,不抛异常。
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
import httpx
|
||||
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
if not client.is_available:
|
||||
logger.info("[vision.v2] mediakit 未配置,跳过 OCR")
|
||||
return []
|
||||
|
||||
url = f"{client.base_url}/tools-sync/ocr"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {client.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {"image_url": img_url}
|
||||
# 部分文档版本用 image_base64,但公网 URL 场景下 image_url 最简
|
||||
resp = httpx.post(url, headers=headers, json=payload, timeout=timeout)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
"[vision.v2] OCR HTTP %d elapsed=%.1fs body=%s",
|
||||
resp.status_code,
|
||||
elapsed,
|
||||
resp.text[:200],
|
||||
)
|
||||
return []
|
||||
data = resp.json()
|
||||
# 兼容几种可能的响应结构
|
||||
code = data.get("code", data.get("status", 0))
|
||||
if code not in (0, "OK", "success", 200):
|
||||
logger.warning("[vision.v2] OCR 业务错误 code=%s elapsed=%.1fs resp=%s", code, elapsed, str(data)[:200])
|
||||
return []
|
||||
texts = _extract_texts(data)
|
||||
# 去重 + 过滤空
|
||||
seen: set[str] = set()
|
||||
out: list[str] = []
|
||||
for t in texts:
|
||||
t = (t or "").strip()
|
||||
if t and t not in seen and len(t) <= 100: # 过滤过长的误识别
|
||||
seen.add(t)
|
||||
out.append(t)
|
||||
logger.info("[vision.v2] OCR 完成 elapsed=%.1fs n=%d texts=%s", elapsed, len(out), out[:5])
|
||||
return out
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] OCR 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def _extract_texts(data: dict) -> list[str]:
|
||||
"""从 OCR 响应中抽取文本,兼容多种结构。"""
|
||||
out: list[str] = []
|
||||
# 常见结构1: data.texts = [{"text": "..."}, ...]
|
||||
d = data.get("data") or data
|
||||
if isinstance(d, dict):
|
||||
for key in ("texts", "lines", "words", "items", "result"):
|
||||
items = d.get(key)
|
||||
if isinstance(items, list):
|
||||
for it in items:
|
||||
if isinstance(it, dict):
|
||||
txt = it.get("text") or it.get("content") or it.get("word")
|
||||
if txt:
|
||||
out.append(str(txt))
|
||||
elif isinstance(it, str):
|
||||
out.append(it)
|
||||
break
|
||||
# 结构2: data.text = "..."
|
||||
if not out:
|
||||
t = d.get("text")
|
||||
if isinstance(t, str):
|
||||
out.append(t)
|
||||
# 结构3: data.ocr_text / data.content
|
||||
if not out:
|
||||
for key in ("ocr_text", "content", "raw_text"):
|
||||
v = d.get(key)
|
||||
if isinstance(v, str) and v.strip():
|
||||
out.append(v)
|
||||
break
|
||||
return out
|
||||
@@ -0,0 +1,226 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""VLM 兜底:专用API路径失败时的最后一道防线,单次调用 doubao-seed-2.1-pro。
|
||||
|
||||
设计原则:简单、直接、无竞速、无复杂超时逻辑。只在 fast_json 结果不可用时调用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_PRO_MODEL = "doubao-seed-2-1-pro-260915"
|
||||
DEFAULT_TIMEOUT = 45
|
||||
DEFAULT_MAX_TOKENS = 800
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
|
||||
def _xml_text(tag: str, xml: str) -> str:
|
||||
m = re.search(rf"<{tag}[^>]*>(.*?)</{tag}>", xml, re.S)
|
||||
return (m.group(1) if m else "").strip()
|
||||
|
||||
|
||||
def _xml_attr(tag: str, attr: str, xml: str) -> str:
|
||||
m = re.search(rf"<{tag}[^>]*\b{attr}\s*=\s*[\"']([^\"']*)[\"']", xml)
|
||||
return (m.group(1) if m else "").strip()
|
||||
|
||||
|
||||
def _xml_to_product(raw: str, idx: int) -> dict[str, Any]:
|
||||
"""解析 VLM 输出的 XML 格式(简化版)。"""
|
||||
scene = _xml_text("scene", raw) or "通用"
|
||||
mood = _xml_text("mood", raw) or ""
|
||||
|
||||
portrait_prompt = "无人像"
|
||||
p_has = _xml_attr("people", "has_person", raw)
|
||||
if p_has and p_has.lower() != "false":
|
||||
gender = _xml_attr("people", "gender", raw) or ""
|
||||
age = _xml_attr("people", "age_range", raw) or ""
|
||||
outfit = _xml_attr("people", "outfit", raw) or ""
|
||||
hair = _xml_attr("people", "hair", raw) or "自然发型"
|
||||
pose = _xml_attr("people", "pose", raw) or ""
|
||||
expr = _xml_attr("people", "expression", raw) or "自然"
|
||||
parts: list[str] = []
|
||||
if gender:
|
||||
parts.append(gender + ("性" if not gender.endswith("性") else ""))
|
||||
if age:
|
||||
parts.append(age)
|
||||
parts.append("人物")
|
||||
parts.append(hair)
|
||||
if outfit:
|
||||
parts.append(f"身着{outfit}")
|
||||
if pose:
|
||||
parts.append(f"姿态{pose}")
|
||||
parts.append(f"表情{expr}")
|
||||
portrait_prompt = ",".join(parts)
|
||||
|
||||
m = re.search(r"<product[^>]*>(.*?)</product>", raw, re.S)
|
||||
if m:
|
||||
pbody = m.group(1)
|
||||
name = _xml_attr("product", "name", raw) or _xml_text("name", pbody) or "未识别"
|
||||
brand = _xml_attr("product", "brand", raw) or _xml_text("brand", pbody) or "无法判断"
|
||||
category = _xml_attr("product", "category", raw) or _xml_text("category", pbody) or "无法判断"
|
||||
appearance = _xml_attr("product", "appearance", raw) or _xml_text("appearance", pbody) or "无法判断"
|
||||
packaging = _xml_attr("product", "packaging", raw) or _xml_text("packaging", pbody) or "无法判断"
|
||||
feat = _xml_attr("product", "features", raw) or _xml_text("features", pbody) or ""
|
||||
feat_list = [x.strip() for x in re.split(r"[,,;;]", feat) if x.strip()] if feat else ["无法判断"]
|
||||
top_text = _xml_attr("product", "text_on_package", raw) or _xml_text("text_on_package", pbody) or ""
|
||||
text_list = [x.strip() for x in re.split(r"[,,;;]", top_text) if x.strip()] if top_text else []
|
||||
summary = _xml_attr("product", "summary", raw) or _xml_text("summary", pbody) or f"{brand} {name}"
|
||||
pp_attr = _xml_attr("product", "portrait_prompt", raw)
|
||||
if pp_attr and pp_attr != "无人像":
|
||||
portrait_prompt = pp_attr
|
||||
return {
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance,
|
||||
"packaging": packaging,
|
||||
"text_on_package": text_list,
|
||||
"key_features": feat_list,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary,
|
||||
"_source": "vlm_pro_xml",
|
||||
}
|
||||
|
||||
if portrait_prompt != "无人像":
|
||||
return {
|
||||
"name": "未识别",
|
||||
"brand": "无法判断",
|
||||
"category": "无法判断",
|
||||
"appearance": "无法判断",
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": ["无法判断"],
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": "未识别",
|
||||
"_source": "vlm_pro_no_product",
|
||||
}
|
||||
return {
|
||||
"name": "未识别",
|
||||
"brand": "无法判断",
|
||||
"category": "无法判断",
|
||||
"appearance": "无法判断",
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": ["无法判断"],
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": "无人像",
|
||||
"summary": "未识别",
|
||||
"_source": "vlm_pro_no_tag",
|
||||
}
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
img_url: str,
|
||||
idx: int,
|
||||
*,
|
||||
model: str | None = None,
|
||||
timeout: int = DEFAULT_TIMEOUT,
|
||||
) -> dict[str, Any] | None:
|
||||
"""单次调用 pro VLM,解析后返回 product dict;失败返回 None。"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
except ImportError as e:
|
||||
logger.warning("[vision.vlm] 导入失败: %s", e)
|
||||
return None
|
||||
|
||||
try:
|
||||
template = get_template("image_analysis")
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(template, image_count=1, industry="通用", image_urls=f"第1张:{img_url}")
|
||||
except Exception as e:
|
||||
logger.warning("[vision.vlm] 模板加载失败: %s", e)
|
||||
return None
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
return None
|
||||
|
||||
use_model = model or DEFAULT_PRO_MODEL
|
||||
_orig_retries = client.max_retries
|
||||
client.max_retries = 0
|
||||
try:
|
||||
raw = client.vision_completion(
|
||||
messages=[{"role": "system", "content": system}, {"role": "user", "content": user}],
|
||||
images=[img_url],
|
||||
temperature=0.3,
|
||||
max_tokens=DEFAULT_MAX_TOKENS,
|
||||
timeout=timeout,
|
||||
model=use_model,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.vlm] 图片 #%d pro VLM 调用失败 elapsed=%.1fs err=%s", idx, time.time() - t0, e)
|
||||
client.max_retries = _orig_retries
|
||||
return None
|
||||
client.max_retries = _orig_retries
|
||||
|
||||
elapsed = time.time() - t0
|
||||
if not raw:
|
||||
logger.warning("[vision.vlm] 图片 #%d pro VLM 返回空 elapsed=%.1fs", idx, elapsed)
|
||||
return None
|
||||
|
||||
text = _strip_code_fence(raw)
|
||||
l, r = text.find("{"), text.rfind("}")
|
||||
if l >= 0 and r > l:
|
||||
try:
|
||||
obj = json.loads(text[l : r + 1])
|
||||
if isinstance(obj, dict):
|
||||
logger.info("[vision.vlm] 图片 #%d pro VLM JSON 完成 elapsed=%.1fs", idx, elapsed)
|
||||
return {
|
||||
"name": obj.get("name") or "未识别",
|
||||
"brand": obj.get("brand") or "无法判断",
|
||||
"category": obj.get("category") or "无法判断",
|
||||
"appearance": obj.get("appearance") or "无法判断",
|
||||
"packaging": obj.get("packaging") or "无法判断",
|
||||
"text_on_package": obj.get("text_on_package") or [],
|
||||
"key_features": obj.get("key_features") or obj.get("features") or ["无法判断"],
|
||||
"scene": obj.get("scene") or "通用",
|
||||
"mood": obj.get("mood") or "",
|
||||
"portrait_prompt": obj.get("portrait_prompt") or "无人像",
|
||||
"summary": obj.get("summary") or f"{obj.get('brand','')} {obj.get('name','')}",
|
||||
"_source": "vlm_pro_json",
|
||||
}
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
try:
|
||||
result = _xml_to_product(text, idx)
|
||||
result["_fallback_used"] = True
|
||||
result["_pro_elapsed"] = round(elapsed, 2)
|
||||
logger.info(
|
||||
"[vision.vlm] 图片 #%d pro VLM XML 完成 elapsed=%.2fs pp=%s",
|
||||
idx,
|
||||
elapsed,
|
||||
(result.get("portrait_prompt") or "")[:40],
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning("[vision.vlm] 图片 #%d 解析失败 elapsed=%.1fs err=%s head=%s", idx, elapsed, e, raw[:200])
|
||||
return None
|
||||
@@ -0,0 +1,148 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""doubao-seed-2.1-lite 强约束 JSON-only 调用。
|
||||
|
||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
||||
设计要点:
|
||||
- system prompt 极致精简,只给字段 schema 和强约束(禁止自然语言、禁止 markdown)
|
||||
- max_tokens=350(比旧 VLM 的 1200 小很多,降低延迟)
|
||||
- temperature=0.1(极低,稳定输出 JSON)
|
||||
- timeout=8s(够快,失败则由外层走 pro VLM 兜底)
|
||||
- 期望返回纯 JSON object(无 ```json 包裹、无解释文字)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 极简 system prompt:只给字段定义 + 硬性输出要求
|
||||
_FAST_SYSTEM = (
|
||||
"你是图片结构化识别器。严格按下方 JSON schema 返回一个对象,不要任何解释、"
|
||||
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填 null 或空数组。\n"
|
||||
"{\n"
|
||||
' "has_person": true/false, // 图中是否有人\n'
|
||||
' "gender": "男"/"女"/null,\n'
|
||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
||||
' "upper_wear": "上装款式,如T恤/衬衫/卫衣/毛衣/西装/夹克/连衣裙/吊带/背心/外套等",\n'
|
||||
' "upper_color": "上装主色",\n'
|
||||
' "lower_wear": "下装款式,如牛仔裤/休闲裤/短裙/长裙/短裤/西裤/运动裤等;穿连衣裙时填null",\n'
|
||||
' "lower_color": "下装主色",\n'
|
||||
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
|
||||
' "accessories": ["眼镜"/"帽子"/"项链"/"耳环"/"背包"/"手表"等数组],\n'
|
||||
' "hairstyle": "发型,如短发/长发/马尾/卷发/丸子头/光头等",\n'
|
||||
' "expression": "表情,如微笑/严肃/酷/开心等",\n'
|
||||
' "pose": "姿势,如站立/坐姿/侧身/行走等",\n'
|
||||
' "scene": "场景,如室内/街拍/户外/办公室/家居/海边/雪景/森林等",\n'
|
||||
' "style": "风格,如休闲/商务/运动/复古/潮流/甜美/酷飒/优雅/街头/法式等",\n'
|
||||
' "has_product": true/false, // 是否有明确商品展示\n'
|
||||
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
||||
' "product_name": "产品名称,非产品图填null",\n'
|
||||
' "brand": "品牌或文字标识,无则null",\n'
|
||||
' "material": "材质,如棉质/牛仔/皮革/真丝/针织/涤纶等",\n'
|
||||
' "pattern": "图案,如纯色/条纹/波点/格子/印花/碎花/Logo等",\n'
|
||||
' "colors": ["主色数组"],\n'
|
||||
' "mood": "整体氛围/情绪,如清新/活力/高级/温暖/冷峻/甜美/复古等"\n'
|
||||
"}"
|
||||
)
|
||||
|
||||
_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
|
||||
|
||||
# 默认模型
|
||||
DEFAULT_LITE_MODEL = "doubao-seed-2-1-lite-260915"
|
||||
DEFAULT_TIMEOUT = 8
|
||||
DEFAULT_MAX_TOKENS = 350
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
"""剥离 ```json ... ``` 包裹(即使要求纯 JSON,模型偶尔仍会包代码块)。"""
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
# 去掉首行 ```json
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
# 去掉尾行 ```
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
|
||||
def call_fast_json(
|
||||
img_url: str,
|
||||
*,
|
||||
model: str | None = None,
|
||||
timeout: int = DEFAULT_TIMEOUT,
|
||||
max_tokens: int = DEFAULT_MAX_TOKENS,
|
||||
) -> dict[str, Any] | None:
|
||||
"""调用 lite VLM 返回结构化 dict;失败/非 JSON 返回 None。
|
||||
|
||||
注意:不做重试(外层竞速/降级逻辑负责),max_retries=0 由外层统一设置。
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.warning("[vision.v2] doubao client 不可用,跳过 fast_json")
|
||||
return None
|
||||
|
||||
use_model = model or DEFAULT_LITE_MODEL
|
||||
# 强制不重试:lite 是快速路径,失败直接走外层 pro 兜底
|
||||
_orig_retries = client.max_retries
|
||||
client.max_retries = 0
|
||||
try:
|
||||
raw = client.vision_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": _FAST_SYSTEM},
|
||||
{"role": "user", "content": _FAST_USER},
|
||||
],
|
||||
images=[img_url],
|
||||
temperature=0.1,
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout,
|
||||
model=use_model,
|
||||
)
|
||||
finally:
|
||||
client.max_retries = _orig_retries
|
||||
elapsed = time.time() - t0
|
||||
if raw is None:
|
||||
logger.warning("[vision.v2] fast_json 返回 None elapsed=%.1fs model=%s", elapsed, use_model)
|
||||
return None
|
||||
|
||||
text = _strip_code_fence(raw)
|
||||
# 截到第一个 { 和最后一个 } 之间,容忍前后偶发文字
|
||||
l = text.find("{")
|
||||
r = text.rfind("}")
|
||||
if l >= 0 and r > l:
|
||||
text = text[l : r + 1]
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(
|
||||
"[vision.v2] fast_json JSON 解析失败 elapsed=%.1fs head=%s",
|
||||
elapsed,
|
||||
raw[:200],
|
||||
)
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
logger.warning("[vision.v2] fast_json 非 dict: %s", type(obj))
|
||||
return None
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs has_person=%s has_product=%s category=%s",
|
||||
use_model,
|
||||
elapsed,
|
||||
obj.get("has_person"),
|
||||
obj.get("has_product"),
|
||||
obj.get("category"),
|
||||
)
|
||||
return obj
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] fast_json 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return None
|
||||
@@ -241,19 +241,35 @@ DOUBAO_API_KEY=${DOUBAO_API_KEY}
|
||||
|
||||
# 模型 Endpoint ID(在 ARK 控制台创建推理接入点后获得)
|
||||
DOUBAO_MODEL=${DOUBAO_MODEL}
|
||||
DOUBAO_FAST_MODEL=${DOUBAO_FAST_MODEL}
|
||||
|
||||
# API Base URL
|
||||
DOUBAO_BASE_URL=${DOUBAO_BASE_URL}
|
||||
|
||||
# 请求超时(秒)
|
||||
DOUBAO_TIMEOUT=60
|
||||
DOUBAO_TIMEOUT=${DOUBAO_TIMEOUT}
|
||||
|
||||
# 最大重试次数
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
DOUBAO_MAX_RETRIES=${DOUBAO_MAX_RETRIES}
|
||||
|
||||
# 视觉模型 Endpoint ID(支持图片/视频理解的模型)
|
||||
# 视觉模型(支持图片/视频理解的模型,model name 格式)
|
||||
DOUBAO_VISION_MODEL=${DOUBAO_VISION_MODEL}
|
||||
|
||||
# 快速视觉模型(viral-video 图片分析 lite 路径)
|
||||
DOUBAO_VISION_LITE_MODEL=${DOUBAO_VISION_LITE_MODEL}
|
||||
|
||||
# 是否启用 lite 视觉路径(true/false)
|
||||
DOUBAO_VISION_USE_LITE=${DOUBAO_VISION_USE_LITE}
|
||||
|
||||
# 信任链文生图模型(Seedream)
|
||||
DOUBAO_IMAGE_MODEL=${DOUBAO_IMAGE_MODEL}
|
||||
|
||||
# 文生图尺寸
|
||||
DOUBAO_IMAGE_SIZE=${DOUBAO_IMAGE_SIZE}
|
||||
|
||||
# 文生图超时(秒)
|
||||
DOUBAO_IMAGE_TIMEOUT=${DOUBAO_IMAGE_TIMEOUT}
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
|
||||
@@ -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 部署。
|
||||
|
||||
@@ -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"]
|
||||
|
||||
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
|
||||
@@ -18,6 +18,17 @@
|
||||
|
||||
set -e
|
||||
|
||||
# #2129: 幂等执行数据库迁移(watchtower 自动部署兜底)
|
||||
# worker 容器独立启动,不能依赖 API 容器先跑迁移
|
||||
cd /app
|
||||
echo "[entrypoint-worker] Running alembic upgrade head..."
|
||||
if alembic upgrade head; then
|
||||
echo "[entrypoint-worker] Migrations ok."
|
||||
else
|
||||
echo "[entrypoint-worker] WARNING: alembic upgrade failed, continuing to start workers..." >&2
|
||||
fi
|
||||
cd - >/dev/null
|
||||
|
||||
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
|
||||
|
||||
# ── 并发计算:显式 env 优先;否则从 WORKER_CONCURRENCY 按比例推导 ──
|
||||
|
||||
@@ -34,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,
|
||||
|
||||
@@ -502,8 +502,19 @@ class SQLAlchemyAssetRepository:
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key(对应 DB 中的 file_url)查找素材。"""
|
||||
model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first()
|
||||
"""按 storage_key 查找素材。
|
||||
|
||||
Bug #2110: 历史数据 file_url 列可能是旧路径(assets/...),新代码统一写入
|
||||
storage_key 列。双列 OR 查询,避免占位 asset 因路径错配导致 ingest 兜底新建
|
||||
第二条 READY 记录,原占位卡 PROCESSING → 前端缩略图出现后消失。
|
||||
"""
|
||||
if not storage_key:
|
||||
return None
|
||||
model = (
|
||||
self.session.query(AssetModel)
|
||||
.filter((AssetModel.storage_key == storage_key) | (AssetModel.file_url == storage_key))
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
@@ -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,90 @@ 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 列表
|
||||
pre_trusted_images = Column(JSON, nullable=True) # #2172 信任链预热结果(Seedream AI 化 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:纯文本 XML 标签模板,运营可直接编辑)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
name = Column(String(128), nullable=False)
|
||||
prompt_type = Column(String(32), nullable=False)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
system_prompt = Column(Text, nullable=False)
|
||||
user_prompt_template = Column(Text, nullable=False)
|
||||
example_output = Column(Text, nullable=True)
|
||||
is_active = Column(Boolean, nullable=False, default=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()
|
||||
|
||||
@@ -39,6 +39,9 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.profile_completed = user.profile_completed
|
||||
model.is_member = user.is_member
|
||||
model.member_type = user.member_type
|
||||
model.member_expires_at = user.member_expires_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -115,5 +118,8 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
profile_completed=model.profile_completed if model.profile_completed is not None else True,
|
||||
is_member=model.is_member if model.is_member is not None else False,
|
||||
member_type=model.member_type,
|
||||
member_expires_at=model.member_expires_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
+246
@@ -0,0 +1,246 @@
|
||||
"""爆款视频任务 SQLAlchemy 仓储实现。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
ViralVideoJobModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
"""ORM → 领域实体。pre_trusted_images 兼容脏数据:双序列化字符串/字符数组/list[str]。"""
|
||||
import json as _pti_json
|
||||
|
||||
_raw_pti = getattr(model, "pre_trusted_images", None)
|
||||
_pti: list[str] | None = None
|
||||
if _raw_pti is not None:
|
||||
if isinstance(_raw_pti, str):
|
||||
try:
|
||||
_p = _pti_json.loads(_raw_pti)
|
||||
if isinstance(_p, list):
|
||||
_pti = [u for u in _p if isinstance(u, str) and u] or None
|
||||
except Exception:
|
||||
_pti = None
|
||||
elif isinstance(_raw_pti, list):
|
||||
_f = [u for u in _raw_pti if isinstance(u, str) and len(u) > 5]
|
||||
_pti = _f if _f else None
|
||||
return ViralVideoJob(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
images=list(model.images or []),
|
||||
pre_trusted_images=_pti,
|
||||
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,
|
||||
pre_trusted_images=job.pre_trusted_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.pre_trusted_images = job.pre_trusted_images
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -0,0 +1,152 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 MEMBERSHIP_PRICES / POINTS_PACKAGES。
|
||||
短 TTL 缓存(30 秒),后台改价/启停后用户端最多 30 秒可见。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
_CACHE_TTL = 30.0
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
|
||||
_QUOTA_LABELS = {
|
||||
"4k": "4K 超清分辨率",
|
||||
"batch_render": "批量渲染",
|
||||
"priority_queue": "优先处理队列",
|
||||
"ai_matting": "AI 智能抠像",
|
||||
"remove_watermark": "去水印",
|
||||
}
|
||||
|
||||
|
||||
def _cached(key: str, loader):
|
||||
now = time.time()
|
||||
hit = _cache.get(key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
with _lock:
|
||||
hit = _cache.get(key)
|
||||
if hit and time.time() - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
value = loader()
|
||||
_cache[key] = (time.time(), value)
|
||||
return value
|
||||
|
||||
|
||||
def _quota_features(quotas: dict[str, Any] | None) -> dict[str, Any]:
|
||||
quotas = quotas or {}
|
||||
features: dict[str, Any] = {}
|
||||
for k, v in quotas.items():
|
||||
if k == "credits_per_month":
|
||||
features["credits_per_month"] = v
|
||||
elif k in _QUOTA_LABELS:
|
||||
features[_QUOTA_LABELS[k]] = v
|
||||
else:
|
||||
features[k] = v
|
||||
return features
|
||||
|
||||
|
||||
def get_membership_plans() -> list[dict[str, Any]]:
|
||||
"""读取 is_enabled=true 的套餐,按年/月周期展开为用户端档位。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT plan_key, name, description, monthly_price, yearly_price,
|
||||
quotas, display_order
|
||||
FROM plans
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY display_order NULLS LAST, created_at
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
base_features = _quota_features(r.quotas if isinstance(r.quotas, dict) else None)
|
||||
if r.yearly_price and float(r.yearly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "yearly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.yearly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.yearly_price) * 100 / 12)),
|
||||
"duration_days": 365,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
if r.monthly_price and float(r.monthly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "monthly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"duration_days": 30,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
return plans
|
||||
|
||||
return _cached("membership_plans", _load)
|
||||
|
||||
|
||||
def get_points_packages() -> list[dict[str, Any]]:
|
||||
"""读取 is_active=true 的积分充值包。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT package_key, name, price, credits, bonus_credits,
|
||||
is_recommended, description, sort_order
|
||||
FROM credit_packages
|
||||
WHERE is_active = TRUE
|
||||
ORDER BY sort_order NULLS LAST, price
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
packages: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
total_points = int(r.credits or 0) + int(r.bonus_credits or 0)
|
||||
price_cents = int(round(float(r.price) * 100))
|
||||
unit = (price_cents / 100 / total_points) if total_points else 0
|
||||
packages.append(
|
||||
{
|
||||
"code": r.package_key,
|
||||
"name": r.name,
|
||||
"points": total_points,
|
||||
"bonus_credits": int(r.bonus_credits or 0),
|
||||
"price_cents": price_cents,
|
||||
"unit_price": f"¥{unit:.3f}/积分",
|
||||
"is_recommended": bool(r.is_recommended),
|
||||
"description": r.description,
|
||||
}
|
||||
)
|
||||
return packages
|
||||
|
||||
return _cached("points_packages", _load)
|
||||
@@ -0,0 +1 @@
|
||||
"""应用层:爆款视频 Prompt 模板系统(#2040)。"""
|
||||
@@ -0,0 +1,427 @@
|
||||
"""爆款视频 5 步编排:图片分析 → 意图解析 → 文案融合 → 分镜 → 审核重写。
|
||||
|
||||
所有 LLM 调用走 DoubaoClient,单测通过 client 参数注入 mock,不真调 API。
|
||||
任何一步解析失败都走规则 fallback,不抛异常阻断。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.viral_video import xml_parser as xp
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
PromptTemplate,
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.application.viral_video.prompts import (
|
||||
FUSION_INSTRUCTIONS,
|
||||
GLOBAL_CONSTRAINTS,
|
||||
NEGATIVE_RULES,
|
||||
)
|
||||
from packages.application.viral_video.reviewer import Reviewer
|
||||
from packages.application.viral_video.schemas import (
|
||||
BodyPoint,
|
||||
Clip,
|
||||
ColorItem,
|
||||
CoreMessage,
|
||||
FusionResult,
|
||||
ImageAnalysis,
|
||||
IntentResult,
|
||||
KenBurns,
|
||||
PersonalBrand,
|
||||
ProductItem,
|
||||
ReviewResult,
|
||||
ScriptSegment,
|
||||
Storyboard,
|
||||
TextItem,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CopyGenerator:
|
||||
"""5 步 Prompt 编排器。"""
|
||||
|
||||
def __init__(self, client=None, reviewer: Optional[Reviewer] = None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
self.client = client
|
||||
self.reviewer = reviewer or Reviewer(client)
|
||||
|
||||
# ── 底层调用 ────────────────────────────────────────────────────────
|
||||
def _chat(self, template: PromptTemplate, system_kwargs: dict | None, **user_kwargs) -> str:
|
||||
system = render_system_prompt(template, **(system_kwargs or {}))
|
||||
user = render_user_prompt(template, **user_kwargs)
|
||||
result = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.7,
|
||||
max_tokens=2048,
|
||||
)
|
||||
return result or ""
|
||||
|
||||
# ── 步骤1:图片多模态分析 ───────────────────────────────────────────
|
||||
def analyze_images(self, images: list[str], industry: str = "") -> ImageAnalysis:
|
||||
template = get_template("image_analysis")
|
||||
image_urls = "\n".join(f"第{i + 1}张:{url}" for i, url in enumerate(images))
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(template, image_count=len(images), industry=industry or "通用", image_urls=image_urls)
|
||||
raw = self.client.vision_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
images=images,
|
||||
max_tokens=2048,
|
||||
temperature=0.3,
|
||||
)
|
||||
analysis = self._parse_image_analysis(raw or "")
|
||||
if not analysis.products and not analysis.key_selling_points:
|
||||
logger.warning("图片分析标签解析失败,走规则 fallback")
|
||||
return self._fallback_image_analysis(images, raw or "")
|
||||
return analysis
|
||||
|
||||
def _parse_image_analysis(self, raw: str) -> ImageAnalysis:
|
||||
products = [
|
||||
ProductItem(
|
||||
name=n["attrs"].get("name", "无法判断"),
|
||||
features=n["attrs"].get("features", "无法判断"),
|
||||
position=n["attrs"].get("position", "secondary"),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "product")
|
||||
]
|
||||
colors = [
|
||||
ColorItem(
|
||||
hex=c["attrs"].get("hex", "#000000"),
|
||||
name=c["attrs"].get("name", "无法判断"),
|
||||
coverage=xp.attr_float(c["attrs"].get("coverage"), 0.0),
|
||||
)
|
||||
for c in xp.find_all(raw, "color")
|
||||
]
|
||||
people = xp.find_first(raw, "people")
|
||||
visible_text = [
|
||||
TextItem(text=t["attrs"].get("text", ""), position=t["attrs"].get("position", ""))
|
||||
for t in xp.find_all(raw, "text_item")
|
||||
]
|
||||
quality_node = xp.find_first(raw, "quality")
|
||||
selling_points = [n["text"] or n["attrs"].get("text", "") for n in xp.find_all(raw, "point")]
|
||||
return ImageAnalysis(
|
||||
products=products,
|
||||
colors=colors,
|
||||
has_person=xp.attr_bool(people["attrs"].get("has_person")) if people else False,
|
||||
person_count=xp.attr_int(people["attrs"].get("count"), 0) if people else 0,
|
||||
people=people["attrs"] if people else {},
|
||||
mood=xp.text_of(raw, "mood"),
|
||||
visible_text=visible_text,
|
||||
scene=xp.text_of(raw, "scene"),
|
||||
quality=quality_node["attrs"] if quality_node else {},
|
||||
key_selling_points=[p for p in selling_points if p],
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_image_analysis(self, images: list[str], raw: str) -> ImageAnalysis:
|
||||
return ImageAnalysis(
|
||||
products=[ProductItem(name="无法判断(视觉分析不可用)", image_index=0)],
|
||||
scene="无法判断",
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤2:意图解析 ─────────────────────────────────────────────────
|
||||
def parse_intent(self, user_copy_text: str, image_analysis: ImageAnalysis, industry: str = "") -> IntentResult:
|
||||
template = get_template("intent_parsing")
|
||||
raw = self._chat(
|
||||
template,
|
||||
None,
|
||||
user_copy_text=user_copy_text or "(用户没有提供文案)",
|
||||
industry=industry or "通用",
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
)
|
||||
intent = self._parse_intent(raw)
|
||||
if not intent.intent_summary and not intent.core_messages:
|
||||
logger.warning("意图解析标签解析失败,走规则 fallback")
|
||||
return self._fallback_intent(user_copy_text, raw)
|
||||
return intent
|
||||
|
||||
def _parse_intent(self, raw: str) -> IntentResult:
|
||||
messages = [
|
||||
CoreMessage(
|
||||
text=n["text"],
|
||||
must_keep=xp.attr_bool(n["attrs"].get("must_keep"), default=False),
|
||||
confidence=xp.attr_float(n["attrs"].get("confidence"), 0.0),
|
||||
)
|
||||
for n in xp.find_all(raw, "message")
|
||||
if n["text"]
|
||||
]
|
||||
brands = [
|
||||
PersonalBrand(text=n["text"], category=n["attrs"].get("category", "brand"))
|
||||
for n in xp.find_all(raw, "brand")
|
||||
if n["text"]
|
||||
]
|
||||
missing = [n["text"] for n in xp.find_all(raw, "info") if n["text"]]
|
||||
return IntentResult(
|
||||
intent_summary=xp.text_of(raw, "intent_summary"),
|
||||
core_messages=messages,
|
||||
personal_brands=brands,
|
||||
emotion_tone=xp.text_of(raw, "emotion_tone"),
|
||||
missing_info=missing,
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_intent(self, user_copy_text: str, raw: str) -> IntentResult:
|
||||
text = (user_copy_text or "").strip()
|
||||
messages = [CoreMessage(text=text[:80], must_keep=True, confidence=1.0)] if text else []
|
||||
return IntentResult(
|
||||
intent_summary=text[:30] or "未提供文案,按产品图片自由创作",
|
||||
core_messages=messages,
|
||||
personal_brands=[],
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤3:文案融合生成(三档)──────────────────────────────────────
|
||||
def fuse(
|
||||
self,
|
||||
fusion_level: str,
|
||||
image_analysis: ImageAnalysis,
|
||||
intent: IntentResult,
|
||||
industry: str = "",
|
||||
target_customer: str = "",
|
||||
marketing_purpose: str = "",
|
||||
duration: int = 15,
|
||||
) -> FusionResult:
|
||||
template = get_template("copy_fusion")
|
||||
system_kwargs = {
|
||||
"fusion_instruction": FUSION_INSTRUCTIONS.get(fusion_level, FUSION_INSTRUCTIONS["ai_polish"]),
|
||||
"global_constraints": GLOBAL_CONSTRAINTS,
|
||||
"negative_rules": NEGATIVE_RULES,
|
||||
}
|
||||
raw = self._chat(
|
||||
template,
|
||||
system_kwargs,
|
||||
industry=industry or "通用",
|
||||
target_customer=target_customer or "通用消费者",
|
||||
marketing_purpose=marketing_purpose or "产品种草",
|
||||
duration=duration,
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
intent_result=self._intent_brief(intent),
|
||||
)
|
||||
result = self._parse_fusion(raw)
|
||||
if not result.title and not result.script_segments:
|
||||
logger.warning("文案融合标签解析失败(fusion=%s),走规则 fallback", fusion_level)
|
||||
return self._fallback_fusion(fusion_level, image_analysis, intent, duration, raw)
|
||||
return result
|
||||
|
||||
def _parse_fusion(self, raw: str) -> FusionResult:
|
||||
body_points = [
|
||||
BodyPoint(
|
||||
text=n["text"],
|
||||
elaboration=n["attrs"].get("elaboration", ""),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "point")
|
||||
if n["text"]
|
||||
]
|
||||
segments = [
|
||||
ScriptSegment(
|
||||
text=n["text"],
|
||||
duration_sec=xp.attr_float(n["attrs"].get("duration_sec"), 0.0),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "segment")
|
||||
if n["text"]
|
||||
]
|
||||
return FusionResult(
|
||||
title=xp.text_of(raw, "title"),
|
||||
hook=xp.text_of(raw, "hook"),
|
||||
body_points=body_points,
|
||||
cta=xp.text_of(raw, "cta"),
|
||||
script_segments=segments,
|
||||
word_count=xp.attr_int(xp.text_of(raw, "word_count"), 0),
|
||||
estimated_duration=xp.attr_int(xp.text_of(raw, "estimated_duration"), 0),
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_fusion(
|
||||
self,
|
||||
fusion_level: str,
|
||||
image_analysis: ImageAnalysis,
|
||||
intent: IntentResult,
|
||||
duration: int,
|
||||
raw: str,
|
||||
) -> FusionResult:
|
||||
product_name = image_analysis.products[0].name if image_analysis.products else "这款产品"
|
||||
selling = image_analysis.key_selling_points[:2]
|
||||
if fusion_level == "ai_full":
|
||||
title = f"{product_name},很多人用完都回购了"
|
||||
hook = f"这个{product_name},我想认真说说"
|
||||
body = selling or ["图片可见的产品卖点"]
|
||||
cta = "感兴趣的可以了解一下"
|
||||
elif fusion_level == "user_primary":
|
||||
user_text = intent.intent_summary or product_name
|
||||
title = user_text[:20]
|
||||
hook = user_text[:15]
|
||||
body = [m.text for m in intent.core_messages] or [user_text]
|
||||
cta = "想了解的可以看看"
|
||||
else:
|
||||
title = intent.intent_summary[:20] or product_name
|
||||
hook = intent.core_messages[0].text[:15] if intent.core_messages else product_name
|
||||
body = [m.text for m in intent.core_messages] or selling or [product_name]
|
||||
cta = "有需要的可以了解一下"
|
||||
|
||||
brand_texts = [b.text for b in intent.personal_brands]
|
||||
points = [BodyPoint(text=b) for b in body]
|
||||
lines = [hook] + body + brand_texts[:2] + [cta]
|
||||
joined = ",".join(lines)
|
||||
per = max(3, duration // max(1, len(lines)))
|
||||
segments = [ScriptSegment(text=line, duration_sec=per, image_index=0) for line in lines]
|
||||
return FusionResult(
|
||||
title=title,
|
||||
hook=hook,
|
||||
body_points=points,
|
||||
cta=cta,
|
||||
script_segments=segments,
|
||||
word_count=len(joined),
|
||||
estimated_duration=duration,
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤4:编导级分镜 ───────────────────────────────────────────────
|
||||
def storyboard(
|
||||
self, fusion: FusionResult, image_analysis: ImageAnalysis, images: list[str], duration: int
|
||||
) -> Storyboard:
|
||||
template = get_template("storyboard")
|
||||
raw = self._chat(
|
||||
template,
|
||||
None,
|
||||
duration=duration,
|
||||
image_count=len(images),
|
||||
fusion_result=self._fusion_brief(fusion),
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
)
|
||||
board = self._parse_storyboard(raw)
|
||||
if not board.clips:
|
||||
logger.warning("分镜标签解析失败,走规则 fallback")
|
||||
return self._fallback_storyboard(fusion, duration, raw)
|
||||
return board
|
||||
|
||||
def _parse_storyboard(self, raw: str) -> Storyboard:
|
||||
clips: list[Clip] = []
|
||||
for node in xp.find_all(raw, "clip"):
|
||||
attrs = node["attrs"]
|
||||
body = node["text"]
|
||||
kb = xp.find_first(node["text"] and f"<root>{node['text']}</root>", "ken_burns")
|
||||
clips.append(
|
||||
Clip(
|
||||
image_index=xp.attr_int(attrs.get("image_index"), 0),
|
||||
transition=attrs.get("transition", "cut"),
|
||||
zoom=(None if attrs.get("zoom") in (None, "null", "None", "") else attrs.get("zoom")),
|
||||
duration_sec=xp.attr_float(attrs.get("duration_sec"), 0.0),
|
||||
bgm_note=attrs.get("bgm_note", ""),
|
||||
voice_text=xp.text_of(body and f"<root>{body}</root>", "voice_text"),
|
||||
subtitle_text=xp.text_of(body and f"<root>{body}</root>", "subtitle_text"),
|
||||
ken_burns=KenBurns(
|
||||
start=kb["attrs"].get("start", "0,0") if kb else "0,0",
|
||||
end=kb["attrs"].get("end", "0,0") if kb else "0,0",
|
||||
ease=kb["attrs"].get("ease", "linear") if kb else "linear",
|
||||
),
|
||||
)
|
||||
)
|
||||
return Storyboard(clips=clips, raw=raw)
|
||||
|
||||
def _fallback_storyboard(self, fusion: FusionResult, duration: int, raw: str) -> Storyboard:
|
||||
segments = fusion.script_segments or [ScriptSegment(text=fusion.hook or fusion.title, duration_sec=duration)]
|
||||
total = sum(s.duration_sec for s in segments) or duration
|
||||
clips = [
|
||||
Clip(
|
||||
image_index=min(s.image_index, 0),
|
||||
transition="cut",
|
||||
duration_sec=max(2.0, s.duration_sec * duration / total if total else duration / len(segments)),
|
||||
voice_text=s.text,
|
||||
subtitle_text=s.text[:20],
|
||||
)
|
||||
for s in segments
|
||||
]
|
||||
return Storyboard(clips=clips, raw=raw)
|
||||
|
||||
# ── 步骤5:审核(不通过自动重写1次)─────────────────────────────────
|
||||
def review_and_rewrite(
|
||||
self, fusion: FusionResult, intent: IntentResult, fusion_level: str
|
||||
) -> tuple[FusionResult, ReviewResult, int]:
|
||||
"""返回最终文案、最后一次审核结果、重写次数(0或1)。"""
|
||||
review = self.reviewer.review(fusion, intent, fusion_level)
|
||||
if review.passed:
|
||||
return fusion, review, 0
|
||||
|
||||
logger.info("文案审核不通过,自动重写 1 次:%s", [i.text for i in review.issues])
|
||||
rewritten = self.reviewer.rewrite(fusion, review, intent, fusion_level)
|
||||
second = self.reviewer.review(rewritten, intent, fusion_level)
|
||||
if second.passed:
|
||||
return rewritten, second, 1
|
||||
# 二次仍不通过:带上重写结果和问题返回,由上游决定是否交给前端
|
||||
return rewritten, second, 1
|
||||
|
||||
# ── 全流程编排 ──────────────────────────────────────────────────────
|
||||
def generate(
|
||||
self,
|
||||
images: list[str],
|
||||
*,
|
||||
industry: str = "",
|
||||
target_customer: str = "",
|
||||
marketing_purpose: str = "",
|
||||
duration: int = 15,
|
||||
user_copy_text: str = "",
|
||||
fusion_level: str = "ai_polish",
|
||||
) -> dict:
|
||||
image_analysis = self.analyze_images(images, industry)
|
||||
intent = self.parse_intent(user_copy_text, image_analysis, industry)
|
||||
fusion = self.fuse(
|
||||
fusion_level,
|
||||
image_analysis,
|
||||
intent,
|
||||
industry=industry,
|
||||
target_customer=target_customer,
|
||||
marketing_purpose=marketing_purpose,
|
||||
duration=duration,
|
||||
)
|
||||
fusion, review, rewrites = self.review_and_rewrite(fusion, intent, fusion_level)
|
||||
board = self.storyboard(fusion, image_analysis, images, duration)
|
||||
return {
|
||||
"image_analysis": image_analysis,
|
||||
"intent_result": intent,
|
||||
"fusion_result": fusion,
|
||||
"review_result": review,
|
||||
"storyboard": board,
|
||||
"rewrite_count": rewrites,
|
||||
}
|
||||
|
||||
# ── 简报工具 ────────────────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _image_brief(a) -> str:
|
||||
if a is None:
|
||||
return "无图片分析信息"
|
||||
lines = [f"产品:{p.name}({p.features})" for p in a.products]
|
||||
lines += [f"卖点:{s}" for s in a.key_selling_points]
|
||||
lines.append(f"场景:{a.scene}")
|
||||
return "\n".join(lines) or "无图片分析信息"
|
||||
|
||||
@staticmethod
|
||||
def _intent_brief(i: IntentResult) -> str:
|
||||
lines = [f"意图:{i.intent_summary}"]
|
||||
lines += [f"核心信息[must_keep={m.must_keep}]:{m.text}" for m in i.core_messages]
|
||||
lines += [f"事实({b.category}):{b.text}" for b in i.personal_brands]
|
||||
return "\n".join(lines)
|
||||
|
||||
@staticmethod
|
||||
def _fusion_brief(f: FusionResult) -> str:
|
||||
lines = [f"标题:{f.title}", f"钩子:{f.hook}"]
|
||||
lines += [f"要点:{p.text}" for p in f.body_points]
|
||||
lines += [f"配音:{s.text}" for s in f.script_segments]
|
||||
lines.append(f"行动号召:{f.cta}")
|
||||
return "\n".join(lines)
|
||||
@@ -0,0 +1,160 @@
|
||||
"""Prompt 模板加载器:从 viral_video_prompt_templates 读模板,30 秒 TTL 热加载。
|
||||
|
||||
DB 不可用或没有数据时自动回落到 prompts.DEFAULT_TEMPLATES,保证流程不阻断。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as _session_mod
|
||||
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
|
||||
|
||||
CACHE_TTL_SECONDS = 30.0
|
||||
|
||||
_VALID_TYPES = {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptTemplate:
|
||||
name: str
|
||||
prompt_type: str
|
||||
version: int
|
||||
system_prompt: str
|
||||
user_prompt_template: str
|
||||
example_output: str = ""
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, PromptTemplate]] = {}
|
||||
|
||||
|
||||
def _fallback(prompt_type: str) -> Optional[PromptTemplate]:
|
||||
for item in DEFAULT_TEMPLATES:
|
||||
if item["prompt_type"] == prompt_type:
|
||||
return PromptTemplate(
|
||||
name=item["name"],
|
||||
prompt_type=item["prompt_type"],
|
||||
version=item["version"],
|
||||
system_prompt=item["system_prompt"],
|
||||
user_prompt_template=item["user_prompt_template"],
|
||||
example_output=item["example_output"] or "",
|
||||
is_active=bool(item["is_active"]),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_lazy_session = None
|
||||
|
||||
|
||||
def _get_session():
|
||||
"""优先用全局 SessionLocal(worker);否则按应用配置懒建同步引擎(api)。"""
|
||||
global _lazy_session
|
||||
if _session_mod.SessionLocal is not None:
|
||||
return _session_mod.SessionLocal()
|
||||
if _lazy_session is not None:
|
||||
return _lazy_session()
|
||||
try:
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
url = str(get_shared_settings().database_url)
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if not url:
|
||||
return None
|
||||
url = url.replace("postgresql+asyncpg://", "postgresql+psycopg://")
|
||||
url = url.replace("postgresql://", "postgresql+psycopg://") if url.startswith("postgresql://") else url
|
||||
engine = sa.create_engine(url, pool_pre_ping=True, pool_size=2, max_overflow=2)
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
_lazy_session = sessionmaker(bind=engine)
|
||||
return _lazy_session()
|
||||
|
||||
|
||||
def _load_from_db(prompt_type: str) -> Optional[PromptTemplate]:
|
||||
session = None
|
||||
try:
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
return None
|
||||
sql = sa.text("""
|
||||
SELECT name, prompt_type, version, system_prompt,
|
||||
user_prompt_template, COALESCE(example_output, '') AS example_output,
|
||||
is_active
|
||||
FROM viral_video_prompt_templates
|
||||
WHERE prompt_type = :pt AND is_active = TRUE
|
||||
ORDER BY version DESC
|
||||
LIMIT 1
|
||||
""")
|
||||
row = session.execute(sql, {"pt": prompt_type}).first()
|
||||
if row is None:
|
||||
return None
|
||||
return PromptTemplate(
|
||||
name=row[0],
|
||||
prompt_type=row[1],
|
||||
version=int(row[2]),
|
||||
system_prompt=row[3],
|
||||
user_prompt_template=row[4],
|
||||
example_output=row[5] or "",
|
||||
is_active=bool(row[6]),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落
|
||||
return None
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def get_template(prompt_type: str, *, force_refresh: bool = False) -> Optional[PromptTemplate]:
|
||||
"""取某类型当前启用模板,30 秒缓存;DB 无数据则回落到代码默认模板。"""
|
||||
if prompt_type not in _VALID_TYPES:
|
||||
raise ValueError(f"未知 prompt_type: {prompt_type}")
|
||||
|
||||
now = time.monotonic()
|
||||
with _lock:
|
||||
cached = _cache.get(prompt_type)
|
||||
if not force_refresh and cached and now - cached[0] < CACHE_TTL_SECONDS:
|
||||
return cached[1]
|
||||
|
||||
template = _load_from_db(prompt_type) or _fallback(prompt_type)
|
||||
if template is not None:
|
||||
with _lock:
|
||||
_cache[prompt_type] = (now, template)
|
||||
return template
|
||||
|
||||
|
||||
def invalidate() -> None:
|
||||
"""清空缓存(测试用)。"""
|
||||
with _lock:
|
||||
_cache.clear()
|
||||
|
||||
|
||||
class _SafeDict(dict):
|
||||
def __missing__(self, key: str) -> str:
|
||||
return "{" + key + "}"
|
||||
|
||||
|
||||
def _safe_format(text: str, kwargs: dict) -> str:
|
||||
try:
|
||||
return text.format_map(_SafeDict(kwargs))
|
||||
except Exception: # noqa: BLE001
|
||||
return text
|
||||
|
||||
|
||||
def render_user_prompt(template: PromptTemplate, **kwargs) -> str:
|
||||
"""填充 user_prompt_template 占位符,缺键原样保留不报错。"""
|
||||
return _safe_format(template.user_prompt_template, kwargs)
|
||||
|
||||
|
||||
def render_system_prompt(template: PromptTemplate, **kwargs) -> str:
|
||||
"""copy_fusion 等 system_prompt 含运行时变量时填充。"""
|
||||
return _safe_format(template.system_prompt, kwargs)
|
||||
@@ -0,0 +1,338 @@
|
||||
"""爆款视频 5 套 Prompt 模板默认值(#2040 核心资产)。
|
||||
|
||||
重要约定(用户明确要求):
|
||||
- 所有 system_prompt / user_prompt_template / example_output 都是**纯文本自然语言 + XML 标签**,
|
||||
运营可直接看懂和编辑,禁止 JSON、禁止 ```json 代码块。
|
||||
- LLM 按 XML 标签输出字段,程序用正则解析(见 xml_parser.py)。
|
||||
- user_prompt_template 中花括号占位符(如 {user_copy_text})在运行时填充。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
TEMPLATE_VERSION = 1
|
||||
|
||||
# 所有文案类 Prompt 自动注入的硬约束
|
||||
GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
||||
1. 不编造时间:不写“今年最新”“2024 爆款”等会过时的时间表述。
|
||||
2. 不承诺效果:不写“保证”“一定”“100%有效”“包治百病”等绝对化用语。
|
||||
3. 不编造价格、销量、认证、奖项:除非用户在文案中明确给出,否则一律不写。
|
||||
4. 符合广告法及平台社区规范。
|
||||
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
|
||||
|
||||
# 反套路化要求
|
||||
NEGATIVE_RULES = """【反套路化要求】
|
||||
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
|
||||
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
|
||||
|
||||
# 输出禁用套路词(测试会检查)
|
||||
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
|
||||
|
||||
# 文案融合三档独立指令段
|
||||
FUSION_INSTRUCTIONS = {
|
||||
"ai_full": """【本次创作模式:AI 全权创作】
|
||||
你是资深短视频编导。用户只提供了产品图片,没有给出具体文案方向。请根据图片内容和营销参数,自由发挥创作完整的爆款短视频文案。充分挖掘产品真实可见的卖点,使用爆款结构,抓人眼球。""",
|
||||
"ai_polish": """【本次创作模式:AI 辅助润色】
|
||||
你是用户的文案助理。用户已经写了草稿/关键词/碎碎念,表达了他想讲的核心意思,但表达不完整、不够吸引人。你的任务是:以用户的意思为主,保留他想表达的所有核心信息点,在此基础上润色扩写、调整语序、增加衔接、优化表达,让文案更流畅更有吸引力。绝对不能改变用户想表达的核心意思,不能把用户的观点换成相反的,不能添加用户没提到的产品卖点。用户提到的品牌名、价格、人名、具体事实必须原样保留。""",
|
||||
"user_primary": """【本次创作模式:以用户原文为主】
|
||||
你是文案润色助手。用户已经写好了明确的文案,这是他最终想表达的内容。你的任务是最小化修改:只做必要的错别字修正、标点调整、语句通顺度优化,以及添加必要的衔接词让口播更自然。用户的核心句子、关键表述、事实信息一律不改。如果用户文案本身已经很好,直接返回,不要为了改而改。personal_brands 中的事实信息必须逐字保留。""",
|
||||
}
|
||||
|
||||
# ── 模板1:图片多模态分析(VLM)────────────────────────────────────────
|
||||
_IMAGE_ANALYSIS_SYSTEM = f"""你是电商商品视觉分析师,负责从商品图片中提取真实可见的商品信息。
|
||||
|
||||
工作方式(分步骤看,不要跳步):
|
||||
1. 先看整体:有哪些产品、什么场景、有没有人物。
|
||||
2. 再看细节:包装文字、颜色构成、人物状态、画面质感。
|
||||
3. 最后提炼卖点:只总结图片里能看到的卖点。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,标签名一个都不能改,不要输出任何解释,不要用代码块:
|
||||
<products> 下面每个产品用一个 <product> 标签,属性 name 是产品名、features 是外观特征、position 是 main 或 secondary、image_index 是第几张图(从0开始)。
|
||||
<colors> 下面每个主要颜色用一个 <color> 标签,属性 hex 是色值、name 是颜色名、coverage 是占比小数。
|
||||
<people> 用一个标签,属性 has_person、count、gender、age_range、hair(发型发色)、skin_tone(肤色)、face_shape(脸型)、outfit(穿着)、pose(姿态)、expression(表情)分别描述人物外貌。有人物时属性尽量具体(如hair="黑色长直发"、outfit="白色衬衫"),无人像时除has_person=false外其他填"无法判断"。
|
||||
<mood> 标签写画面整体情绪氛围。
|
||||
<visible_text> 下面每处可见文字用一个 <text_item> 标签,属性 text 是文字内容、position 是位置。
|
||||
<scene> 标签写场景描述。
|
||||
<quality> 用一个标签,属性 resolution、lighting、composition、blur 描述画质。
|
||||
<key_selling_points> 下面每个卖点用一个 <point> 标签。
|
||||
|
||||
【人物属性硬性要求(has_person=true时必须遵守)】
|
||||
hair/skin_tone/face_shape/outfit四项绝对禁止填“无法判断”,必须基于图片可见特征给出具体中文描述:
|
||||
- hair:必须描述发型+发色,如“黑色齐肩直发”“棕色微卷中长发”“深棕色短发”
|
||||
- skin_tone:必须描述肤色,如“暖调自然肤色”“白皙肤色”“小麦色”
|
||||
- face_shape:必须描述脸型,如“鹅蛋脸”“圆脸”“瓜子脸”“方脸”
|
||||
- outfit:必须描述可见穿着,如“米色翻领衬衫”“白色T恤”“黑色连衣裙”
|
||||
即使局部被遮挡也要根据可见部分合理推断;确实看不清时按最接近的直观印象描述。
|
||||
|
||||
其他非人物属性看不到或无法判断时填“无法判断”,布尔值填false,不要留空标签。
|
||||
|
||||
【有人物场景输出参考(女性手持商品示例,必须写全10个属性,禁止省略)】
|
||||
<people has_person="true" count="1" gender="女" age_range="青年" hair="黑色齐肩直发" skin_tone="暖调自然肤色" face_shape="鹅蛋脸" outfit="米色翻领衬衫" pose="正面半身,手持商品" expression="面带微笑"/>"""
|
||||
|
||||
_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。
|
||||
所属行业:{industry}
|
||||
图片地址:
|
||||
{image_urls}
|
||||
|
||||
按约定的标签格式输出分析结果。"""
|
||||
|
||||
_IMAGE_ANALYSIS_EXAMPLE = """<products>
|
||||
<product name="大公鸡头 多功能油污净 625ml" features="红色瓶盖白色瓶身,鸡头图案Logo" position="main" image_index="0"/>
|
||||
</products>
|
||||
<colors>
|
||||
<color hex="#D32F2F" name="红色" coverage="0.4"/>
|
||||
<color hex="#FFFFFF" name="白色" coverage="0.5"/>
|
||||
</colors>
|
||||
<people has_person="false" count="0" gender="无法判断" age_range="无法判断" hair="无法判断" skin_tone="无法判断" face_shape="无法判断" outfit="无法判断" pose="无法判断" expression="无法判断"/>
|
||||
<mood>干净、实用</mood>
|
||||
<visible_text>
|
||||
<text_item text="多功能油污净" position="瓶身正面"/>
|
||||
</visible_text>
|
||||
<scene>白底棚拍产品图</scene>
|
||||
<quality resolution="高清" lighting="均匀柔和" composition="主体居中" blur="false"/>
|
||||
<key_selling_points>
|
||||
<point>针对重油污设计</point>
|
||||
<point>大容量625ml</point>
|
||||
</key_selling_points>"""
|
||||
|
||||
# ── 模板2:用户文案意图解析(LLM)──────────────────────────────────────
|
||||
_INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案可能只是几个关键词、碎碎念或者不完整的短句,你要读懂他真正想讲什么。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<intent_summary> 用用户的语言风格,一句话、30字以内概括核心意图。
|
||||
<core_messages> 下面每个核心信息点用一个 <message> 标签,属性 must_keep 为 true 或 false、confidence 为 0 到 1 的小数,标签内容写信息点。
|
||||
<personal_brands> 把用户提到的具体事实——品牌名、价格、人名、地名、时间、产品名——每条用一个 <brand> 标签,属性 category 取 brand、price、person、place、time、product 之一。这些事实必须原样引用,一个字都不能改。
|
||||
<emotion_tone> 写文案的情绪调性。
|
||||
<missing_info> 把你认为缺失、后续生成时需要合理推断的信息,每条用一个 <info> 标签;没有就输出空标签。"""
|
||||
|
||||
_INTENT_USER = """用户原始文案:{user_copy_text}
|
||||
所属行业:{industry}
|
||||
图片分析结果(供参考):
|
||||
{image_analysis}
|
||||
|
||||
请理解用户意图,按标签格式输出。"""
|
||||
|
||||
_INTENT_EXAMPLE = """<intent_summary>一款厨房去油污神器,喷一喷油污就掉</intent_summary>
|
||||
<core_messages>
|
||||
<message must_keep="true" confidence="0.97">去油污效果好,喷上等几分钟再擦</message>
|
||||
<message must_keep="false" confidence="0.7">适合厨房重油污场景</message>
|
||||
</core_messages>
|
||||
<personal_brands>
|
||||
<brand category="product">大公鸡头多功能油污净</brand>
|
||||
<brand category="price">39块钱一瓶</brand>
|
||||
</personal_brands>
|
||||
<emotion_tone>亲切、真实、带分享感</emotion_tone>
|
||||
<missing_info>
|
||||
<info>没有说明具体容量,按图片读出的625ml处理</info>
|
||||
</missing_info>"""
|
||||
|
||||
# ── 模板3:文案融合生成(LLM)──────────────────────────────────────────
|
||||
_FUSION_SYSTEM = """你负责为短视频生成营销文案。请按思维链分步完成:先定人设和目标客户,再找卖点,再搭结构,再安排情绪,最后写行动号召,不要一步到位乱写。
|
||||
|
||||
{fusion_instruction}
|
||||
|
||||
{global_constraints}
|
||||
|
||||
{negative_rules}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<title> 视频标题。
|
||||
<hook> 开头3秒钩子,5到15字。
|
||||
<body_points> 每个要点用一个 <point> 标签,属性 elaboration 是展开说明、image_index 是对应第几张图(从0开始),标签内容写要点。
|
||||
<cta> 口语化的行动号召。
|
||||
<script_segments> 每段配音用一个 <segment> 标签,属性 duration_sec 是秒数、image_index 是对应图片,标签内容写配音文案(纯口播文本,不加旁白标注、不加镜头标注、不加"主播:"之类前缀)。
|
||||
<voiceover_script> 把所有 segment 的配音文案按顺序自然拼接成一段完整的纯口播文本(无标记、无括号、无前缀),长度要适配 {duration} 秒,约 {approx_chars} 字。
|
||||
<overview_theme> 视频主题(一句话概括)。
|
||||
<scene_and_lighting> 整体场景描述+光线设定(100-200字,要具体:在哪拍、什么光线、什么色调、什么氛围)。
|
||||
<word_count> 配音总字数,只写数字。
|
||||
<estimated_duration> 预计时长秒数,只写数字。
|
||||
|
||||
用户在 personal_brands 中提到的品牌名、价格、人名、地名、时间、产品名等事实信息,必须原样出现在文案里,一个字都不能改。"""
|
||||
|
||||
_FUSION_USER = """所属行业:{industry}
|
||||
目标客户:{target_customer}
|
||||
营销目的:{marketing_purpose}
|
||||
视频时长:{duration}秒
|
||||
图片分析结果:
|
||||
{image_analysis}
|
||||
用户意图解析结果:
|
||||
{intent_result}
|
||||
|
||||
请按标签格式生成文案。"""
|
||||
|
||||
_FUSION_EXAMPLE = """<title>厨房重油污,别再用洗洁精硬擦了</title>
|
||||
<hook>这油污,我真的忍很久了</hook>
|
||||
<body_points>
|
||||
<point elaboration="喷在油污上等几分钟,一擦就干净" image_index="0">大公鸡头油污净去油快</point>
|
||||
<point elaboration="39块钱625ml,能用很久" image_index="0">39块钱一瓶,性价比高</point>
|
||||
</body_points>
|
||||
<cta>厨房油污重的,真的可以试一瓶</cta>
|
||||
<script_segments>
|
||||
<segment duration_sec="3" image_index="0">这油污我真的忍很久了,用洗洁精擦半天都没用</segment>
|
||||
<segment duration_sec="6" image_index="0">后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净</segment>
|
||||
<segment duration_sec="4" image_index="0">39块钱625ml,厨房重油污的可以试一瓶</segment>
|
||||
</script_segments>
|
||||
<voiceover_script>这油污我真的忍很久了,用洗洁精擦半天都没用。后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净。39块钱625ml,厨房重油污的可以试一瓶。</voiceover_script>
|
||||
<overview_theme>厨房油污清洁好物分享</overview_theme>
|
||||
<scene_and_lighting>简洁明亮的厨房台面场景,自然光从窗户洒入,色调温暖柔和,突出产品白色瓶身与去油污对比效果。</scene_and_lighting>
|
||||
<word_count>58</word_count>
|
||||
<estimated_duration>13</estimated_duration>"""
|
||||
|
||||
# ── 模板4:编导级分镜(LLM)────────────────────────────────────────────
|
||||
_STORYBOARD_SYSTEM = """你是短视频编导,负责把文案拆成可拍摄的分镜,为 Seedance 2.5 视频模型写编导分镜脚本。脚本将整体作为 prompt 一次性传给视频模型,必须让模型在连贯镜头流中清楚每段时间拍什么、画面如何、人物说什么。
|
||||
|
||||
工作方式:
|
||||
1. 按文案的 script_segments 顺序分配镜头。
|
||||
2. 每个镜头确定景别/角度/运镜、画面场景与对白、人物动作细节、音效/BGM、转场。
|
||||
3. 检查所有镜头时长加起来接近目标时长,误差不超过2秒。
|
||||
4. image_index 必须在已上传图片范围内,第一张主图必须用在第一个镜头。
|
||||
|
||||
{fusion_instruction}
|
||||
|
||||
{global_constraints}
|
||||
|
||||
{negative_rules}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<clips> 下面每个镜头用一个 <clip> 标签,属性 image_index 是图片序号(从0开始)、transition 取 fade/cut/zoom_in/slide_left/dissolve/wipe 之一、zoom 取 in/out/null、duration_sec 是该镜头秒数、bgm_note 是该段BGM情绪。每个 <clip> 里面包含:
|
||||
<voice_text> 该镜头配音文本(纯口播文本,不加旁白标注);
|
||||
<subtitle_text> 字幕文本,可与配音一致或更精简;
|
||||
<shot_type_angle_movement> 景别+角度+运镜(例:近景俯拍45度,缓慢推镜;中景平视,固定镜头;特写平视,快速拉镜);
|
||||
<scene_and_dialogue> 画面场景描述 + 人物口播台词(对白要自然口语化,像朋友聊天,不要硬广推销腔);
|
||||
<action_details> 人物动作、表情、物品操作细节(手怎么动、表情变化、产品怎么展示);
|
||||
<audio_bgm> 环境音+BGM提示(例:轻快流行BGM,环境嘈杂咖啡店背景音);
|
||||
<transition> 硬切/淡入淡出/叠化(最后一镜写『结束』即可);
|
||||
<reference_image_index> 参考图片索引(0-based,对应第几张产品图,无则空);
|
||||
<ken_burns> 用一个空标签,属性 start、end 写"x,y"坐标、ease 写缓动方式;不需要运镜时坐标相同。"""
|
||||
|
||||
_STORYBOARD_USER = """目标时长:{duration}秒
|
||||
上传图片数量:{image_count}张(第1张是主图/封面)
|
||||
文案内容:
|
||||
{fusion_result}
|
||||
图片分析结果:
|
||||
{image_analysis}
|
||||
|
||||
请按标签格式输出分镜。"""
|
||||
|
||||
_STORYBOARD_EXAMPLE = """<clips>
|
||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常、轻微烦躁">
|
||||
<voice_text>这油污我真的忍很久了</voice_text>
|
||||
<subtitle_text>这油污忍很久了</subtitle_text>
|
||||
<shot_type_angle_movement>近景俯拍45度,缓慢推镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>厨房台面,主妇皱眉看着灶台油污。对白:这油污我真的忍很久了</scene_and_dialogue>
|
||||
<action_details>右手拿着脏抹布,无奈摇头</action_details>
|
||||
<audio_bgm>轻快日常BGM,带一点烦躁感</audio_bgm>
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="zoom_in" zoom="in" duration_sec="6" bgm_note="轻快、出现转机">
|
||||
<voice_text>后来换了大公鸡头油污净,喷上等几分钟,一擦就干净</voice_text>
|
||||
<subtitle_text>喷上等几分钟,一擦就干净</subtitle_text>
|
||||
<shot_type_angle_movement>特写平视,固定镜头</shot_type_angle_movement>
|
||||
<scene_and_dialogue>手部特写,喷油污净在油污处。对白:后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净</scene_and_dialogue>
|
||||
<action_details>左手拿产品瓶身,右手按压喷头,等待片刻后用抹布轻擦</action_details>
|
||||
<audio_bgm>轻快转折BGM,带清爽感</audio_bgm>
|
||||
<transition>淡入淡出</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="20,20" end="80,80" ease="ease-in-out"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="fade" zoom="null" duration_sec="4" bgm_note="温暖、推荐">
|
||||
<voice_text>39块钱625ml,厨房重油污的可以试一瓶</voice_text>
|
||||
<subtitle_text>39元625ml,可以试一瓶</subtitle_text>
|
||||
<shot_type_angle_movement>中景平视,缓慢拉镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>产品正面展示,明亮背景。对白:39块钱625ml,厨房重油污的可以试一瓶</scene_and_dialogue>
|
||||
<action_details>产品置于画面中央,轻微转动展示瓶身</action_details>
|
||||
<audio_bgm>温暖收尾BGM</audio_bgm>
|
||||
<transition>结束</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="50,50" end="20,20" ease="ease-in-out"/>
|
||||
</clip>
|
||||
</clips>"""
|
||||
|
||||
# ── 模板5:文案审核(LLM)──────────────────────────────────────────────
|
||||
_REVIEW_SYSTEM = f"""你是短视频文案合规审核员,从6个维度逐条检查文案:
|
||||
1. 违规词:有没有平台禁用词、敏感词。
|
||||
2. 夸大承诺:有没有“包治百病”“100%有效”“保证赚钱”等绝对化、夸大表述。
|
||||
3. 事实一致性:有没有编造价格、数据、认证,或者用户没提到的产品特性。
|
||||
4. 用户意图保留:在 ai_polish 和 user_primary 模式下,core_messages 中 must_keep=true 的点是否都保留了。
|
||||
5. 结构完整性:标题、钩子、正文、行动号召是否齐全。
|
||||
6. 语气人设:是否符合选定的人设语气,有没有“家人们谁懂啊”“绝绝子”“宝子们”等套路词。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<passed> 整体是否通过,只写 true 或 false。
|
||||
<issues> 每个问题用一个 <issue> 标签,属性 dimension 是维度名、severity 取 error 或 warning、location 是问题所在(如 hook、body_points、cta),标签内容写问题描述;没有问题就输出空标签。
|
||||
<rewrite_suggestions> 每条具体修改建议用一个 <suggestion> 标签;没有就输出空标签。"""
|
||||
|
||||
_REVIEW_USER = """本次创作模式:{fusion_level}
|
||||
待审核文案:
|
||||
{fusion_result}
|
||||
用户意图解析(用于核对核心信息是否保留):
|
||||
{intent_result}
|
||||
|
||||
请按6个维度审核,按标签格式输出。"""
|
||||
|
||||
_REVIEW_EXAMPLE = """<passed>false</passed>
|
||||
<issues>
|
||||
<issue dimension="夸大承诺" severity="error" location="body_points">出现了“一喷100%掉光”的绝对化表述,违反广告法</issue>
|
||||
<issue dimension="用户意图保留" severity="warning" location="cta">用户强调的“39块钱”没有保留</issue>
|
||||
</issues>
|
||||
<rewrite_suggestions>
|
||||
<suggestion>把“一喷100%掉光”改为“喷上等几分钟,大部分油污能擦掉”</suggestion>
|
||||
<suggestion>在结尾补回“39块钱625ml”</suggestion>
|
||||
</rewrite_suggestions>"""
|
||||
|
||||
|
||||
# 5 套模板默认数据(seed 数据源与 loader 的兜底)
|
||||
DEFAULT_TEMPLATES: list[dict] = [
|
||||
{
|
||||
"name": "图片多模态分析",
|
||||
"prompt_type": "image_analysis",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
|
||||
"user_prompt_template": _IMAGE_ANALYSIS_USER,
|
||||
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "用户文案意图解析",
|
||||
"prompt_type": "intent_parsing",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _INTENT_SYSTEM,
|
||||
"user_prompt_template": _INTENT_USER,
|
||||
"example_output": _INTENT_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "文案融合生成",
|
||||
"prompt_type": "copy_fusion",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _FUSION_SYSTEM,
|
||||
"user_prompt_template": _FUSION_USER,
|
||||
"example_output": _FUSION_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "编导级分镜",
|
||||
"prompt_type": "storyboard",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _STORYBOARD_SYSTEM,
|
||||
"user_prompt_template": _STORYBOARD_USER,
|
||||
"example_output": _STORYBOARD_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "文案审核",
|
||||
"prompt_type": "review",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _REVIEW_SYSTEM,
|
||||
"user_prompt_template": _REVIEW_USER,
|
||||
"example_output": _REVIEW_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
]
|
||||
@@ -0,0 +1,312 @@
|
||||
"""文案审核 + 自动重写(#2040 第5套 Prompt)。
|
||||
|
||||
6 维度:违规词 / 夸大承诺 / 事实一致性 / 用户意图保留 / 结构完整性 / 语气人设。
|
||||
LLM 审核之外叠加本地规则预检(保证即使 LLM 不可用也能兜住广告法红线)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.viral_video import xml_parser as xp
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.application.viral_video.schemas import (
|
||||
FusionResult,
|
||||
IntentResult,
|
||||
ReviewIssue,
|
||||
ReviewResult,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 本地规则:绝对化/夸大词
|
||||
_EXAGGERATION_PATTERNS = [
|
||||
r"100\s*%",
|
||||
r"百分百",
|
||||
r"包治百病",
|
||||
r"保证.{0,8}(有效|赚钱|瘦|好)",
|
||||
r"绝对(有效|安全|靠谱)",
|
||||
r"全网第一",
|
||||
r"国家级",
|
||||
r"特效",
|
||||
r"立刻见效",
|
||||
r"一喷(就|全|100)",
|
||||
]
|
||||
|
||||
# 本地规则:平台违规/套路词
|
||||
_VIOLATION_PHRASES = [
|
||||
"家人们谁懂啊",
|
||||
"绝绝子",
|
||||
"宝子们",
|
||||
"yyds",
|
||||
"最(好|强|牛|便宜)", # 广告法极限词
|
||||
"第一(名|品牌)?",
|
||||
]
|
||||
|
||||
_LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
|
||||
|
||||
|
||||
class Reviewer:
|
||||
def __init__(self, client=None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
self.client = client
|
||||
|
||||
# ── 审核 ────────────────────────────────────────────────────────────
|
||||
def review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> ReviewResult:
|
||||
local = self._rule_check(fusion, intent, fusion_level)
|
||||
llm_result = self._llm_review(fusion, intent, fusion_level)
|
||||
if llm_result is None:
|
||||
return ReviewResult(
|
||||
passed=not local,
|
||||
issues=local,
|
||||
rewrite_suggestions=[],
|
||||
raw="",
|
||||
)
|
||||
# LLM 与本地规则合并去重
|
||||
issues = self._merge_issues(llm_result.issues, local)
|
||||
return ReviewResult(
|
||||
passed=llm_result.passed and not local,
|
||||
issues=issues,
|
||||
rewrite_suggestions=llm_result.rewrite_suggestions,
|
||||
raw=llm_result.raw,
|
||||
)
|
||||
|
||||
def _llm_review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> Optional[ReviewResult]:
|
||||
template = get_template("review")
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
fusion_level=fusion_level,
|
||||
fusion_result=self._fusion_text(fusion),
|
||||
intent_result=self._intent_text(intent),
|
||||
)
|
||||
raw = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.2,
|
||||
max_tokens=1024,
|
||||
)
|
||||
if not raw:
|
||||
return None
|
||||
passed = xp.text_of(raw, "passed").strip().lower()
|
||||
issues = [
|
||||
ReviewIssue(
|
||||
dimension=n["attrs"].get("dimension", "未知维度"),
|
||||
severity=n["attrs"].get("severity", "warning"),
|
||||
location=n["attrs"].get("location", ""),
|
||||
text=n["text"],
|
||||
)
|
||||
for n in xp.find_all(raw, "issue")
|
||||
if n["text"]
|
||||
]
|
||||
suggestions = [n["text"] for n in xp.find_all(raw, "suggestion") if n["text"]]
|
||||
parsed = ReviewResult(
|
||||
passed=passed == "true" and not issues,
|
||||
issues=issues,
|
||||
rewrite_suggestions=suggestions,
|
||||
raw=raw,
|
||||
)
|
||||
return parsed
|
||||
|
||||
# ── 本地规则预检 ────────────────────────────────────────────────────
|
||||
def _rule_check(self, fusion: FusionResult, intent, fusion_level: str) -> list[ReviewIssue]:
|
||||
issues: list[ReviewIssue] = []
|
||||
for location, text in self._segments(fusion):
|
||||
for pattern in _EXAGGERATION_PATTERNS:
|
||||
if re.search(pattern, text):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="夸大承诺",
|
||||
severity="error",
|
||||
location=location,
|
||||
text=f"出现夸大/绝对化表述:{self._hit(text, pattern)}",
|
||||
)
|
||||
)
|
||||
for phrase in _VIOLATION_PHRASES:
|
||||
if re.search(phrase, text, flags=re.IGNORECASE):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="违规词",
|
||||
severity="error",
|
||||
location=location,
|
||||
text=f"出现违规或套路词:{self._hit(text, phrase)}",
|
||||
)
|
||||
)
|
||||
|
||||
# 结构完整性
|
||||
if not fusion.title:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="title", text="缺少标题"))
|
||||
if not fusion.hook:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="hook", text="缺少开头钩子"))
|
||||
if not fusion.cta:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="cta", text="缺少行动号召"))
|
||||
|
||||
# 用户意图保留(must_keep)
|
||||
full_text = self._fusion_text(fusion)
|
||||
if fusion_level in {"ai_polish", "user_primary"} and intent is not None:
|
||||
for message in intent.core_messages:
|
||||
if message.must_keep:
|
||||
key = self._compact(message.text)
|
||||
if key and key[:10] not in self._compact(full_text):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="用户意图保留",
|
||||
severity="warning",
|
||||
location="script_segments",
|
||||
text=f"用户核心信息被丢失:{message.text[:30]}",
|
||||
)
|
||||
)
|
||||
for brand in intent.personal_brands:
|
||||
if brand.text and brand.text not in full_text:
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="事实一致性",
|
||||
severity="error",
|
||||
location="script_segments",
|
||||
text=f"personal_brands 事实信息未原样保留:{brand.text[:30]}",
|
||||
)
|
||||
)
|
||||
return issues
|
||||
|
||||
@staticmethod
|
||||
def _hit(text: str, pattern: str) -> str:
|
||||
match = re.search(pattern, text, flags=re.IGNORECASE)
|
||||
return match.group(0) if match else pattern
|
||||
|
||||
@staticmethod
|
||||
def _compact(text: str) -> str:
|
||||
return re.sub(r"[\s,。!?、,.!?;;::\"'“”‘’()()【】\[\]]", "", text)
|
||||
|
||||
@staticmethod
|
||||
def _merge_issues(llm_issues: list[ReviewIssue], local: list[ReviewIssue]) -> list[ReviewIssue]:
|
||||
merged = list(local)
|
||||
seen = {(i.dimension, Reviewer._compact(i.text)[:20]) for i in local}
|
||||
for issue in llm_issues:
|
||||
key = (issue.dimension, Reviewer._compact(issue.text)[:20])
|
||||
if key not in seen:
|
||||
merged.append(issue)
|
||||
seen.add(key)
|
||||
return merged
|
||||
|
||||
# ── 自动重写(1 次)─────────────────────────────────────────────────
|
||||
def rewrite(
|
||||
self,
|
||||
fusion: FusionResult,
|
||||
review: ReviewResult,
|
||||
intent: IntentResult,
|
||||
fusion_level: str,
|
||||
) -> FusionResult:
|
||||
from packages.application.viral_video.generator import CopyGenerator
|
||||
|
||||
template = get_template("copy_fusion")
|
||||
system_kwargs = {
|
||||
"fusion_instruction": (
|
||||
"【本次任务:按审核意见修正文案】只修改指出的问题,其他内容尽量原样保留;"
|
||||
"personal_brands 事实信息逐字保留;修正后按原标签格式完整输出。"
|
||||
),
|
||||
"global_constraints": "",
|
||||
"negative_rules": "",
|
||||
}
|
||||
issue_text = "\n".join(f"- [{i.dimension}/{i.location}] {i.text}" for i in review.issues)
|
||||
suggestion_text = "\n".join(f"- {s}" for s in review.rewrite_suggestions)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
industry="",
|
||||
target_customer="",
|
||||
marketing_purpose="",
|
||||
duration=fusion.estimated_duration or 15,
|
||||
image_analysis="(沿用原图片分析)",
|
||||
intent_result=self._intent_text(intent),
|
||||
)
|
||||
user = (
|
||||
f"{user}\n\n原文案:\n{self._fusion_text(fusion)}\n\n"
|
||||
f"审核发现的问题:\n{issue_text}\n\n修改建议:\n{suggestion_text or '(无)'}\n"
|
||||
"请输出修正后的完整文案。"
|
||||
)
|
||||
system = render_system_prompt(template, **system_kwargs)
|
||||
raw = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.5,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return self._rule_fix(fusion, review)
|
||||
rewritten = CopyGenerator._parse_fusion(CopyGenerator(self.client), raw)
|
||||
if not rewritten.title and not rewritten.script_segments:
|
||||
return self._rule_fix(fusion, review)
|
||||
# 保底:personal_brands 必须保留
|
||||
full = self._fusion_text(rewritten)
|
||||
for brand in intent.personal_brands:
|
||||
if brand.text and brand.text not in full:
|
||||
rewritten.cta = (rewritten.cta + brand.text).strip()
|
||||
return rewritten
|
||||
|
||||
def _rule_fix(self, fusion: FusionResult, review: ReviewResult) -> FusionResult:
|
||||
"""LLM 重写不可用时的本地兜底:删除/替换明显违规表述。"""
|
||||
replacements = [
|
||||
(re.compile(r"100\s*%|百分百"), "大部分"),
|
||||
(re.compile(r"绝对(有效|安全|靠谱)"), "比较\\1"),
|
||||
(re.compile(r"包治百病"), "适用多种情况"),
|
||||
(re.compile(r"立刻见效"), "坚持使用会有改善"),
|
||||
(re.compile(r"一喷(就|全|100%)"), "喷上等一会儿可以"),
|
||||
(re.compile(r"家人们谁懂啊|绝绝子|宝子们|yyds", re.IGNORECASE), ""),
|
||||
(re.compile(r"最好|最强|最牛|最便宜"), "很不错"),
|
||||
]
|
||||
|
||||
def fix(text: str) -> str:
|
||||
for pattern, repl in replacements:
|
||||
text = pattern.sub(repl, text)
|
||||
return text
|
||||
|
||||
fusion.title = fix(fusion.title)
|
||||
fusion.hook = fix(fusion.hook)
|
||||
fusion.cta = fix(fusion.cta)
|
||||
for point in fusion.body_points:
|
||||
point.text = fix(point.text)
|
||||
point.elaboration = fix(point.elaboration)
|
||||
for segment in fusion.script_segments:
|
||||
segment.text = fix(segment.text)
|
||||
fusion.raw = ""
|
||||
return fusion
|
||||
|
||||
# ── 文本工具 ────────────────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _segments(fusion: FusionResult):
|
||||
yield "title", fusion.title
|
||||
yield "hook", fusion.hook
|
||||
for point in fusion.body_points:
|
||||
yield "body_points", f"{point.text} {point.elaboration}"
|
||||
yield "cta", fusion.cta
|
||||
for segment in fusion.script_segments:
|
||||
yield "script_segments", segment.text
|
||||
|
||||
@staticmethod
|
||||
def _fusion_text(fusion: FusionResult) -> str:
|
||||
parts = [fusion.title, fusion.hook]
|
||||
parts += [p.text for p in fusion.body_points]
|
||||
parts += [s.text for s in fusion.script_segments]
|
||||
parts.append(fusion.cta)
|
||||
return "\n".join(p for p in parts if p)
|
||||
|
||||
@staticmethod
|
||||
def _intent_text(intent) -> str:
|
||||
if intent is None:
|
||||
return "无意图信息"
|
||||
parts = [f"意图:{intent.intent_summary}"]
|
||||
parts += [f"核心信息[must_keep={m.must_keep}]:{m.text}" for m in intent.core_messages]
|
||||
parts += [f"事实({b.category}):{b.text}" for b in intent.personal_brands]
|
||||
return "\n".join(parts)
|
||||
@@ -0,0 +1,116 @@
|
||||
"""内部 Pydantic 校验模型(不暴露给运营,运营只看 DB 里的纯文本)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class ProductItem(BaseModel):
|
||||
name: str = "无法判断"
|
||||
features: str = "无法判断"
|
||||
position: str = "secondary"
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class ColorItem(BaseModel):
|
||||
hex: str = "#000000"
|
||||
name: str = "无法判断"
|
||||
coverage: float = 0.0
|
||||
|
||||
|
||||
class TextItem(BaseModel):
|
||||
text: str = ""
|
||||
position: str = ""
|
||||
|
||||
|
||||
class ImageAnalysis(BaseModel):
|
||||
products: list[ProductItem] = Field(default_factory=list)
|
||||
colors: list[ColorItem] = Field(default_factory=list)
|
||||
has_person: bool = False
|
||||
person_count: int = 0
|
||||
people: dict[str, str] = Field(default_factory=dict)
|
||||
mood: str = ""
|
||||
visible_text: list[TextItem] = Field(default_factory=list)
|
||||
scene: str = ""
|
||||
quality: dict[str, str] = Field(default_factory=dict)
|
||||
key_selling_points: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class CoreMessage(BaseModel):
|
||||
text: str
|
||||
must_keep: bool = False
|
||||
confidence: float = 0.0
|
||||
|
||||
|
||||
class PersonalBrand(BaseModel):
|
||||
text: str
|
||||
category: str = "brand"
|
||||
|
||||
|
||||
class IntentResult(BaseModel):
|
||||
intent_summary: str = ""
|
||||
core_messages: list[CoreMessage] = Field(default_factory=list)
|
||||
personal_brands: list[PersonalBrand] = Field(default_factory=list)
|
||||
emotion_tone: str = ""
|
||||
missing_info: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class BodyPoint(BaseModel):
|
||||
text: str
|
||||
elaboration: str = ""
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class ScriptSegment(BaseModel):
|
||||
text: str
|
||||
duration_sec: float = 0
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class FusionResult(BaseModel):
|
||||
title: str = ""
|
||||
hook: str = ""
|
||||
body_points: list[BodyPoint] = Field(default_factory=list)
|
||||
cta: str = ""
|
||||
script_segments: list[ScriptSegment] = Field(default_factory=list)
|
||||
word_count: int = 0
|
||||
estimated_duration: int = 0
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class KenBurns(BaseModel):
|
||||
start: str = "0,0"
|
||||
end: str = "0,0"
|
||||
ease: str = "linear"
|
||||
|
||||
|
||||
class Clip(BaseModel):
|
||||
image_index: int = 0
|
||||
transition: str = "cut"
|
||||
zoom: str | None = None
|
||||
duration_sec: float = 0
|
||||
bgm_note: str = ""
|
||||
voice_text: str = ""
|
||||
subtitle_text: str = ""
|
||||
ken_burns: KenBurns = Field(default_factory=KenBurns)
|
||||
|
||||
|
||||
class Storyboard(BaseModel):
|
||||
clips: list[Clip] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class ReviewIssue(BaseModel):
|
||||
dimension: str
|
||||
severity: str = "warning"
|
||||
location: str = ""
|
||||
text: str = ""
|
||||
|
||||
|
||||
class ReviewResult(BaseModel):
|
||||
passed: bool = True
|
||||
issues: list[ReviewIssue] = Field(default_factory=list)
|
||||
rewrite_suggestions: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
@@ -0,0 +1,104 @@
|
||||
"""XML 标签式输出解析器(替代 json.loads)。
|
||||
|
||||
LLM 按 ``<tag attr="x">内容</tag>`` 输出,本模块解析,解析失败不抛异常,
|
||||
由调用方走规则 fallback。采用栈式扫描,嵌套标签全部可提取(内外层都保留)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from html import unescape
|
||||
from typing import Optional
|
||||
|
||||
_OPEN_RE = re.compile(r"<(?P<tag>[\w-]+)(?P<attrs>(?:\s(?:[^>]*?\S)?)?)(?P<self>/?)>")
|
||||
_CLOSE_RE = re.compile(r"</(?P<tag>[\w-]+)\s*>")
|
||||
_ATTR_RE = re.compile(r"""([\w:-]+)\s*=\s*(?:"([^"]*)"|'([^']*)')""")
|
||||
|
||||
|
||||
def parse_attributes(raw: str) -> dict[str, str]:
|
||||
"""解析标签属性字符串。"""
|
||||
attrs: dict[str, str] = {}
|
||||
for match in _ATTR_RE.finditer(raw or ""):
|
||||
value = match.group(2) if match.group(2) is not None else match.group(3)
|
||||
attrs[match.group(1)] = value
|
||||
return attrs
|
||||
|
||||
|
||||
def parse_tags(text: Optional[str]) -> list[dict]:
|
||||
"""提取全部标签(含嵌套内外层),返回 [{tag, attrs, text}],按开标签出现顺序。"""
|
||||
if not text:
|
||||
return []
|
||||
results: list[dict] = []
|
||||
stack: list[dict] = []
|
||||
token_re = re.compile(r"<[^>]+>")
|
||||
for token in token_re.finditer(text):
|
||||
raw_token = token.group(0)
|
||||
# 先按开/闭标签匹配
|
||||
open_match = _OPEN_RE.match(raw_token)
|
||||
close_match = _CLOSE_RE.match(raw_token)
|
||||
is_close_tag = raw_token.startswith("</")
|
||||
if not is_close_tag and open_match:
|
||||
is_self_close = open_match.group("self") == "/"
|
||||
node = {
|
||||
"tag": open_match.group("tag"),
|
||||
"attrs": parse_attributes(open_match.group("attrs")),
|
||||
"text": "",
|
||||
"_start": token.end(),
|
||||
}
|
||||
if is_self_close:
|
||||
node.pop("_start")
|
||||
results.append(node)
|
||||
else:
|
||||
stack.append(node)
|
||||
results.append(node)
|
||||
elif is_close_tag and close_match:
|
||||
tag = close_match.group("tag")
|
||||
# 弹出到最近同名开标签
|
||||
for idx in range(len(stack) - 1, -1, -1):
|
||||
if stack[idx]["tag"] == tag:
|
||||
node = stack[idx]
|
||||
node["text"] = unescape(text[node["_start"] : token.start()].strip())
|
||||
node.pop("_start", None)
|
||||
del stack[idx:]
|
||||
break
|
||||
# 未闭合标签:给剩余部分作为文本
|
||||
for node in stack:
|
||||
if "_start" in node:
|
||||
node["text"] = unescape(text[node["_start"] :].strip())
|
||||
node.pop("_start", None)
|
||||
return results
|
||||
|
||||
|
||||
def find_all(text: Optional[str], tag: str) -> list[dict]:
|
||||
"""提取指定标签的全部节点。"""
|
||||
return [n for n in parse_tags(text) if n["tag"] == tag]
|
||||
|
||||
|
||||
def find_first(text: Optional[str], tag: str) -> Optional[dict]:
|
||||
nodes = find_all(text, tag)
|
||||
return nodes[0] if nodes else None
|
||||
|
||||
|
||||
def text_of(text: Optional[str], tag: str, default: str = "") -> str:
|
||||
node = find_first(text, tag)
|
||||
return node["text"] if node else default
|
||||
|
||||
|
||||
def attr_bool(value: Optional[str], default: bool = False) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
return value.strip().lower() in {"true", "1", "yes", "是"}
|
||||
|
||||
|
||||
def attr_float(value: Optional[str], default: float = 0.0) -> float:
|
||||
try:
|
||||
return float(value) if value is not None and value.strip() else default
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def attr_int(value: Optional[str], default: int = 0) -> int:
|
||||
try:
|
||||
return int(float(value)) if value is not None and value.strip() else default
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
+31
-5
@@ -90,12 +90,38 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-1-6-250615"
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # #2181: lite方舟侧100%超时,默认fast_model也走pro;方舟恢复lite后通过ENV DOUBAO_FAST_MODEL切回
|
||||
)
|
||||
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_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_timeout: int = 45 # #2180: 方舟LLM高峰期响应6-8s,原30s太紧提到45s
|
||||
doubao_max_retries: int = 1 # #2180: timeout调大后一次调用就够,1次重试防偶发抖动;避免6次重试叠加到351s
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_use_lite: bool = True # #2188: lite恢复稳定,爆款视频默认lite-first提速(20-30s)
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
doubao_image_model: str = (
|
||||
"doubao-seedream-5-0-flash-260915" # #2173: 信任链 Seedream 改 flash 模型(实测 pro 46.5s→flash 13s;pro AI化图仍被Seedance拦截)
|
||||
)
|
||||
doubao_image_size: str = "1K" # #2173: 1K 已足够做 Seedance 参考图,2K 在 flash 下也 22s,1K 13s
|
||||
doubao_image_timeout: int = 60 # #2173: flash+1K 通常15s内,给60s余量
|
||||
doubao_trust_chain_enabled: bool = (
|
||||
True # #2173: 信任链总开关;若Seedream产物仍被Seedance拦截,可配 False 关闭直接t2v降级
|
||||
)
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
|
||||
dashscope_video_poll_interval: int = 10
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user