Compare commits

..

1 Commits

Author SHA1 Message Date
CI Bot 3f226cb165 style: auto-format with black + isort + ruff + prettier [skip ci-format-check] 2026-10-05 10:17:18 +00:00
60 changed files with 2018 additions and 6589 deletions
@@ -1,61 +0,0 @@
"""功能计费积分字段(爆款/对口型/智能剪辑 DB 化计费)。
给 gpu_lipsync_tasks / generation_tasks / lipsync_jobs 三张表加积分字段:
- credits_prepaid: 提交任务时预扣积分
- credits_cost: 最终结算积分
- credits_transaction_id: 预扣流水 ID
注意:feature_pricing_configs 配置表由 xiaoxia-admin 侧 migration 建立,
本仓库只读,不在此创建。
Revision ID: 096_feature_billing_fields
Revises: 095_viral_video_prompt_templates
Create Date: 2026-10-05
"""
import sqlalchemy as sa
from alembic import op
revision = "096_feature_billing_fields"
down_revision = "095_viral_video_prompt_templates"
branch_labels = None
depends_on = None
_TABLES = ("gpu_lipsync_tasks", "generation_tasks", "lipsync_jobs")
_COLUMNS = (
("credits_prepaid", sa.Float(), "0"),
("credits_cost", sa.Float(), "0"),
("credits_transaction_id", sa.String(36), ""),
)
def _table_exists(conn, name: str) -> bool:
return name in sa.inspect(conn).get_table_names()
def upgrade() -> None:
conn = op.get_bind()
for table in _TABLES:
if not _table_exists(conn, table):
continue
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
for col_name, col_type, default in _COLUMNS:
if col_name in existing:
continue
op.add_column(
table,
sa.Column(col_name, col_type, nullable=False, server_default=default),
)
def downgrade() -> None:
conn = op.get_bind()
for table in _TABLES:
if not _table_exists(conn, table):
continue
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
for col_name, _col_type, _default in _COLUMNS:
if col_name not in existing:
continue
op.drop_column(table, col_name)
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -1,222 +0,0 @@
# -*- coding: utf-8 -*-
"""099: AI 模型路由层 seed — 补齐缺失模型和能力配置.
幂等:所有 INSERT 先检查存在性。
- ai_models: 补齐 qwen3.7-plus, seedream, seedance, embedding, wan3.0 等
- ai_capability_configs: 补齐 image_generation, video_generation, embedding
- 更新已有 capability 的 lite_model_id
"""
import json
import sqlalchemy as sa
from alembic import op
revision = "099_ai_model_router_seed"
down_revision = "098_viral_video_image_analysis_v5"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
# CI 环境下 ai_models 表可能尚未创建(由 ORM 自动建表,非 migration)
# 如果表不存在则跳过 seed,由应用启动时 ORM 建表后首次访问时生效
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
if not table_check:
# ai_models 表不存在,跳过所有 seed(CI 环境)
return
# ── 1. 补齐 ai_models 缺失记录 ────────────────────────────────────────────
existing_models = {
row[0]
for row in conn.execute(
sa.text("SELECT model_key FROM ai_models WHERE deleted_at IS NULL")
).fetchall()
}
# 从已有 active 记录获取 API key(复用,不硬编码)
dashscope_key_row = conn.execute(
sa.text(
"SELECT api_key FROM ai_models WHERE provider='dashscope' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
)
).first()
dashscope_key = dashscope_key_row[0] if dashscope_key_row else ""
volcengine_key_row = conn.execute(
sa.text(
"SELECT api_key FROM ai_models WHERE provider='volcengine' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
)
).first()
volcengine_key = volcengine_key_row[0] if volcengine_key_row else ""
new_models = [
{
"model_key": "qwen3.7-plus",
"name": "通义千问3.7 Plus(VLM 兜底)",
"provider": "dashscope",
"api_key": dashscope_key,
"api_base": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"description": "阿里云百炼 Qwen3.7 Plus 多模态模型,用于 VLM 兜底分析",
},
{
"model_key": "doubao-seedream-5-0-flash-260915",
"name": "Seedream 5.0 Flash(图片生成)",
"provider": "volcengine",
"api_key": volcengine_key,
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
"description": "火山引擎 Seedream 5.0 Flash 文生图模型",
},
{
"model_key": "doubao-seedance-2-5-260628",
"name": "Seedance 2.5(视频生成)",
"provider": "volcengine",
"api_key": volcengine_key,
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
"description": "火山引擎 Seedance 2.5 图/文生视频模型",
},
{
"model_key": "doubao-embedding-vision-251215",
"name": "豆包多模态向量嵌入",
"provider": "volcengine",
"api_key": volcengine_key,
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
"description": "火山引擎豆包多模态向量嵌入模型",
},
{
"model_key": "wan3.0-video",
"name": "Wan 3.0 视频生成",
"provider": "dashscope",
"api_key": dashscope_key,
"api_base": "https://dashscope.aliyuncs.com/api/v1",
"description": "阿里云百炼 Wan 3.0 视频生成模型",
},
{
"model_key": "doubao-seed-2-1-pro-260915",
"name": "豆包 Seed 2.1 Pro(高精度推理)",
"provider": "volcengine",
"api_key": volcengine_key,
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
"description": "火山引擎豆包 Seed 2.1 Pro 深度思考+多模态",
},
]
for m in new_models:
if m["model_key"] not in existing_models:
conn.execute(
sa.text(
"""
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base, description, status, is_default, usage_today, created_at, updated_at)
VALUES (gen_random_uuid()::text, :name, :provider, :model_key, :api_key, :api_base, :description, 'active', false, 0, now(), now())
"""
),
m,
)
# ── 2. 补齐 ai_capability_configs 缺失项 ──────────────────────────────────
cap_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
if not cap_table_check:
return
existing_caps = {
row[0]
for row in conn.execute(
sa.text("SELECT capability_key FROM ai_capability_configs")
).fetchall()
}
def _get_model_id(model_key: str) -> str | None:
row = conn.execute(
sa.text(
"SELECT id FROM ai_models WHERE model_key = :key AND deleted_at IS NULL AND status = 'active' LIMIT 1"
),
{"key": model_key},
).first()
return row[0] if row else None
# image_generation
if "image_generation" not in existing_caps:
mid = _get_model_id("doubao-seedream-5-0-flash-260915")
if mid:
conn.execute(
sa.text(
"""
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 60, 1, 2, :ep, true, now(), now())
"""
),
{
"ck": "image_generation",
"cn": "图片生成(Seedream)",
"pm": mid,
"ep": json.dumps({"size": "1K"}),
},
)
# video_generation
if "video_generation" not in existing_caps:
mid = _get_model_id("doubao-seedance-2-5-260628")
fb_mid = _get_model_id("wan3.0-video")
if mid:
conn.execute(
sa.text(
"""
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, fallback_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, :fm, 600, 1, 1, :ep, true, now(), now())
"""
),
{
"ck": "video_generation",
"cn": "视频生成(Seedance/Wan)",
"pm": mid,
"fm": fb_mid,
"ep": json.dumps({}),
},
)
# embedding
if "embedding" not in existing_caps:
mid = _get_model_id("doubao-embedding-vision-251215")
if mid:
conn.execute(
sa.text(
"""
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 30, 2, 5, :ep, true, now(), now())
"""
),
{
"ck": "embedding",
"cn": "向量嵌入",
"pm": mid,
"ep": json.dumps({}),
},
)
# ── 3. 更新 image_analysis 的 lite_model_id ─────────────────────────────
lite_model_id = _get_model_id("qwen3.8-flash")
if lite_model_id:
conn.execute(
sa.text(
"UPDATE ai_capability_configs SET lite_model_id = :lite WHERE capability_key = 'image_analysis' AND lite_model_id IS NULL"
),
{"lite": lite_model_id},
)
def downgrade() -> None:
conn = op.get_bind()
# 安全检查表是否存在
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
if not table_check:
return
conn.execute(
sa.text("DELETE FROM ai_capability_configs WHERE capability_key IN ('image_generation', 'video_generation', 'embedding')")
)
conn.execute(
sa.text(
"DELETE FROM ai_models WHERE model_key IN ('qwen3.7-plus', 'doubao-seedream-5-0-flash-260915', 'doubao-seedance-2-5-260628', 'doubao-embedding-vision-251215', 'wan3.0-video', 'doubao-seed-2-1-pro-260915') AND deleted_at IS NULL"
)
)
@@ -1,107 +0,0 @@
# -*- coding: utf-8 -*-
"""100: 修正已有 capability 的模型绑定.
幂等:仅当 primary_model_id 当前绑定到旧模型 (doubao-seed-1-6) 时才更新,
避免覆盖用户在后台的自定义配置。
- 更新 5 个 LLM capability (intent_parsing, copy_fusion, storyboard, copy_review, asset_classify)
的 primary_model_id 从 doubao-seed-1-6 改为 doubao-seed-2-1-pro-260915
- 更新 image_analysis 的 primary/lite/fallback 模型绑定
"""
import sqlalchemy as sa
from alembic import op
revision = "100_fix_capability_model_bindings"
down_revision = "099_ai_model_router_seed"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
# Check tables exist
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
if not table_check:
return
config_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
if not config_table_check:
return
# Look up model IDs by model_key (not hardcoded UUIDs)
pro_model_row = conn.execute(
sa.text(
"SELECT id FROM ai_models WHERE model_key = 'doubao-seed-2-1-pro-260915' AND deleted_at IS NULL LIMIT 1"
)
).first()
if not pro_model_row:
return
pro_model_id = pro_model_row[0]
old_model_row = conn.execute(
sa.text("SELECT id FROM ai_models WHERE model_key = 'doubao-seed-1-6-250615' LIMIT 1")
).first()
old_model_id = old_model_row[0] if old_model_row else None
llm_capabilities = [
"intent_parsing",
"copy_fusion",
"storyboard",
"copy_review",
"asset_classify",
]
for cap_key in llm_capabilities:
if old_model_id:
conn.execute(
sa.text(
"UPDATE ai_capability_configs SET primary_model_id = :new_id, updated_at = NOW() "
"WHERE capability_key = :cap_key AND primary_model_id = :old_id"
),
{"new_id": pro_model_id, "old_id": old_model_id, "cap_key": cap_key},
)
# Update image_analysis
qwen38_row = conn.execute(
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.8-flash' AND deleted_at IS NULL LIMIT 1")
).first()
qwen37_row = conn.execute(
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' AND deleted_at IS NULL LIMIT 1")
).first()
if qwen38_row and qwen37_row:
qwen38_id = qwen38_row[0]
qwen37_id = qwen37_row[0]
current_ia = conn.execute(
sa.text(
"SELECT primary_model_id, lite_model_id, fallback_model_id "
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
)
).first()
if current_ia:
current_primary, current_lite, current_fallback = current_ia
updates = {}
if current_primary != qwen38_id:
updates["primary_model_id"] = qwen38_id
if current_lite != qwen38_id:
updates["lite_model_id"] = qwen38_id
if current_fallback != qwen37_id:
updates["fallback_model_id"] = qwen37_id
if updates:
set_clause = ", ".join([f"{k} = :{k}" for k in updates.keys()])
set_clause += ", updated_at = NOW()"
updates["cap_key"] = "image_analysis"
conn.execute(
sa.text(f"UPDATE ai_capability_configs SET {set_clause} WHERE capability_key = :cap_key"),
updates,
)
def downgrade() -> None:
pass
@@ -1,153 +0,0 @@
# -*- coding: utf-8 -*-
"""101: 补齐 qwen-vl-plus 视觉模型并修正 image_analysis 绑定与 max_tokens.
背景:
- qwen-vl-plus 做图片识别时返回 JSON 约 500-600 tokens,旧硬编码
max_tokens=350 导致 JSON 被截断、解析失败返回"未识别"。
- 代码侧已移除硬编码,改由 capability 的 DB 配置决定 max_tokens。
幂等:
- qwen-vl-plus 已存在则不插入;
- 仅当 image_analysis 当前 primary_model 不是 qwen-vl-plus 时才更新绑定,
避免覆盖后台手动配置。
"""
import sqlalchemy as sa
from alembic import op
revision = "101_qwen_vl_plus_and_max_tokens"
down_revision = "100_fix_capability_model_bindings"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
models_table = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
if not models_table:
return
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
if not caps_table:
return
# ── c. 补全其他 capability 的 max_tokens 默认值(幂等)──────────────────
# 放在 image_analysis 特定逻辑之前,确保任何分支 return 都不会跳过本段。
# 仅在当前值为 NULL 或过小 (<100) 时更新,不覆盖已有合理配置。
# embedding / tts / voice_clone 不走 chat 接口,无需设置。
default_max_tokens = {
"intent_parsing": 500,
"copy_fusion": 2500,
"storyboard": 4000,
"copy_review": 1000,
"asset_classify": 500,
"image_generation": 500,
"video_generation": 500,
}
for cap_key, mt in default_max_tokens.items():
conn.execute(
sa.text(
"UPDATE ai_capability_configs "
"SET max_tokens = :mt, updated_at = now() "
"WHERE capability_key = :key "
"AND (max_tokens IS NULL OR max_tokens < 100)"
),
{"mt": mt, "key": cap_key},
)
# ── a. 确保 qwen-vl-plus 模型存在 ────────────────────────────────────────
conn.execute(sa.text("""
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base,
description, status, is_default, usage_today,
created_at, updated_at)
SELECT gen_random_uuid()::text,
'通义千问VL Plus',
'dashscope',
'qwen-vl-plus',
COALESCE(
(SELECT api_key FROM ai_models
WHERE provider = 'dashscope' AND deleted_at IS NULL
AND api_key IS NOT NULL AND api_key != ''
LIMIT 1),
''
),
'https://dashscope.aliyuncs.com/compatible-mode/v1',
'阿里云视觉理解模型(图片识别/分析)',
'active', false, 0, now(), now()
WHERE NOT EXISTS (
SELECT 1 FROM ai_models
WHERE model_key = 'qwen-vl-plus' AND deleted_at IS NULL
)
"""))
qwen_vl_row = conn.execute(
sa.text(
"SELECT id FROM ai_models WHERE model_key = 'qwen-vl-plus' "
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
)
).first()
if not qwen_vl_row:
return
qwen_vl_id = qwen_vl_row[0]
qwen37_row = conn.execute(
sa.text(
"SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' "
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
)
).first()
qwen37_id = qwen37_row[0] if qwen37_row else None
# ── b. 仅当当前 primary 不是 qwen-vl-plus 时修正绑定与 max_tokens ───────
current = conn.execute(
sa.text(
"SELECT primary_model_id, lite_model_id, fallback_model_id, max_tokens "
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
)
).first()
if current is None:
# capability 不存在则创建
conn.execute(
sa.text("""
INSERT INTO ai_capability_configs
(id, capability_key, capability_name, primary_model_id,
lite_model_id, fallback_model_id, timeout_seconds,
max_retries, max_tokens, concurrency, extra_params,
is_enabled, created_at, updated_at)
VALUES (gen_random_uuid()::text, 'image_analysis', '图片分析',
:primary, :primary, :fallback, 30, 1, 1000, 2,
'{}'::jsonb, true, now(), now())
"""),
{"primary": qwen_vl_id, "fallback": qwen37_id},
)
return
current_primary = current[0]
if current_primary == qwen_vl_id:
# 已经绑定 qwen-vl-plus:视为后台/数据迁移已处理,不覆盖任何配置
return
set_parts = [
"primary_model_id = :vl_id",
"lite_model_id = :vl_id",
"max_tokens = 1000",
"updated_at = now()",
]
params: dict = {"vl_id": qwen_vl_id}
if qwen37_id is not None:
set_parts.insert(2, "fallback_model_id = :qwen37_id")
params["qwen37_id"] = qwen37_id
conn.execute(
sa.text(
"UPDATE ai_capability_configs SET " + ", ".join(set_parts) + " WHERE capability_key = 'image_analysis'"
),
params,
)
def downgrade() -> None:
pass
@@ -1,36 +0,0 @@
# -*- coding: utf-8 -*-
"""102: image_analysis max_tokens 1200 -> 1500.
v6 prompt 更长、字段更多,旧 max_tokens 容易截断 JSON。
仅在 image_analysis 当前 max_tokens < 1500 时更新(幂等,不覆盖后台已调到 >=1500 的配置)。
"""
import sqlalchemy as sa
from alembic import op
revision = "102_image_analysis_max_tokens_1500"
down_revision = "101_qwen_vl_plus_and_max_tokens"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
if not caps_table:
return
conn.execute(
sa.text(
"UPDATE ai_capability_configs "
"SET max_tokens = 1500, updated_at = now() "
"WHERE capability_key = 'image_analysis' "
"AND (max_tokens IS NULL OR max_tokens < 1500)"
)
)
def downgrade() -> None:
pass
@@ -1,194 +0,0 @@
# -*- coding: utf-8 -*-
"""image_analysis v7 prompt + max_tokens 3000 + max_retries 3
Revision ID: 103_v7_prompt_and_tokens_3000
Revises: 102_image_analysis_max_tokens_1500
Create Date: 2026-10-07
变更:
1. 插入v7精简prompt(~1KB,v6 ~4.5KB,删除few-shot/冗长规则,减少输出token占用),设为active
2. v6停用(is_active=False),保留历史
3. image_analysis capability: max_tokens 1500→3000,max_retries 1→3
ai_capability_configs 由应用 create_all 创建,全新 alembic-only 库可能不存在,
故第3步做 to_regclass 守卫(同 102)。
"""
from sqlalchemy import text
from alembic import op
revision = "103_v7_prompt_and_tokens_3000"
down_revision = "102_image_analysis_max_tokens_1500"
branch_labels = None
depends_on = None
V7_SYSTEM = """# 角色
你是一位专业的图片分析师,擅长准确识别图片中的场景、人物、物体、文字、氛围。
# 任务
对用户上传的图片逐张分析,描述你看到的内容,输出JSON格式。
## 技能
### 技能1:判断图片类型
判断图片属于哪种类型,type字段填对应的英文值:
- 商品图(product):单个或多个商品、产品包装
- 门店场景图(store):店铺内部、门头招牌、货架陈列
- 人物图(person):人物形象、穿搭造型、肖像照片
- 风景图(scene):风景、动物、美食、街景
- 其他(other):以上都不是
### 技能2:描述通用信息
不管什么图都要描述:
- type:图片类型,填product/store/person/scene/other其中一个
- scene:一句话描述场景,例如"理疗养生店内部,摆着多张理疗床和产品货架"
- mood:整体氛围,2-4个词,例如"整洁专业"、"热闹温馨"
- colors:主要颜色,最多5个,写具体颜色名(亮红色/米白色/深蓝色,不写笼统的红色蓝色)
- visible_text:图片里看到的文字,说明什么字、在什么位置,最多5条;没看到就空数组
- lighting:光线情况,例如"明亮柔光"、"自然光"、"室内暖黄灯"
- composition:怎么拍的,例如"居中特写"、"中景平视"、"俯拍"
- has_person:有没有人,true或false
### 技能3:描述门店场景
如果是门店场景图(type="store"),还要描述:
- store_type:什么类型的店,例如"养生馆"、"便利店"、"餐饮店"、"母婴店"
- brand_signage:招牌上写了什么字、有什么品牌标识
- visual_elements:看到哪些显眼的东西(招牌样式、灯光、货架、商品陈列、海报、收银台等),最多8个
- product_categories:看到哪些品类的商品,例如"饮料零食"、"养生产品"
- promotion_elements:有没有促销活动(打折海报、满减吊旗等),没有就空数组
- atmosphere:店内什么氛围,例如"亲民生活化"、"老字号专业感"
- cleanliness:店内干净程度,例如"干净整洁"、"货架整齐"
- 看到顾客或店员要描述他们在做什么,has_person填true
### 技能4:描述商品
如果是商品图(type="product"),逐个商品描述:
- product_name:商品名称,尽量具体,例如"OMO奥妙除菌除螨洗衣液";看不出来填null
- brand:什么牌子,看不出来填null
- category:类目,从以下选一个:服饰鞋包/美妆/数码/食品/家居清洁/母婴/配饰/其他
- package_type:什么包装,例如"瓶装"、"盒装"、"罐装"、"袋装"、"多瓶装"
- package_color:包装主要颜色,写具体色(亮红色不写红色)
- body_shape:瓶身或包装形状,例如"圆润胖瓶"、"竖款带把手瓶身"
- label_design:标签设计,例如"红色标签印白色品牌logo"
- key_text_on_package:包装上最显眼的文字(品牌名、功能词、卖点词),最多5个
- product_features:包装特征,3-6个短语,包含颜色、瓶盖、形状、标签图案
- key_selling_points:核心卖点,1-3个短语
### 技能5:描述人物
如果是人物图(type="person"),描述:
- person_count:几个人
- gender:性别(男/女/无法判断)
- age_range:年龄段(儿童/青少年/青年/中年/老年/无法判断)
- outfit_style:穿搭风格,例如"休闲日常"、"通勤商务"、"街头潮流"
- upper_wear:上装(颜色+款式+材质),穿裙装不填
- lower_wear:下装(颜色+款式+版型),穿裙装不填
- dress_wear:裙装描述,穿上下装不填
- outerwear:外套
- shoes:鞋子
- bag:包袋,没有填null
- accessories:配饰(眼镜/帽子/项链/耳环/手表/手链/围巾/腰带等),没有填空数组
- hairstyle:发型
- makeup:妆容,男生或看不出填null
- expression:表情,例如"微笑看镜头"、"冷酷无表情"
- pose:姿势动作,例如"身直立正对镜头"、"单手撩发"
- body_type:身材,例如"纤细苗条"、"高挑身材"、"丰满匀称"
- portrait_prompt:80-150字详细描述人物形象(后面用来AI生成肖像图),要写清年龄段、穿搭完整细节、发型发色、妆容、表情、姿势、场景、光线、风格感觉,语言要有画面感
### 技能6:描述风景
如果是风景图(type="scene"),描述:
- scene_type:什么场景,例如"自然风景"、"城市街景"、"动物"、"美食"
- main_subject:画面主体是什么
- key_elements:关键元素,最多8个
- environment_objects:周围环境物体,最多8个
- atmosphere:整体氛围,例如"秋日慵懒氛围感"、"清新自然氧气感"
- 有人物就描述人物特征
## 限制
- 只输出JSON,不要任何解释文字,不要markdown代码块包裹,不要写"好的""以下是分析结果"这种废话
- 颜色写具体色调(亮红色/米白色/深蓝色/翠绿色),不写笼统词汇
- 瓶身、包装、招牌上的文字尽量识别出来(品牌名、功能词、卖点词)
- 多个商品、多个人物分开描述,不要合并
- 看不出来、不确定的字段填null或空数组,布尔值填true/false,绝对不要瞎编
- 确保JSON格式合法,所有大括号、中括号、引号正确闭合
- 数组字段控制数量:colors最多5个,visible_text最多5条,visual_elements最多8个,accessories最多10个"""
V7_USER = "请分析这张图片,按系统消息的JSON结构输出。"
def _capability_table_exists(bind) -> bool:
return bool(bind.execute(text("SELECT to_regclass('public.ai_capability_configs')")).scalar())
def upgrade() -> None:
bind = op.get_bind()
# 1. 停用旧的active image_analysis prompt(含v6)
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
"WHERE prompt_type = 'image_analysis' AND is_active = TRUE"
)
)
# 2. 幂等插入v7(存在则更新并重新激活)
existing = bind.execute(
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
).fetchone()
if existing:
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
"system_prompt = :sys, user_prompt_template = :usr, "
"name = 'v7 精简结构化分析', updated_at = NOW() "
"WHERE prompt_type = 'image_analysis' AND version = 7"
),
{"sys": V7_SYSTEM, "usr": V7_USER},
)
else:
bind.execute(
text(
"INSERT INTO viral_video_prompt_templates "
"(prompt_type, version, name, system_prompt, user_prompt_template, "
"is_active, created_at, updated_at) "
"VALUES ('image_analysis', 7, 'v7 精简结构化分析', "
":sys, :usr, TRUE, NOW(), NOW())"
),
{"sys": V7_SYSTEM, "usr": V7_USER},
)
# 3. capability max_tokens=3000、max_retries=3(表不存在则跳过)
if _capability_table_exists(bind):
bind.execute(
text(
"UPDATE ai_capability_configs SET max_tokens = 3000, "
"updated_at = NOW() "
"WHERE capability_key = 'image_analysis' AND "
"(max_tokens IS NULL OR max_tokens < 3000)"
)
)
bind.execute(
text(
"UPDATE ai_capability_configs SET max_retries = 3, updated_at = NOW() "
"WHERE capability_key = 'image_analysis' AND "
"(max_retries IS NULL OR max_retries < 3)"
)
)
def downgrade() -> None:
bind = op.get_bind()
# 删除v7
bind.execute(
text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
)
# 恢复v6为active
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
"WHERE prompt_type = 'image_analysis' AND version = 6"
)
)
# tokens/retries回退
if _capability_table_exists(bind):
bind.execute(
text(
"UPDATE ai_capability_configs SET max_tokens = 1500, max_retries = 1, "
"updated_at = NOW() WHERE capability_key = 'image_analysis'"
)
)
@@ -1,242 +0,0 @@
# -*- coding: utf-8 -*-
"""image_analysis v8 prompt + storyboard v3 prompt - 用户端展示格式 markdown 控制
Revision ID: 104_v8_display_markdown
Revises: 103_v7_prompt_and_tokens_3000
Create Date: 2026-10-07
变更:
1. image_analysis v8: 在 v7 基础上 system_prompt 末尾追加「## 用户端展示格式」章节,
要求 VLM 在每张图的 JSON 里输出 summary_markdown 字段(markdown 格式的图片描述),
v8 设 is_active=true,v7 设 is_active=false。
2. storyboard v3: 在 v2 基础上 system_prompt 追加要求 LLM 在 copy_result 中
输出 copy_display_markdown 字段(markdown 格式的完整文案展示),
v3 设 is_active=true,v2 设 is_active=false。
"""
from sqlalchemy import text
from alembic import op
revision = "104_v8_display_markdown"
down_revision = "103_v7_prompt_and_tokens_3000"
branch_labels = None
depends_on = None
# ── v8 追加的 system prompt 内容 ──────────────────────────────────────
V8_SYSTEM_APPEND = """
## 用户端展示格式
对于每张分析的图片,在 JSON 中额外输出一个 **summary_markdown** 字段,用 markdown 格式写出给用户看的图片描述。
格式要求(根据图片类型自适应):
**商品图(type=product)**示例:
### 商品名称
**品牌**:品牌名 | **类目**:服饰鞋包/美妆/数码/...
**核心特征**
- 特征1:描述
- 特征2:描述
**外观**:颜色+材质+设计描述
**包装**:包装类型描述
**文字信息**:包装上看到的文字
**门店场景图(type=store)**示例:
### 门店名称/类型
**类型**:奶茶店/便利店/养生馆/...
**品牌标识**:招牌文字描述
**环境氛围**:店内整体感觉
**陈列亮点**
- 亮点1
- 亮点2
**氛围**:亲民/专业/时尚/...
**人物图(type=person)**示例:
### 人物描述
**形象**:年龄段 + 风格
**穿搭**
- 上装:颜色+款式
- 下装:颜色+款式
- 配饰:...
**气质**:表情+姿势+整体感觉
**风景/场景图(type=scene)**示例:
### 场景名称
**类型**:自然风景/城市街景/动物/美食
**主体**:画面主要元素
**氛围**:整体感觉描述
要求:
- 内容真实具体,从实际图片分析得出
- 用 markdown 语法:**加粗**、列表、标题
- 控制在 100-200 字
- 不要编造图片中没有的信息
"""
# ── storyboard v3 追加的 system prompt 内容 ──────────────────────────
V3_STORYBOARD_APPEND = """
## 用户端展示格式
在输出分镜脚本的同时,在顶层输出一个 **copy_display_markdown** 字段(用 XML 标签 <copy_display_markdown> 包裹),用 markdown 格式写出完整文案展示。
格式示例:
# 标题/主题
## 整体概要
一句话描述视频内容
## 分镜预览
### 镜头1(0-3秒)
**景别**:近景俯拍,缓慢推镜
**画面**:场景描述
**台词**:口播文本
**动作**:人物动作描述
### 镜头2(3-9秒)
...
## 完整口播
完整口播文案文本
要求:
- 把所有分镜按时间顺序整理成易读的格式
- 用 markdown 语法组织,**加粗**标签、##二级标题、列表等
- 控制在 300-500 字
- 让用户一眼看懂视频会拍成什么样
"""
def upgrade() -> None:
bind = op.get_bind()
# ── 1. image_analysis v8 ──────────────────────────────────────────
# 停用所有 active image_analysis prompt
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
"WHERE prompt_type = 'image_analysis' AND is_active = TRUE"
)
)
# 读取 v7 的 prompt 内容作为基础
v7_row = bind.execute(
text(
"SELECT system_prompt, user_prompt_template, COALESCE(example_output, '') "
"FROM viral_video_prompt_templates "
"WHERE prompt_type = 'image_analysis' "
"ORDER BY version DESC LIMIT 1"
)
).fetchone()
if v7_row:
v7_system = v7_row[0] or ""
v8_system = v7_system + V8_SYSTEM_APPEND
v8_user = v7_row[1] or "{image_url}"
v8_example = v7_row[2] or ""
# 幂等:已有 v8 则更新,否则插入
existing_v8 = bind.execute(
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 8")
).fetchone()
if existing_v8:
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
"system_prompt = :sys, user_prompt_template = :usr, "
"example_output = :ex, name = 'v8 用户端展示格式', "
"updated_at = NOW() "
"WHERE prompt_type = 'image_analysis' AND version = 8"
),
{"sys": v8_system, "usr": v8_user, "ex": v8_example},
)
else:
bind.execute(
text(
"INSERT INTO viral_video_prompt_templates "
"(prompt_type, version, name, system_prompt, user_prompt_template, "
"example_output, is_active, created_at, updated_at) "
"VALUES ('image_analysis', 8, 'v8 用户端展示格式', "
":sys, :usr, :ex, TRUE, NOW(), NOW())"
),
{"sys": v8_system, "usr": v8_user, "ex": v8_example},
)
# ── 2. storyboard v3 ─────────────────────────────────────────────
# 停用所有 active storyboard prompt
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
"WHERE prompt_type = 'storyboard' AND is_active = TRUE"
)
)
# 读取当前 storyboard prompt
sb_row = bind.execute(
text(
"SELECT system_prompt, user_prompt_template, COALESCE(example_output, '') "
"FROM viral_video_prompt_templates "
"WHERE prompt_type = 'storyboard' "
"ORDER BY version DESC LIMIT 1"
)
).fetchone()
if sb_row:
sb_system = sb_row[0] or ""
v3_system = sb_system + V3_STORYBOARD_APPEND
v3_user = sb_row[1] or ""
v3_example = sb_row[2] or ""
existing_v3 = bind.execute(
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'storyboard' AND version = 3")
).fetchone()
if existing_v3:
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
"system_prompt = :sys, user_prompt_template = :usr, "
"example_output = :ex, name = 'v3 用户端展示格式', "
"updated_at = NOW() "
"WHERE prompt_type = 'storyboard' AND version = 3"
),
{"sys": v3_system, "usr": v3_user, "ex": v3_example},
)
else:
bind.execute(
text(
"INSERT INTO viral_video_prompt_templates "
"(prompt_type, version, name, system_prompt, user_prompt_template, "
"example_output, is_active, created_at, updated_at) "
"VALUES ('storyboard', 3, 'v3 用户端展示格式', "
":sys, :usr, :ex, TRUE, NOW(), NOW())"
),
{"sys": v3_system, "usr": v3_user, "ex": v3_example},
)
def downgrade() -> None:
bind = op.get_bind()
# 删除 v8
bind.execute(
text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 8")
)
# 恢复 v7 active
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE, updated_at = NOW() "
"WHERE prompt_type = 'image_analysis' AND version = 7"
)
)
# 删除 v3
bind.execute(text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'storyboard' AND version = 3"))
# 恢复 storyboard v2 active
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE, updated_at = NOW() "
"WHERE prompt_type = 'storyboard' AND version = 2"
)
)
-116
View File
@@ -1,116 +0,0 @@
# -*- coding: utf-8 -*-
"""image_analysis v8 + storyboard v3 叙述优先重写版(架构大简化)
Revision ID: 105_narration_first
Revises: 104_v8_display_markdown
Create Date: 2026-10-07
变更:
1. image_analysis v8:用「叙述优先」版整体替换 104 的 append 版——VLM 主交付物是
自然叙述 summary_markdown,结构化字段仅保留 type/name/brand/has_person,
顶层 products 改名 images;v8 active,其余 image_analysis 全部 deactivate。
2. storyboard v3:整体替换为风格重写版(口播口语化、画面有画面感、
copy_display_markdown 流畅叙述);v3 active,其余 storyboard deactivate。
3. intent_parsing 类型模板全部 deactivate(意图解析步骤已删除)。
模板内容直接取自 packages.application.viral_video.prompts.DEFAULT_TEMPLATES,
保证代码默认值与 DB seed 完全一致。
"""
from sqlalchemy import text
from alembic import op
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
revision = "105_narration_first"
down_revision = "104_v8_display_markdown"
branch_labels = None
depends_on = None
def _tpl(prompt_type: str, version: int) -> dict:
for t in DEFAULT_TEMPLATES:
if t["prompt_type"] == prompt_type and t["version"] == version:
return t
raise RuntimeError("default template missing: %s v%s" % (prompt_type, version))
def _upsert(bind, t: dict) -> None:
existing = bind.execute(
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = :pt AND version = :ver"),
{"pt": t["prompt_type"], "ver": t["version"]},
).fetchone()
params = {
"pt": t["prompt_type"],
"ver": t["version"],
"name": t["name"],
"sys": t["system_prompt"],
"usr": t["user_prompt_template"],
"ex": t.get("example_output", "") or "",
}
if existing:
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET name = :name, "
"system_prompt = :sys, user_prompt_template = :usr, "
"example_output = :ex, is_active = TRUE, updated_at = NOW() "
"WHERE prompt_type = :pt AND version = :ver"
),
params,
)
else:
bind.execute(
text(
"INSERT INTO viral_video_prompt_templates "
"(prompt_type, version, name, system_prompt, user_prompt_template, "
"example_output, is_active, created_at, updated_at) "
"VALUES (:pt, :ver, :name, :sys, :usr, :ex, TRUE, NOW(), NOW())"
),
params,
)
def upgrade() -> None:
bind = op.get_bind()
# 1. image_analysis:停用全部后写入叙述优先 v8
bind.execute(
text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'image_analysis'")
)
_upsert(bind, _tpl("image_analysis", 8))
# 2. storyboard:停用全部后写入重写版 v3
bind.execute(text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'storyboard'"))
_upsert(bind, _tpl("storyboard", 3))
# 3. intent_parsing 已废弃:全部停用
bind.execute(
text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'intent_parsing'")
)
# 4. review 模板确保 active
bind.execute(text("UPDATE viral_video_prompt_templates SET is_active = TRUE " "WHERE prompt_type = 'review'"))
def downgrade() -> None:
bind = op.get_bind()
# 恢复 104 的 v8/v3 无法重建(内容已替换),仅把版本 active 状态回退:
# 停用新版,尝试恢复 v7 / v2
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
"WHERE prompt_type IN ('image_analysis','storyboard') "
"AND version IN (8, 3)"
)
)
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
"WHERE prompt_type = 'image_analysis' AND version = 7"
)
)
bind.execute(
text(
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
"WHERE prompt_type = 'storyboard' AND version = 2"
)
)
+1 -48
View File
@@ -44,7 +44,6 @@ from packages.application import (
GetGenerationTaskUseCase,
ListGeneratedVideosByTaskUseCase,
)
from packages.domain import feature_pricing_service
from packages.domain.smart_match import smart_select_assets
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
@@ -164,6 +163,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
return matched or None
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -700,17 +700,6 @@ def create_generation_task(
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
effective_strategy_id = "one_take"
# ── smart_edit 计费预扣(全局 points 开关 + 功能开关均开才扣) ──
# 首期固定价:dynamic_cost=0,price=(0+fixed_cost)×multiplier,price_cap 封顶。
# 预览任务不扣费;按任务条数扣费,任一任务预扣失败(余额不足)整体拒绝。
smart_edit_charge = 0.0
charged_task_count = 0
if not request.is_preview and feature_pricing_service.is_feature_enabled("smart_edit"):
unit_credits, _bd = feature_pricing_service.calculate_price("smart_edit", 0.0)
if unit_credits > 0:
smart_edit_charge = round(unit_credits * count, 2)
charged_task_count = count
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。
# - 变体 0:clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant
# 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用);
@@ -941,42 +930,6 @@ def create_generation_task(
)
# 变体序号写入 extra_meta(响应/排查时可辨识)
task.extra_meta["variant_index"] = task_index
# smart_edit 逐条预扣(首期固定价,credits_cost=prepaid,不做结算)
task_txn_id = ""
if charged_task_count > 0:
from packages.domain.points_service import PointsService
unit_credits = round(smart_edit_charge / count, 2)
res = PointsService().deduct_points(
user_id=user_id,
amount=unit_credits,
source="smart_edit",
db=db,
description="智能剪辑生成预扣",
ref_id=task.id,
)
if not res.get("success"):
# 余额不足:退还本次请求已扣积分后整体拒绝
already_charged = round(unit_credits * task_index, 2)
if already_charged > 0:
PointsService().refund_points(
user_id=user_id,
amount=already_charged,
source="smart_edit",
db=db,
ref_id=task.id,
description="智能剪辑批量提交失败退回",
)
raise HTTPException(
status_code=402,
detail=(f"积分不足:智能剪辑每条需 {unit_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"),
)
task_txn_id = str(res.get("transaction_id") or "")
task.credits_prepaid = unit_credits
task.credits_cost = unit_credits
task.credits_transaction_id = task_txn_id
generation_task_repository.update(task)
try:
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
+3 -13
View File
@@ -280,18 +280,11 @@ def generate_copy(
raise HTTPException(status_code=404, detail="任务不存在")
if job.user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="无权操作此任务")
# 允许首次进入(IMAGE_ANALYZED/PENDING)、失败重试(FAILED)、文案重新生成(COPY_GENERATED/COMPLETED)
if job.status not in (
ViralVideoStatus.IMAGE_ANALYZED,
ViralVideoStatus.PENDING,
ViralVideoStatus.FAILED,
ViralVideoStatus.COPY_GENERATED,
ViralVideoStatus.COMPLETED,
):
if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED):
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案")
# 失败重试 / 重新生成:retry_count 自增
if job.status in (ViralVideoStatus.FAILED, ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
# 允许失败任务重试:重置
if job.status == ViralVideoStatus.FAILED:
job.retry_count += 1
job.error_msg = ""
@@ -348,9 +341,6 @@ def confirm_copy(
raise HTTPException(status_code=403, detail="无权操作此任务")
if job.status != ViralVideoStatus.COPY_GENERATED:
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
# #2218: 额外校验 copy_result 完整性,防止孤儿/脏数据进入渲染
if not isinstance(job.copy_result, dict) or not job.copy_result:
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
# 积分预扣(已扣过/重试任务跳过)
from app.config import settings as _settings
@@ -228,8 +228,6 @@ class GpuLipsyncService:
lipsync_job_id: str = "",
user_id: str = "",
project_id: str = "",
credits_prepaid: float = 0.0,
credits_transaction_id: str = "",
) -> GpuLipsyncTaskModel:
task_id = str(uuid.uuid4())
now = datetime.now(UTC)
@@ -242,8 +240,6 @@ class GpuLipsyncService:
audio_url=audio_url,
status="pending",
attempt=0,
credits_prepaid=float(credits_prepaid or 0.0),
credits_transaction_id=str(credits_transaction_id or ""),
created_at=now,
updated_at=now,
)
+1 -198
View File
@@ -38,7 +38,6 @@ from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
from packages.application.cosyvoice_service import CosyVoiceError
from packages.config import get_api_settings
from packages.domain import feature_pricing_service
from packages.domain.sentence_timings import (
compute_sentence_timings,
probe_audio_duration,
@@ -223,47 +222,7 @@ class LipsyncService:
if timings:
job.sentence_timings = timings
# 4. 检查是否走 Ditto(蚂蚁数字人,#2076):开关 + 配置完整
use_ditto = False
if self.settings.use_ditto_lipsync:
try:
from packages.application.ditto_service import get_ditto_client
ditto = get_ditto_client()
if ditto.is_configured:
use_ditto = True
logger.info("[lipsync] 优先走 Ditto 蚂蚁数字人: job_id=%s", job.id)
else:
logger.info(
"[lipsync] Ditto 开关已开但配置不完整(base_url=%s, template=%s),继续判断 GPU: job_id=%s",
bool(ditto.base_url),
bool(ditto.default_video_url),
job.id,
)
except Exception as exc:
logger.warning("[lipsync] Ditto 初始化失败,继续判断 GPU: job_id=%s err=%s", job.id, exc)
if use_ditto:
try:
# Ditto 使用预置人物模板视频,不用用户上传的 video_url;
# 但保留用户 video_url 以便失败回退到 GPU/MediaKit。
job.status = "processing"
job.mediakit_task_id = "ditto:submitted"
job.updated_at = datetime.now(UTC)
self.db.commit()
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
lipsync_ditto_process_async.apply_async(args=(job.id, job.user_id))
logger.info("[lipsync] Ditto 任务已异步派发: job_id=%s", job.id)
return
except Exception as exc:
logger.warning("[lipsync] Ditto 派发失败,回退 GPU/MediaKit: job_id=%s err=%s", job.id, exc)
try:
self.db.rollback()
except Exception:
pass
# 5. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
# 4. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
use_gpu = False
if self.settings.use_gpu_lipsync:
try:
@@ -409,8 +368,6 @@ class LipsyncService:
lipsync_job_id=job.id,
user_id=job.user_id,
project_id=job.project_id,
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
credits_transaction_id=str(getattr(job, "credits_transaction_id", "") or ""),
)
logger.info(
"[lipsync] 已创建 GPU 任务(异步): job_id=%s gpu_task=%s",
@@ -458,121 +415,6 @@ class LipsyncService:
job.output_duration,
)
# ── lip_sync 计费辅助 ────────────────────────────────────────────────
@staticmethod
def _estimate_duration(
*,
audio_duration: Optional[float] = None,
sentence_timings: Optional[list] = None,
script_text: str = "",
) -> float:
"""预估音频/成片秒数。
优先级:audio_duration(预合成前端已 ffprobe)> timings 末句 end_time >
脚本字数 / 5 字每秒 > 默认 10 秒。
"""
if audio_duration and float(audio_duration) > 0:
return float(audio_duration)
if sentence_timings:
max_end = 0.0
for item in sentence_timings:
if isinstance(item, dict):
end = item.get("end_time") or item.get("end") or 0.0
else:
end = 0.0
try:
max_end = max(max_end, float(end))
except (TypeError, ValueError):
continue
if max_end > 0:
return max_end
text = (script_text or "").strip()
if text:
return max(1.0, len(text) / 5.0)
return 10.0
def _settle_lip_sync(self, job: LipsyncJobModel, actual_duration: float) -> None:
"""按实际时长结算(首期只退不补:final < prepaid 退差额,> 不补)。
幂等:credits_cost 已 > 0 说明结算过,直接跳过。
结算失败不阻塞业务(结果已产出),仅记录日志。
"""
try:
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
if prepaid <= 0:
return
if float(getattr(job, "credits_cost", 0) or 0) > 0:
return
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
duration = float(actual_duration or 0.0)
if duration <= 0:
duration = self._estimate_duration(
sentence_timings=job.sentence_timings,
script_text=job.script_text,
)
final_price, _bd = feature_pricing_service.calculate_price("lip_sync", duration * unit_cost)
final_price = round(float(final_price), 2)
job.credits_cost = final_price
if final_price < prepaid - 0.009:
refund = round(prepaid - final_price, 2)
from packages.domain.points_service import PointsService
res = PointsService().refund_points(
user_id=job.user_id,
amount=refund,
source="lip_sync",
db=self.db,
ref_id=str(job.credits_transaction_id or job.id),
description="对口型结算退费",
)
if not res.get("success"):
logger.warning(
"[lip_sync] 结算退费失败 job_id=%s refund=%.2f(不阻塞)",
job.id,
refund,
)
# final > prepaid:首期只退不补,不补扣
self.db.commit()
except Exception: # noqa: BLE001
logger.exception("[lip_sync] 结算异常 job_id=%s(不阻塞结果)", job.id)
try:
self.db.rollback()
except Exception: # noqa: BLE001
pass
def _refund_lip_sync(self, job: LipsyncJobModel) -> None:
"""任务失败/取消时全额退还预扣积分(credits_cost 已结算则退实际未消耗部分)。"""
try:
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
if prepaid <= 0:
return
txn_id = str(getattr(job, "credits_transaction_id", "") or "")
cost = float(getattr(job, "credits_cost", 0) or 0)
refund = round(prepaid - cost, 2) if cost > 0 else round(prepaid, 2)
if refund <= 0:
return
from packages.domain.points_service import PointsService
res = PointsService().refund_points(
user_id=job.user_id,
amount=refund,
source="lip_sync",
db=self.db,
ref_id=txn_id or job.id,
description="对口型失败/取消退款",
)
if res.get("success"):
job.credits_cost = prepaid # 标记已全额退回,防重复退
self.db.commit()
except Exception: # noqa: BLE001
logger.exception("[lip_sync] 退款异常 job_id=%s", job.id)
try:
self.db.rollback()
except Exception: # noqa: BLE001
pass
# ── 创建任务 ──────────────────────────────────────────────────────────
def create_job(
@@ -624,35 +466,6 @@ class LipsyncService:
if not isinstance(sentence_timings, list) or len(sentence_timings) == 0:
raise MediaKitError("预合成模式 sentence_timings 不能为空", code="InvalidInput")
# 0.5 lip_sync 计费预扣(全局 points 开关 + 功能开关均开才扣)
prepaid_credits = 0.0
prepaid_txn_id = ""
if feature_pricing_service.is_feature_enabled("lip_sync"):
est_duration = self._estimate_duration(
audio_duration=audio_duration,
sentence_timings=sentence_timings,
script_text=script_text,
)
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
dynamic_cost = est_duration * unit_cost
prepaid_credits, _bd = feature_pricing_service.calculate_price("lip_sync", dynamic_cost)
if prepaid_credits > 0:
from packages.domain.points_service import PointsService
res = PointsService().deduct_points(
user_id=user_id,
amount=prepaid_credits,
source="lip_sync",
db=self.db,
description="对口型生成预扣",
)
if not res.get("success"):
raise ValueError(
f"积分不足:本次对口型需 {prepaid_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"
)
prepaid_txn_id = str(res.get("transaction_id") or "")
# 1. 创建数据库记录
job_id = str(uuid.uuid4())
job = LipsyncJobModel(
@@ -669,8 +482,6 @@ class LipsyncService:
emotion=emotion or "",
# 音频直传(含预合成)直接进入 pending(后续同步改为 submitted);TTS 模式进入 tts_processing
status="tts_processing" if is_tts_mode else "pending",
credits_prepaid=prepaid_credits,
credits_transaction_id=prepaid_txn_id,
)
self.db.add(job)
self.db.flush()
@@ -866,8 +677,6 @@ class LipsyncService:
job.completed_at = _now
job.updated_at = _now
self.db.commit()
# lip_sync 超时全额退款
self._refund_lip_sync(job)
return job
# 未提交的任务不轮询
@@ -893,8 +702,6 @@ class LipsyncService:
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
self.db.commit()
# lip_sync 结算(只退不补)
self._settle_lip_sync(job, float(job.output_duration or 0.0))
# 异步转存自家 OSS
try:
from app.tasks.lipsync_tts import persist_output_video_task
@@ -912,8 +719,6 @@ class LipsyncService:
job.error_message = error.get("message", "任务执行失败")
job.error_code = error.get("code", "TaskFailed")
job.completed_at = datetime.now(UTC)
# lip_sync 失败全额退款(先退款再统一 commit)
self._refund_lip_sync(job)
else:
# 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted
if isinstance(mk_status, str) and mk_status:
@@ -1007,8 +812,6 @@ class LipsyncService:
job.status = "cancelled"
job.updated_at = datetime.now(UTC)
self.db.commit()
# lip_sync 取消全额退款
self._refund_lip_sync(job)
self.db.refresh(job)
return job
-318
View File
@@ -1,318 +0,0 @@
"""Ditto 蚂蚁数字人口型异步任务 — #2076.
把 Ditto 同步 HTTP 调用(30-120s)从 API 请求移到 Celery 后台执行:
1. 加载 LipsyncJob
2. 调 DittoClient.generate_and_persist(video_url=默认模板, audio_url=job.audio_url, script=job.script_text)
3. 成功:标记 completed,写入 output_video_url(Ditto 输出自带音频,无需二次混流/超分)
4. 失败:回退 GPU MuseTalk → 再失败回退 MediaKit
注意:
- 保留 MuseTalk 代码不动;Ditto 优先,失败按原链路兜底
- Ditto 使用预置的人物模板视频(settings.ditto_default_video_url),不用用户上传的 video_url
- 不传 GFPGAN 超分,不需要 ffmpeg 音视频混流
"""
from __future__ import annotations
import logging
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Optional
from celery import shared_task
from sqlalchemy.orm import Session
if TYPE_CHECKING:
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
logger = logging.getLogger(__name__)
_DITTO_URL_TTL_SECONDS = 7 * 24 * 3600 # Ditto 结果 OSS URL 7 天有效
def _get_db_session() -> Session:
try:
from worker_app.db import SessionLocal # type: ignore
except ImportError:
from app.db import SessionLocal # type: ignore
return SessionLocal()
def _sign_media_url(url: str) -> str:
"""对自家 OSS URL 签 7 天预签名。"""
if not url:
return url
try:
from urllib.parse import urlparse
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
public_base = getattr(storage, "public_url", "")
if not isinstance(public_base, str) or not public_base:
return url
own_host = urlparse(public_base).netloc.lower()
host = urlparse(url).netloc.lower()
if not own_host or host != own_host:
return url
return storage.get_download_url(url, expires_seconds=_DITTO_URL_TTL_SECONDS)
except Exception:
return url
def _probe_video_duration(video_bytes: bytes) -> float:
"""用 ffprobe 探测视频时长(秒);失败返回 0。"""
try:
import os
import subprocess
import tempfile
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp:
tmp.write(video_bytes)
tmp_path = tmp.name
try:
out = subprocess.check_output(
[
"ffprobe",
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
tmp_path,
],
stderr=subprocess.STDOUT,
timeout=10,
)
return float(out.decode().strip() or 0)
finally:
os.unlink(tmp_path)
except Exception as exc:
logger.warning("[ditto_task] ffprobe 失败: %s", exc)
return 0.0
def _refund_lip_sync(db: Session, job: "LipsyncJobModel") -> None:
"""Ditto 失败/取消时全额退款(复用 lipsync_service 的退款逻辑)。"""
try:
from app.services.lipsync_service import LipsyncService
LipsyncService(db)._refund_lip_sync(job)
except Exception:
logger.exception("[ditto_task] lip_sync 退款异常 job_id=%s", job.id)
def _settle_lip_sync(db: Session, job: "LipsyncJobModel", duration: float) -> None:
"""Ditto 成功后按实际时长结算。"""
try:
from app.services.lipsync_service import LipsyncService
LipsyncService(db)._settle_lip_sync(job, duration)
except Exception:
logger.exception("[ditto_task] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
def _fallback_to_gpu_then_mediakit(db: Session, job: "LipsyncJobModel") -> None:
"""Ditto 失败后:优先回退 GPU MuseTalk,再回退 MediaKit 云端。
复用 lipsync_service 现有路径逻辑以保证兜底一致性。
"""
# 先尝试走 GPU MuseTalk(若可用)
try:
from app.services.gpu_lipsync_service import GpuLipsyncService
from app.tasks.lipsync_gpu import lipsync_gpu_process_async
gpu_svc = GpuLipsyncService(db)
if gpu_svc.has_available_worker():
logger.info("[ditto_task] 回退 GPU MuseTalk: job_id=%s", job.id)
# 复用 lipsync_service._submit_to_gpu_create 逻辑
from app.services.lipsync_service import LipsyncService
svc = LipsyncService(db)
storage = _shared_storage()
persisted_audio = None
try:
persisted_audio = svc._persist_external_audio_for_gpu(job=job, storage=storage)
except Exception as exc:
logger.warning("[ditto_task] GPU 外部音频转存失败: %s", exc)
audio_url_for_task = persisted_audio or job.audio_url
gpu_task = gpu_svc.create_task(
video_url=job.video_url,
audio_url=audio_url_for_task,
lipsync_job_id=job.id,
user_id=job.user_id,
)
if gpu_task is not None:
job.mediakit_task_id = f"gpu:{gpu_task.id}"
job.status = "processing"
job.updated_at = datetime.now(UTC)
db.commit()
lipsync_gpu_process_async.apply_async(args=(job.id, job.user_id, gpu_task.id))
return
db.rollback()
except Exception as exc:
logger.warning("[ditto_task] GPU MuseTalk 回退失败,转 MediaKit: %s", exc)
try:
db.rollback()
except Exception:
pass
# 最后兜底:MediaKit 云端
try:
from app.services.mediakit_client import get_mediakit_client
client = get_mediakit_client()
video_url = _sign_media_url(job.video_url)
signed_audio_url = _sign_media_url(job.audio_url)
result = client.submit_lipsync(
video_url=video_url,
audio_url=signed_audio_url,
enable_video_loop=job.enable_video_loop,
client_token=job.id,
)
job.mediakit_task_id = result["task_id"]
job.status = "submitted"
job.submitted_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
db.commit()
logger.info("[ditto_task] 已回退 MediaKit: job_id=%s task_id=%s", job.id, result["task_id"])
except Exception as exc:
job.status = "failed"
job.error_message = f"Ditto/GPU/MediaKit 均失败: {exc}"
job.error_code = "AllBackendsFailed"
job.updated_at = datetime.now(UTC)
db.commit()
logger.error("[ditto_task] 所有兜底均失败: job_id=%s err=%s", job.id, exc)
def _shared_storage():
from packages.shared.storage import get_shared_storage_service
return get_shared_storage_service()
@shared_task(
name="lipsync_ditto_process_async",
bind=True,
max_retries=0,
acks_late=True,
time_limit=600,
soft_time_limit=540,
)
def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> None:
"""异步调用 Ditto 生成口型视频。
Args:
job_id: LipsyncJob ID
user_id: 用户 ID
"""
from packages.application.ditto_service import DittoError, get_ditto_client
db: Session = _get_db_session()
job: Optional[LipsyncJobModel] = None
try:
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
job = db.query(LipsyncJobModel).filter_by(id=job_id, user_id=user_id).first()
if job is None:
logger.error("[ditto_task] job 不存在: job_id=%s", job_id)
return
if job.status != "processing":
logger.warning(
"[ditto_task] job 状态异常(非 processing),跳过: job_id=%s status=%s",
job_id,
job.status,
)
return
audio_url = job.audio_url or ""
script = job.script_text or ""
if not audio_url:
raise DittoError("job.audio_url 为空,无法调用 Ditto", code="InvalidParam")
logger.info(
"[ditto_task] 开始 Ditto 生成: job_id=%s audio=%s script_len=%d",
job_id,
audio_url[:100],
len(script),
)
client = get_ditto_client()
result = client.generate_and_persist(
job_id=job_id,
user_id=user_id,
audio_url=audio_url,
script=script,
# video_url 不传则用默认模板
)
# Ditto 返回的 MP4 自带音频,直接标记完成
job.output_video_url = result.video_url
# 探测时长(用于计费)
duration = _probe_video_duration(result.video_bytes)
if duration <= 0:
# 兜底:按音频时长估算(1秒≈1秒)
try:
from packages.domain.sentence_timings import probe_audio_duration
from packages.shared.url_security import safe_download_bytes
audio_data = safe_download_bytes(
audio_url, allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav"), timeout=30
)
duration = probe_audio_duration(audio_data)
except Exception:
duration = 0.0
job.output_duration = duration
job.status = "completed"
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
db.commit()
logger.info(
"[ditto_task] Ditto 完成: job_id=%s url=%s duration=%.2fs rtf=%.2f frames=%d",
job_id,
result.video_url[:100],
duration,
result.rtf,
result.frames,
)
_settle_lip_sync(db, job, duration)
except DittoError as exc:
logger.error("[ditto_task] Ditto 失败,回退: job_id=%s code=%s err=%s", job_id, exc.code, exc)
if job is not None:
try:
db.rollback()
job = db.query(type(job)).filter_by(id=job_id).first() if hasattr(job, "id") else job
# 回退 GPU/MediaKit
_fallback_to_gpu_then_mediakit(db, job)
except Exception as fallback_exc:
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
try:
if job:
job.status = "failed"
job.error_message = f"Ditto 失败且回退异常: {exc}; fallback: {fallback_exc}"
job.error_code = "FallbackError"
job.updated_at = datetime.now(UTC)
db.commit()
except Exception:
pass
except Exception as exc:
logger.exception("[ditto_task] 未预期异常: job_id=%s err=%s", job_id, exc)
if job is not None:
try:
db.rollback()
job = db.query(type(job)).filter_by(id=job_id).first()
_fallback_to_gpu_then_mediakit(db, job)
except Exception as fallback_exc:
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
try:
if job:
job.status = "failed"
job.error_message = f"Ditto 异常: {exc}"
job.error_code = "DittoAsyncError"
job.updated_at = datetime.now(UTC)
db.commit()
except Exception:
pass
finally:
db.close()
-29
View File
@@ -104,7 +104,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
job.updated_at = datetime.now(UTC)
db.commit()
logger.info("[lipsync_gpu_async] GPU 任务已被用户取消: job_id=%s", job_id)
_refund_lip_sync(db, job)
return
if final_task.status != "done":
@@ -142,7 +141,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
job_id,
job.output_duration,
)
_settle_lip_sync(db, job, final_task)
except Exception as exc:
logger.exception("[lipsync_gpu_async] 异常: job_id=%s err=%s", job_id, exc)
try:
@@ -159,33 +157,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
db.close()
def _settle_lip_sync(db: Session, job: LipsyncJobModel, gpu_task) -> None:
"""GPU 成功后结算:同步 credits_cost 到 gpu 任务并按实际时长多退少不补。"""
try:
from app.services.lipsync_service import LipsyncService
# GPU 任务表先同步结算结果(标记用)
LipsyncService._settle_lip_sync(job, float(getattr(gpu_task, "result_duration", 0) or 0.0))
gpu_task.credits_cost = float(job.credits_cost or 0.0)
db.commit()
except Exception: # noqa: BLE001
logger.exception("[lipsync_gpu_async] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
try:
db.rollback()
except Exception: # noqa: BLE001
pass
def _refund_lip_sync(db: Session, job: LipsyncJobModel) -> None:
"""GPU 取消/失败路径全额退款。"""
try:
from app.services.lipsync_service import LipsyncService
LipsyncService(db)._refund_lip_sync(job)
except Exception: # noqa: BLE001
logger.exception("[lipsync_gpu_async] lip_sync 退款异常 job_id=%s", job.id)
def _fallback_to_mediakit(db: Session, job: LipsyncJobModel) -> None:
"""GPU 失败时回退到 MediaKit 云端渲染。"""
try:
+1 -27
View File
@@ -261,33 +261,7 @@ def tts_synthesize_and_submit(
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
)
# 3. 优先走 Ditto(#2076):开关打开且配置完整时,派发 Ditto 异步任务,不再走 MediaKit
ditto_dispatched = False
try:
from packages.config import get_api_settings as _get_settings
_settings = _get_settings()
if _settings.use_ditto_lipsync and _settings.ditto_api_base_url and _settings.ditto_default_video_url:
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
job.status = "processing"
job.mediakit_task_id = "ditto:tts-submitted"
job.updated_at = datetime.now(UTC)
db.commit()
lipsync_ditto_process_async.apply_async(args=(job_id, user_id))
logger.info("[lipsync_tts] TTS 完成,已派发 Ditto 任务: job_id=%s", job_id)
ditto_dispatched = True
except Exception as _ditto_err:
logger.warning("[lipsync_tts] Ditto 派发失败,回退 MediaKit: job_id=%s err=%s", job_id, _ditto_err)
try:
db.rollback()
except Exception:
pass
if ditto_dispatched:
return
# 4. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
audio_url = _sign_media_url(job.audio_url)
video_url = _sign_media_url(job.video_url)
-12
View File
@@ -14,7 +14,6 @@
"axios": "^1.7.2",
"classnames": "^2.5.1",
"dayjs": "^1.11.23",
"marked": "^12.0.2",
"mp4box": "^2.4.1",
"react": "^18.3.1",
"react-dom": "^18.3.1",
@@ -4503,17 +4502,6 @@
"url": "https://github.com/sponsors/sindresorhus"
}
},
"node_modules/marked": {
"version": "12.0.2",
"resolved": "https://registry.npmmirror.com/marked/-/marked-12.0.2.tgz",
"integrity": "sha512-qXUm7e/YKFoqFPYPa3Ukg9xlI5cyAtGmyEIzMfW//m6kXwCy2Ps9DYf5ioijFKQ8qyuscrHoY04iJGctu2Kg0Q==",
"bin": {
"marked": "bin/marked.js"
},
"engines": {
"node": ">= 18"
}
},
"node_modules/math-intrinsics": {
"version": "1.1.0",
"resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz",
-1
View File
@@ -25,7 +25,6 @@
"axios": "^1.7.2",
"classnames": "^2.5.1",
"dayjs": "^1.11.23",
"marked": "^12.0.2",
"mp4box": "^2.4.1",
"react": "^18.3.1",
"react-dom": "^18.3.1",
+15 -13
View File
@@ -65,23 +65,27 @@ export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
return isImageAnalysisStage(stage) || isCopyStage(stage)
}
/** 单张图片 VLM 识别结果(v8 叙述优先,仅保留最少结构化字段) */
/** 单张图片 VLM 识别出的商品信息 */
export interface ImageProductAnalysis {
/** store / product / person / scene */
type?: string
name?: string
brand?: string
has_person?: boolean
/** v8: 用户端展示用的叙述 markdown(由提示词控制排版) */
summary_markdown?: 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 {
/** v8 字段 */
images?: ImageProductAnalysis[]
/** 老数据兼容 */
products?: ImageProductAnalysis[]
}
@@ -128,8 +132,6 @@ export interface CopyResult {
/** 向后兼容:= voiceover_script */
suggested_copy?: string
title?: string
/** v3 storyboard: 用户端展示用的 markdown 文案(由提示词控制排版) */
copy_display_markdown?: string
/** v1.5 旧字段兼容(老数据降级时可能出现) */
scenes?: Array<{ shot: string; narration: string; duration?: number }>
}
@@ -1050,13 +1050,10 @@
flex-direction: column;
align-items: center;
justify-content: center;
height: 360px;
padding: 28px 16px;
gap: 10px;
background: #fff;
border: 1px solid #e5e7eb;
border-radius: 10px;
margin-top: 8px;
}
.vv-copy-loading .vv-spinner {
width: 28px;
@@ -1085,38 +1082,18 @@
/* ── Storyboard (linear doc style) ── */
.vv-storyboard {
display: flex;
flex-direction: column;
height: 360px;
padding: 10px 12px;
background: #fff;
border: 1px solid #e5e7eb;
border-radius: 10px;
margin-top: 8px;
overflow: hidden;
padding: 6px 2px;
background: transparent;
border: none;
}
.vv-sb-doc {
flex: 1 1 auto;
display: flex;
flex-direction: column;
gap: 3px;
color: #1f2937;
font-size: 13px;
line-height: 1.55;
overflow-y: auto;
padding-right: 4px;
margin-right: -4px;
}
.vv-sb-doc::-webkit-scrollbar {
width: 6px;
}
.vv-sb-doc::-webkit-scrollbar-thumb {
background: #d8c4ff;
border-radius: 3px;
}
.vv-sb-doc::-webkit-scrollbar-track {
background: transparent;
}
.vv-sb-h {
margin: 6px 0 2px;
@@ -1406,14 +1383,12 @@
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
.vv-sb-actions {
flex-shrink: 0;
display: flex;
align-items: center;
gap: 8px;
margin-top: 8px;
padding-top: 8px;
border-top: 1px solid #e5e7eb;
background: #fff;
}
.vv-sb-actions .vv-btn-ghost {
padding: 6px 14px;
@@ -1973,99 +1948,3 @@
padding-bottom: 6px;
border-bottom: 1px dashed #e5e7eb;
}
/* ─────────── markdown 渲染(提示词控制展示格式) ─────────── */
.vv-recog-md {
padding: 4px 0;
}
.vv-copy-preview {
margin-bottom: 14px;
padding: 12px 14px;
background: linear-gradient(180deg, #faf7ff 0%, #f6f2ff 100%);
border: 1px solid #ece4fb;
border-radius: 10px;
}
.vv-copy-preview-h {
margin: 0 0 8px;
border-bottom: none;
padding-bottom: 0;
}
.vv-md-body {
font-size: 13px;
line-height: 1.7;
color: #374151;
word-break: break-word;
}
.vv-md-body h1,
.vv-md-body h2,
.vv-md-body h3,
.vv-md-body h4 {
margin: 10px 0 6px;
font-weight: 600;
color: #1f2937;
line-height: 1.4;
}
.vv-md-body h1 {
font-size: 18px;
}
.vv-md-body h2 {
font-size: 16px;
}
.vv-md-body h3 {
font-size: 15px;
}
.vv-md-body h4 {
font-size: 14px;
}
.vv-md-body p {
margin: 6px 0;
}
.vv-md-body ul,
.vv-md-body ol {
margin: 6px 0;
padding-left: 20px;
}
.vv-md-body li {
margin: 3px 0;
}
.vv-md-body strong {
color: #111827;
font-weight: 600;
}
.vv-md-body blockquote {
margin: 8px 0;
padding: 4px 12px;
border-left: 3px solid #7c3aed;
background: rgba(124, 58, 237, 0.05);
color: #4b5563;
}
.vv-md-body code {
padding: 1px 5px;
background: #f3f4f6;
border-radius: 4px;
font-size: 12px;
color: #be185d;
}
.vv-md-body a {
color: #7c3aed;
text-decoration: none;
}
.vv-md-body a:hover {
text-decoration: underline;
}
.vv-md-body table {
border-collapse: collapse;
margin: 8px 0;
width: 100%;
}
.vv-md-body th,
.vv-md-body td {
border: 1px solid #e5e7eb;
padding: 6px 10px;
text-align: left;
}
.vv-md-body hr {
border: none;
border-top: 1px solid #e5e7eb;
margin: 12px 0;
}
@@ -1,6 +1,5 @@
import React, { useCallback, useEffect, useRef, useState } from "react"
import axios from "axios"
import { marked } from "marked"
import {
PlusOutlined,
CloseOutlined,
@@ -151,16 +150,6 @@ type TabTask = {
audioInst: HTMLAudioElement | null
}
/* ── marked 配置:禁用 mangle/headerIds,输出干净 HTML ── */
marked.setOptions({ gfm: true, breaks: false })
const renderMarkdown = (md: string): string => {
try {
return marked.parse(md ?? "", { async: false }) as string
} catch {
return (md ?? "").replace(/&/g, "&amp;").replace(/</g, "&lt;")
}
}
/* ─────────── 常量 ─────────── */
const LANGUAGES = ["中文(普通话)", "粤语", "英语", "日语", "韩语"]
@@ -307,8 +296,6 @@ interface Storyboard {
hard_constraints: string[]
negative_prompts: string[]
voiceover_script: string
/** v3: 用户端展示用 markdown 文案(由提示词控制排版) */
copy_display_markdown: string
}
/** 兼容旧 copy_result(final_copy/title/scenes)→ 新 Storyboard 结构 */
@@ -336,7 +323,6 @@ function copyResultToStoryboard(cr: CopyResult | null | undefined): Storyboard |
hard_constraints: Array.isArray(cr.hard_constraints) ? cr.hard_constraints : [],
negative_prompts: Array.isArray(cr.negative_prompts) ? cr.negative_prompts : [],
voiceover_script: cr.voiceover_script || cr.final_copy || cr.suggested_copy || "",
copy_display_markdown: cr.copy_display_markdown || "",
}
}
// 兜底:旧结构转简单分镜
@@ -372,7 +358,6 @@ function copyResultToStoryboard(cr: CopyResult | null | undefined): Storyboard |
hard_constraints: [],
negative_prompts: [],
voiceover_script: finalCopy,
copy_display_markdown: cr.copy_display_markdown || "",
}
}
@@ -427,7 +412,6 @@ const MOCK_STORYBOARD: Storyboard = {
negative_prompts: ["冷色调", "模糊", "变形", "水印文字", "卡通风格", "空无一人"],
voiceover_script:
"还在为餐桌选不到好桌子发愁?这张北美黑胡桃木餐桌,一家人坐下来吃饭刚刚好。全实木、无贴皮,纹理好看又耐刮。点小黄车,给家里添一张好桌子。",
copy_display_markdown: "",
}
const fmtSize = (bytes: number | undefined) => {
@@ -1208,10 +1192,8 @@ const ViralVideoPage: React.FC = () => {
/* ── 识别描述汇览渲染 ── */
const renderRecognition = () => {
const images: ImageProductAnalysis[] =
(task.imageAnalysis?.images as ImageProductAnalysis[] | undefined) ||
(task.imageAnalysis?.products as ImageProductAnalysis[] | undefined) ||
[]
const products: ImageProductAnalysis[] =
(task.imageAnalysis?.products as ImageProductAnalysis[] | undefined) || []
if (task.uiStep === "step1_analyzing") {
return (
<div className="vv-recog">
@@ -1219,36 +1201,83 @@ const ViralVideoPage: React.FC = () => {
<LoadingOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
识别描述汇览
</div>
<div className="vv-muted">AI 正在识别画面…</div>
<div className="vv-muted">AI 正在识别商品特征…</div>
</div>
)
}
if (images.length === 0) return null
if (products.length === 0) return null
const featureText = (f: string[] | string | undefined) => {
if (!f) return ""
if (Array.isArray(f)) return f.join(";")
return f
}
return (
<div className="vv-recog">
<div className="vv-recog-title">
<CheckCircleFilled style={{ color: "#10b981" }} />
识别描述汇览
</div>
{images.map((p, i) => {
const meta = [p.name || "未识别", p.brand, p.category].filter(Boolean)
return (
<div key={i} className="vv-recog-item vv-recog-md">
<div className="vv-recog-line">
<span className="vv-recog-k">图片{i + 1}:</span>
<span>{meta.join(" · ")}</span>
</div>
{p.summary_markdown ? (
<div
className="vv-md-body"
dangerouslySetInnerHTML={{ __html: renderMarkdown(p.summary_markdown) }}
/>
) : (
<div className="vv-muted">(暂无叙述描述)</div>
)}
{products.map((p, i) => (
<div key={i} className="vv-recog-item">
<div className="vv-recog-line">
<span className="vv-recog-k">图片{i + 1}:</span>
<span>
{p.name || "未识别"}
{p.spec && <span className="vv-recog-meta">({p.spec})</span>}
{p.brand && <span className="vv-recog-meta"> · {p.brand}</span>}
{p.category && <span className="vv-recog-meta"> · {p.category}</span>}
</span>
</div>
)
})}
{featureText(p.key_features ?? p.features) && (
<div className="vv-recog-line">
<span className="vv-recog-k">核心特征:</span>
<span className="vv-recog-v">{featureText(p.key_features ?? p.features)}</span>
</div>
)}
{p.colors && p.colors.length > 0 && (
<div className="vv-recog-line">
<span className="vv-recog-k">主色调:</span>
<span className="vv-recog-v">{p.colors.join(" / ")}</span>
</div>
)}
{p.material_or_texture && (
<div className="vv-recog-line">
<span className="vv-recog-k">材质/纹理:</span>
<span className="vv-recog-v">{p.material_or_texture}</span>
</div>
)}
{p.visual_style && (
<div className="vv-recog-line">
<span className="vv-recog-k">视觉风格:</span>
<span className="vv-recog-v">{p.visual_style}</span>
</div>
)}
{p.scene && (
<div className="vv-recog-line">
<span className="vv-recog-k">场景:</span>
<span className="vv-recog-v">{p.scene}</span>
</div>
)}
{p.target_audience_hint && (
<div className="vv-recog-line">
<span className="vv-recog-k">目标人群:</span>
<span className="vv-recog-v">{p.target_audience_hint}</span>
</div>
)}
{(p.text_on_image || p.label_text) && (
<div className="vv-recog-line">
<span className="vv-recog-k">包装文字:</span>
<span className="vv-recog-v">{p.text_on_image || p.label_text}</span>
</div>
)}
{p.selling_points && (
<div className="vv-recog-line">
<span className="vv-recog-k">卖点:</span>
<span className="vv-recog-v">{p.selling_points}</span>
</div>
)}
</div>
))}
</div>
)
}
@@ -1387,19 +1416,6 @@ const ViralVideoPage: React.FC = () => {
return (
<div className="vv-copy-box vv-storyboard">
<div className="vv-sb-doc">
{/* 文案预览(提示词控制排版,只读;编辑在下方分镜字段中进行) */}
{sb.copy_display_markdown && (
<div className="vv-copy-preview">
<h4 className="vv-sb-h vv-copy-preview-h">
<FileTextOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
文案预览
</h4>
<div
className="vv-md-body"
dangerouslySetInnerHTML={{ __html: renderMarkdown(sb.copy_display_markdown) }}
/>
</div>
)}
{/* 视频总览 */}
<h4 className="vv-sb-h">视频总览</h4>
<p className="vv-sb-inline-row">
-3
View File
@@ -53,9 +53,6 @@ celery_app.conf.imports = (
# #1998 GPU MuseTalk 异步推理:wait_for_result→签名 URL→回写 lipsync_jobs
# 必须在 Worker 侧注册,否则 apply_async 消息无人消费,job 永远卡在 processing
"app.tasks.lipsync_gpu",
# #2076 Ditto 蚂蚁数字人异步推理:同步 HTTP 调用 Ditto → MP4 流转存 OSS → 回写 lipsync_jobs
# 必须在 Worker 侧注册;失败回退 GPU MuseTalk → MediaKit
"app.tasks.lipsync_ditto",
)
# Celery Beat 定时任务调度
@@ -387,41 +387,6 @@ BATCH_RENDER_SIMILARITY_LIMIT = 0.20
"""批次内成片查重相似度阈值:超过则重选独立 plan 重渲一次(20%)。"""
def _refund_smart_edit_prepaid(task_id: str) -> None:
"""智能剪辑任务最终失败时退还预扣积分(幂等)。"""
session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.domain.points_service import PointsService
repo = SQLAlchemyGenerationTaskRepository(session)
task = repo.get(task_id)
if not task:
return
prepaid = float(getattr(task, "credits_prepaid", 0) or 0)
if prepaid <= 0:
return
txn_id = getattr(task, "credits_transaction_id", "") or ""
res = PointsService().refund_points(
user_id=task.user_id,
amount=prepaid,
source="smart_edit",
db=session,
ref_id=task.id,
related_transaction_id=txn_id or None,
description="智能剪辑任务失败退回",
)
task.credits_cost = 0.0
task.credits_prepaid = 0.0
repo.update(task)
if not res.get("success"):
logger.warning("[task_id=%s] 失败退积分未成功: %s", task_id, res)
finally:
session.close()
def should_rerender_for_batch_dedup(*, batch_id: str, render_attempt: int, batch_similarity) -> bool:
"""批次内查重后判定是否需要重选 plan 重渲。
@@ -1202,10 +1167,6 @@ def generate_video(self, task_id: str) -> dict:
"mark_failed",
error_message="source_edit_plan_id is required. Please create a preview task first.",
)
try:
_refund_smart_edit_prepaid(task_id)
except Exception:
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
return {
"status": "failed",
"task_id": task_id,
@@ -1244,7 +1205,6 @@ def generate_video(self, task_id: str) -> dict:
)
# ── 自动重试逻辑 ──────────────────────────────────────────────────
will_retry = False
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
@@ -1257,7 +1217,6 @@ def generate_video(self, task_id: str) -> dict:
if _task and _task.auto_retry_enabled and _task.auto_retry_max > 0:
current_retry = _task.retry_count or 0
if current_retry < _task.auto_retry_max:
will_retry = True
logger.info(
"[task_id=%s] 触发自动重试: 当前重试次数=%d, 最大重试次数=%d",
task_id,
@@ -1291,13 +1250,6 @@ def generate_video(self, task_id: str) -> dict:
exc_info=True,
)
# 最终失败(不再重试):退还 smart_edit 预扣积分
if not will_retry:
try:
_refund_smart_edit_prepaid(task_id)
except Exception:
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
return {
"status": "failed",
"task_id": task_id,
File diff suppressed because it is too large Load Diff
@@ -1,95 +0,0 @@
# -*- coding: utf-8 -*-
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates(prompt_type='image_analysis'
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到 prompts.py 的
image_analysis v8 默认 system/user。
规则:
- DB 有 is_active=true 的 image_analysis 记录:system 原样用 DB.system_prompt
(自带完整输出格式,不追加任何硬编码 schema),user 用 DB.user_prompt_template
渲染(填入 image_url / ocr_text);
- DB 无记录/异常:system/user 用 prompts.py 里的 v8 默认模板。
"""
from __future__ import annotations
import logging
import threading
import time
logger = logging.getLogger(__name__)
def _default_template() -> dict:
# 延迟导入:避免模块加载时拉起整个 packages 依赖链(也便于旧 Python 收集测试)
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
for item in DEFAULT_TEMPLATES:
if item["prompt_type"] == "image_analysis":
return item
raise RuntimeError("image_analysis 默认模板缺失")
_cache_lock = threading.Lock()
_cache: dict[str, tuple[float, tuple[str, str]]] = {}
_CACHE_TTL = 30.0
def _load_db_template():
"""查 DB is_active=true 的 image_analysis 记录;不可达/无记录返回 None。"""
try:
from packages.application.viral_video.prompt_loader import _load_from_db
return _load_from_db("image_analysis")
except Exception as e: # noqa: BLE001
logger.warning("[vision.v2] 查询DB image_analysis prompt失败: %s", e)
return None
def _render_user(user_tpl: str, image_url: str, ocr_text: str) -> str:
try:
return user_tpl.format(image_url=image_url, ocr_text=ocr_text or "无")
except Exception: # noqa: BLE001
return user_tpl
def _resolve(kind: str, image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
now = time.time()
cache_key = f"prompt_{kind}"
with _cache_lock:
hit = _cache.get(cache_key)
if hit and now - hit[0] < _CACHE_TTL:
sys_prompt, usr_prompt = hit[1]
return sys_prompt, _render_user(usr_prompt, image_url, ocr_text)
default = _default_template()
sys_prompt = default["system_prompt"]
usr_prompt = default["user_prompt_template"]
tpl = _load_db_template()
if tpl is not None:
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
if db_sys:
sys_prompt = db_sys
db_usr = getattr(tpl, "user_prompt_template", "") or usr_prompt
usr_prompt = db_usr or usr_prompt
logger.info(
"[vision.v2] 使用DB image_analysis prompt version=%s",
getattr(tpl, "version", "?"),
)
with _cache_lock:
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
return sys_prompt, _render_user(usr_prompt, image_url, ocr_text)
def resolve_fast_prompt(image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
return _resolve("fast", image_url, ocr_text)
def resolve_pro_prompt(image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
return _resolve("pro", image_url, ocr_text)
def invalidate_cache() -> None:
with _cache_lock:
_cache.clear()
+285 -84
View File
@@ -1,106 +1,307 @@
"""V2 结果组装(v8 叙述优先,大幅精简)。
# -*- coding: utf-8 -*-
"""把 fast_json VLM 输出 + OCR 文本组装为与旧 _normalize() 完全一致的 dict。
设计原则:VLM 直接输出最终给用户看的 summary_markdown,assembler 只负责
- 解析 fast JSON(兼容顶层 {"images":[...]} 与 {"products":[...]} 两种键);
- 补齐 5 个必备字段(type/name/brand/has_person/summary_markdown);
- summary_markdown 缺失(异常)时才拼一句最基础的兜底文字。
正常情况下不改写、不“润色”VLM 输出,不做 brand 多级兜底,不处理任何
colors/material/key_features 等细分字段。
目标:下游(信任链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
import logging
from typing import Any
logger = logging.getLogger(__name__)
_VALID_TYPES = ("store", "product", "person", "scene")
# ---------- portrait_prompt 模板 ----------
# 目标:60-100 字的人物穿搭描述,用于 Seedream 纯文生图。要求具体、风格化、视觉细节丰富。
# 旧 VLM 输出格式参考:"一位25岁左右的亚洲女性,身穿白色V领短袖T恤,黑色高腰阔腿裤,
# 搭配银色项链,长发披肩,表情自信,街拍风格,阳光明媚的城市街头"
def _coerce_bool(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
if isinstance(value, str):
return value.strip().lower() in ("true", "1", "yes", "是")
return False
def _join_parts(*parts: str | None) -> str:
return "".join(p for p in parts if p)
def _basic_markdown(image: dict[str, Any]) -> str:
"""异常兜底:VLM 没给 summary_markdown 时只拼一句基础文字。"""
name = (image.get("name") or "").strip() or "未识别"
brand = (image.get("brand") or "").strip()
typ = image.get("type") or "scene"
label = f"{brand}{name}" if brand and brand not in name else (brand or name)
if typ == "store":
return f"这是{label}的门店场景,画面细节识别不完整。"
if typ == "person":
return f"画面中的人物与{label}相关,细节识别不完整。"
if typ == "product":
return f"这是{label}的商品图片,具体外观细节识别不完整。"
return f"画面内容为{label},细节识别不完整。"
_AGE_PREFIX = {
"青年": "年轻",
"中年": "中年",
"老年": "老年",
}
# gender 后缀
_GENDER_WORD = {"男": "男性", "女": "女性"}
def _normalize_image(raw: Any, idx: int) -> dict[str, Any]:
"""把一条 VLM 输出归一化为 5 字段 dict。"""
if not isinstance(raw, dict):
raw = {}
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 "人物"
typ = str(raw.get("type") or "").strip().lower()
if typ not in _VALID_TYPES:
typ = "scene"
name = str(raw.get("name") or "").strip() or "未识别"
brand = str(raw.get("brand") or "").strip()
has_person = _coerce_bool(raw.get("has_person"))
if typ == "person" and not has_person:
# type=person 通常意味着主体是人,保持一致(仅异常补全)
has_person = True
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 ""
summary = raw.get("summary_markdown")
summary = summary.strip() if isinstance(summary, str) else ""
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}"
image: dict[str, Any] = {
"type": typ,
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,
"has_person": has_person,
"summary_markdown": summary,
"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",
}
if not summary:
image["summary_markdown"] = _basic_markdown(image)
image["_source"] = "summary_missing"
logger.info("[assembler] 图片 #%s 缺少 summary_markdown,使用基础兜底", idx)
return image
def _extract_items(fast_json: Any) -> list[Any]:
"""从 fast JSON 中取出图片条目:优先 images,兼容 products。"""
if not isinstance(fast_json, dict):
return []
items = fast_json.get("images")
if not isinstance(items, list):
items = fast_json.get("products")
return items if isinstance(items, list) else []
def assemble_result(idx: int, fast_json: Any, ocr_texts: list[str] | None = None) -> dict[str, Any]:
"""组装单张图片分析结果。
每次调用对应一张图片;fast_json 形如 {"images": [{...}]}(v8)。
返回单条 image dict(5 字段,必要时带 _source)。
"""
items = _extract_items(fast_json)
if items:
image = _normalize_image(items[0], idx)
else:
# 极端异常:fast 无任何可用条目,OCR 文字可作为名称线索
ocr_hint = ""
if ocr_texts:
ocr_hint = "、".join(t for t in ocr_texts if t)[:40]
image = _normalize_image({"name": ocr_hint or "未识别", "summary_markdown": ""}, idx)
image["_source"] = "empty_fast_json"
return image
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
@@ -1,11 +1,10 @@
# -*- coding: utf-8 -*-
"""V2 图片分析主路径(v8 叙述优先):每图并行 OCR(火山 MediaKit,未配置自动跳过)
+ fast VLM 强约束 JSON;失败时单次 pro VLM 兜底。
"""V2 图片分析主路径:每图并行 OCR(火山专用API)+ lite JSON VLM,失败时单次 pro VLM 兜底。
架构:
- 单图 2 路并行(OCR + fast VLM),外层 N 图全并发(workers=8);
- 兜底单次 pro VLM,无竞速/复杂重试;
- 输出统一为 5 字段 image dict(type/name/brand/has_person/summary_markdown)。
设计原则(灵应10-05要求):
- 主力路径简洁:单图2路并行,外层N图全并发
- 兜底简单:单次 pro VLM 调用,无竞速/重试/复杂超时
- 输出 dict 格式与旧版完全一致,下游零改动
"""
from __future__ import annotations
@@ -20,85 +19,111 @@ 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", "20"))
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "20"))
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "8"))
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "8"))
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
def _is_usable(r: dict[str, Any] | None) -> bool:
if not isinstance(r, dict):
return False
return bool((r.get("summary_markdown") or "").strip())
_FALLBACK_RESULT = {
"name": "未识别",
"brand": "无法判断",
"category": "非产品图",
"appearance": "无法判断",
"packaging": "无法判断",
"text_on_package": [],
"key_features": ["无法判断"],
"scene": "通用",
"mood": "",
"portrait_prompt": "无法判断",
"summary": "未识别",
}
def _basic_failure(ocr_result: list[str], fast_elapsed: float, source: str) -> dict[str, Any]:
image = assembler.assemble_result(-1, {}, ocr_result)
image["_source"] = source
image["_fast_elapsed"] = round(fast_elapsed, 2)
return image
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] = []
pool = ThreadPoolExecutor(max_workers=2)
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: # noqa: BLE001
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:
for f in (f_fj, f_ocr):
if not f.done():
f.cancel()
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
finally:
fast_elapsed = time.time() - t0
pool.shutdown(wait=False)
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", idx, fast_elapsed)
logger.info(
"[vision.v2] 图片 #%d fast命中 elapsed=%.2fs pp=%s",
idx,
fast_elapsed,
(assembled.get("portrait_prompt") or "")[:40],
)
return assembled
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, ocr_hint=ocr_result, timeout=_PRO_TIMEOUT)
if _is_usable(pro_result):
# 第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)
return _basic_failure(ocr_result, fast_elapsed, "v2_all_failed")
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 pro_timeout=%.0fs",
len(img_urls),
workers,
_FAST_TIMEOUT,
_PRO_TIMEOUT,
)
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)}
@@ -106,19 +131,14 @@ def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
idx = future_to_idx[fut]
try:
results[idx] = fut.result()
except Exception as e: # noqa: BLE001
except Exception as e:
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
results[idx] = assembler.assemble_result(idx, {}, [])
results[idx]["_source"] = "v2_future_exception" # type: ignore[index]
r = dict(_FALLBACK_RESULT)
r["_source"] = "v2_future_exception"
results[idx] = r
elapsed = time.time() - t0
succ = sum(1 for r in results if _is_usable(r))
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] # type: ignore[misc]
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]
@@ -1,138 +0,0 @@
# -*- coding: utf-8 -*-
"""VLM 返回文本的稳健 JSON 提取工具。
背景:复杂门店图 VLM 输出经常被 max_tokens 截断(finish_reason=length),
json.loads 失败后整个结果被丢弃,导致"未识别"。本工具提供:
1. markdown 代码块剥离(含只开不闭的截断场景)
2. 最外层 { } 切片
3. 非法控制字符清理
4. 直接 json.loads
5. 截断 JSON 括号/引号栈补全修复
6. 尾部逐字符截断重试(去除最后一个不完整 token 后修复)
成功返回 dict;截断修复产物带 _partial=True 标记;彻底失败返回 None。
"""
from __future__ import annotations
import json
import logging
import re
logger = logging.getLogger(__name__)
_CODE_FENCE_RE = re.compile(r"^```(?:json)?\s*\n?(.*?)\n?```\s*$", re.DOTALL)
def _strip_code_fence(s: str) -> str:
s = s.strip()
m = _CODE_FENCE_RE.match(s)
if m:
return m.group(1).strip()
# 兼容开头 ```json 但结尾无 ```(截断场景)
if s.startswith("```"):
lines = s.split("\n")
if lines and lines[0].startswith("```"):
lines = lines[1:]
s = "\n".join(lines).strip()
return s
def _repair_truncated_json(text: str) -> str:
"""尝试补全被截断的JSON:维护 bracket/quote 栈,在末尾补闭合符。"""
stack: list[str] = []
in_string = False
escape = False
for ch in text:
if escape:
escape = False
continue
if ch == "\\" and in_string:
escape = True
continue
if ch == '"':
in_string = not in_string
continue
if in_string:
continue
if ch in "{[":
stack.append(ch)
elif ch == "}":
if stack and stack[-1] == "{":
stack.pop()
elif ch == "]":
if stack and stack[-1] == "[":
stack.pop()
repair = ""
if in_string:
repair += '"'
for opener in reversed(stack):
repair += "}" if opener == "{" else "]"
if repair:
logger.info(
"[json_utils] 截断JSON修复: 补全%d个闭合符 in_string=%s",
len(repair),
in_string,
)
return text + repair
def _clean_invalid_chars(text: str) -> str:
"""清理JSON中非法的控制字符(tab/newline 之外的 0x00-0x1f 段)。"""
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]", "", text)
def extract_json_object(text: str) -> dict | None:
"""从VLM返回文本中稳健提取JSON对象。
返回 dict 或 None。成功的 dict 可能带 _partial=True 标记,
表示原始文本被截断、经括号补全后得到的产物。
"""
if not text or not isinstance(text, str):
return None
# 1. 剥离 markdown
text = _strip_code_fence(text)
# 2. 找最外层 { }
lpos = text.find("{")
if lpos < 0:
return None
rpos = text.rfind("}")
if rpos > lpos:
text = text[lpos : rpos + 1]
else:
# 截断场景:无任何闭合 },取到末尾交给修复器
text = text[lpos:]
# 3. 清理非法控制字符
text = _clean_invalid_chars(text)
# 4. 直接 loads
try:
obj = json.loads(text)
return obj if isinstance(obj, dict) else None
except json.JSONDecodeError:
pass
# 5. 尝试截断修复
repaired = _repair_truncated_json(text)
try:
obj = json.loads(repaired)
if isinstance(obj, dict):
obj["_partial"] = True
return obj
except json.JSONDecodeError:
pass
# 6. 尾部逐字符截断重试(去除最后一个不完整 token)
for _ in range(50):
last_comma = repaired.rfind(",")
last_brace = max(repaired.rfind("}"), repaired.rfind("]"))
cut = max(last_comma, last_brace)
if cut < 10:
break
repaired = repaired[: cut + 1]
repaired = _repair_truncated_json(repaired)
try:
obj = json.loads(repaired)
if isinstance(obj, dict):
obj["_partial"] = True
return obj
except json.JSONDecodeError:
continue
return None
@@ -1,109 +1,226 @@
# -*- coding: utf-8 -*-
"""V2 兜底路径:vision client(fallback 变体)单图调用,走 v8 叙述优先 prompt。
"""VLM 兜底:专用API路径失败时的最后一道防线,单次调用 doubao-seed-2.1-pro。
fast 超时/非 JSON/为空时单次调用;输出统一走 assembler.assemble_result 组装,
与 fast 路径同为 5 字段 image dict。
设计原则:简单、直接、无竞速、无复杂超时逻辑。只在 fast_json 结果不可用时调用。
"""
from __future__ import annotations
import json
import logging
import re
import time
from typing import Any
from . import _prompt, assembler
logger = logging.getLogger(__name__)
_DEFAULT_TIMEOUT = 45
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,
*,
ocr_hint: list[str] | None = None,
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int | None = None,
model: str | None = None,
timeout: int = DEFAULT_TIMEOUT,
) -> dict[str, Any] | None:
"""单次调用 pro VLM,解析后返回 product dict;失败返回 None。"""
t0 = time.time()
ocr_text = "、".join(t for t in (ocr_hint or []) if t)[:200]
try:
from packages.shared.ai_router import ai_router
client = ai_router.get_vision_client("image_analysis", variant="fallback")
if not client or not client.is_available:
logger.warning("[vision.v2] pro vision client 不可用,跳过")
return None
except Exception as e: # noqa: BLE001
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
return None
system_prompt, user_prompt = _prompt.resolve_pro_prompt(img_url, ocr_text)
messages: list[dict[str, Any]] = [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
]
try:
call_kwargs: dict[str, Any] = {
"messages": messages,
"images": None,
"temperature": 0.3,
"timeout": timeout,
"enable_thinking": False,
"response_format": {"type": "json_object"},
"max_tokens": max_tokens if max_tokens is not None else 4000,
}
from .json_utils import extract_json_object
obj = None
for _outer in range(2):
kw = dict(call_kwargs)
if _outer == 1:
kw.pop("response_format", None)
msgs2 = [dict(messages[0]), dict(messages[1])]
cont = [dict(c) for c in msgs2[1]["content"]]
cont[-1] = {
"type": "text",
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
}
msgs2[1] = {"role": "user", "content": cont}
kw["messages"] = msgs2
raw = client.vision_completion(**kw)
if not raw:
continue
obj = extract_json_object(raw)
if obj is not None:
break
logger.warning("[vision.v2] pro 非JSON(100字) outer=%s: %s", _outer, raw[:100])
elapsed = time.time() - t0
if obj is None:
logger.warning("[vision.v2] pro 两次均未得到JSON elapsed=%.1fs", elapsed)
return None
if obj.get("_partial"):
logger.warning("[vision.v2] pro 截断JSON(partial) elapsed=%.1fs", elapsed)
result = assembler.assemble_result(idx, obj, ocr_hint or [])
result["_source"] = "vlm_pro"
result["_fallback_used"] = True
logger.info("[vision.v2] pro 完成 model=%s elapsed=%.1fs", client.model, elapsed)
return result
except Exception as e: # noqa: BLE001
logger.warning(
"[vision.v2] pro 异常 elapsed=%.1fs err=%s",
time.time() - t0,
e,
exc_info=True,
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
@@ -1,111 +1,206 @@
# -*- coding: utf-8 -*-
"""V2 快速路径:vision client(默认 image_analysis 能力)强约束 JSON-only 调用。
"""doubao-seed-2.1-lite 强约束 JSON-only 调用。
要点:
- 通过 ai_router.get_vision_client() 获取 client;
- enable_thinking=False 关闭推理链,response_format=json_object 强约束 JSON;
- system/user prompt 优先读后台模板(v8 叙述优先),DB 不可用时用 prompts.py 默认;
- temperature=0.1(稳定输出 JSON);两次尝试(第二次去 json_object 约束)。
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 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
from . import _prompt
logger = logging.getLogger(__name__)
_DEFAULT_TIMEOUT = 20
# 极简 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,
*,
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int | None = None,
model: str | None = None,
timeout: int = DEFAULT_TIMEOUT,
max_tokens: int = DEFAULT_MAX_TOKENS,
) -> dict[str, Any] | None:
"""调用 lite VLM 返回结构化 dict;失败/非 JSON 返回 None。
直接用 httpx 发最小 payload(关闭 thinking),不走 ai_client 包装:
- 关闭 thinking/推理链(reasoning_tokens 是延迟主因,单次要10-12s)
- 单次调用不重试(失败由外层走 pro 兜底)
- 温度=0.1 稳定输出 JSON
"""
t0 = time.time()
import httpx
try:
from packages.shared.ai_router import ai_router
from packages.shared import get_shared_settings
client = ai_router.get_vision_client("image_analysis", variant="primary")
if not client or not client.is_available:
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
settings = get_shared_settings()
api_key = settings.doubao_api_key
base_url = (settings.doubao_base_url or "https://ark.cn-beijing.volces.com/api/v3").rstrip("/")
if not api_key:
logger.warning("[vision.v2] doubao api_key 未配置,跳过 fast_json")
return None
except Exception as e: # noqa: BLE001
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
return None
system_prompt, user_prompt = _prompt.resolve_fast_prompt(img_url, "")
messages: list[dict[str, Any]] = [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
use_model = model or DEFAULT_LITE_MODEL
url = f"{base_url}/chat/completions"
payload: dict[str, Any] = {
"model": use_model,
"messages": [
{"role": "system", "content": _FAST_SYSTEM},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": _FAST_USER},
],
},
],
},
]
try:
call_kwargs: dict[str, Any] = {
"messages": messages,
"images": None,
"temperature": 0.1,
"timeout": timeout,
"enable_thinking": False,
"response_format": {"type": "json_object"},
"max_tokens": max_tokens,
"stream": False,
}
if max_tokens is not None:
call_kwargs["max_tokens"] = max_tokens
from .json_utils import extract_json_object
obj = None
for _outer in range(2):
kw = dict(call_kwargs)
if _outer == 1:
kw.pop("response_format", None)
msgs2 = [dict(messages[0]), dict(messages[1])]
cont = [dict(c) for c in msgs2[1]["content"]]
cont[-1] = {
"type": "text",
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
}
msgs2[1] = {"role": "user", "content": cont}
kw["messages"] = msgs2
raw = client.vision_completion(**kw)
if not raw:
continue
obj = extract_json_object(raw)
if obj is not None:
break
logger.warning("[vision.v2] fast_json 非JSON(100字) outer=%s: %s", _outer, raw[:100])
# 关键:关闭 thinking(避免产生 reasoning_tokens 拖慢响应)
# 方舟/豆包 2.x 模型支持 thinking.type=disabled
try:
payload["thinking"] = {"type": "disabled"}
except Exception:
pass
# 部分模型用 reasoning_effort 控制思考深度
payload["reasoning_effort"] = "low"
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if obj is None:
logger.warning("[vision.v2] fast_json 两次均未得到JSON elapsed=%.1fs", elapsed)
if resp.status_code != 200:
logger.warning(
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
)
# 如果400说明不支持thinking参数,降级重试一次
if resp.status_code == 400 and "thinking" in resp.text.lower():
payload.pop("thinking", None)
payload.pop("reasoning_effort", None)
time.time()
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if resp.status_code != 200:
logger.warning("[vision.v2] fast_json 降级后 HTTP %d elapsed=%.1fs", resp.status_code, elapsed)
return None
else:
return None
data = resp.json()
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
if raw is None:
logger.warning("[vision.v2] fast_json 返回 None elapsed=%.1fs", elapsed)
return None
if obj.get("_partial"):
logger.warning("[vision.v2] fast_json 截断JSON(partial) elapsed=%.1fs", elapsed)
usage = data.get("usage") or {}
logger.info(
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs type=%s",
client.model,
"[vision.v2] fast_json 直连完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
use_model,
elapsed,
obj.get("type"),
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
usage.get("reasoning_tokens", 0),
)
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: # noqa: BLE001
logger.warning(
"[vision.v2] fast_json 异常 elapsed=%.1fs err=%s",
time.time() - t0,
e,
exc_info=True,
)
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
@@ -335,10 +335,6 @@ class GenerationTaskModel(Base):
bgm_config = Column(JSON, nullable=False, default=dict)
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
logs = Column(Text, nullable=False, default="[]", server_default="[]")
# 功能计费(smart_edit):预扣积分 / 最终积分 / 预扣流水 ID
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
updated_at = Column(
DateTime,
@@ -731,11 +727,6 @@ class LipsyncJobModel(Base):
# 精确句子时间戳(TTS 合成后由 silencedetect 计算,用于 B-roll 精确定位)
sentence_timings = Column(JSON, nullable=True) # list[{index,text,start_time,end_time}]
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
# 时间戳
submitted_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
@@ -914,11 +905,6 @@ class GpuLipsyncTaskModel(Base):
# 心跳:worker 最近一次 poll/result 的时间,用于判定 worker 失联
last_heartbeat_at = Column(DateTime, nullable=True)
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
class GpuWorkerModel(Base):
"""GPU Worker 注册表 — 反向轮询模式下用于心跳与监控."""
+4 -17
View File
@@ -351,25 +351,12 @@ class CosyVoiceService:
用于私有 bucket 下,将裸 URL 转为预签名 URL,
确保 CosyVoice 服务器能下载参考音频.
"""
# 优先从 ai_router 获取 DB 配置
_router_key, _router_url, _router_model = "", "", ""
try:
from packages.shared.ai_router import ai_router
tts_client = ai_router.get_tts_client("tts")
if tts_client and tts_client.is_available:
_router_key = tts_client.api_key
_router_url = tts_client.base_url
_router_model = tts_client.model
except Exception:
pass
settings = get_shared_settings()
self._api_key = api_key or _router_key or settings.cosyvoice_api_key
self._base_url = base_url or _router_url or settings.cosyvoice_base_url
self._model = model or _router_model or settings.cosyvoice_model
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "")
self._api_key = api_key or settings.cosyvoice_api_key
self._base_url = base_url or settings.cosyvoice_base_url
self._model = model or settings.cosyvoice_model
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "voice-enrollment")
self._audio_url_signer = audio_url_signer
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
-269
View File
@@ -1,269 +0,0 @@
"""蚂蚁 Ditto 数字人口型 API 客户端 — #2076.
封装 Ditto FastAPI(部署在 5060Ti GPU 节点,Tailscale 内网可达):
- GET /health 健康检查
- POST /generate 生成口型视频(同步返回 MP4 流)
关键特性:
- 入参:video_url(人物模板视频 URL) + audio_url(TTS 音频 URL) + script(文案原文)
- 出参:直接返回 video/mp4 字节流(自带音频,无需二次混流)
- 429 时指数退避重试(最多 ditto_max_retries 次)
- 500/超时视为失败
- 输出 MP4 字节流转存到自家 OSS,返回公网 URL
注意:
- 保留 MuseTalk/GPU 路径不变;本服务作为更高优先级的第三条口型路径
- 不传 emotion/表情精细控制,使用默认 emo_global=4(中性)+ use_script_emo=true(关键词驱动表情)
- Ditto 输出自带音视频,不需要 GFPGAN 超分,不需要 ffmpeg 音视频混流
"""
from __future__ import annotations
import io
import logging
import time
from dataclasses import dataclass
from typing import Optional
import httpx
from packages.config import get_api_settings
logger = logging.getLogger(__name__)
class DittoError(Exception):
"""Ditto API 调用失败."""
def __init__(self, message: str, code: str = "DittoError", status_code: int = 0):
self.code = code
self.status_code = status_code
super().__init__(message)
@dataclass
class DittoResult:
"""Ditto 生成结果."""
video_bytes: bytes
video_url: str = "" # 转存 OSS 后填充
elapsed_seconds: float = 0.0
rtf: float = 0.0 # 实时率(响应头 X-RTF)
frames: int = 0 # 帧数(响应头 X-Frames)
class DittoClient:
"""蚂蚁 Ditto 数字人口型 API 客户端."""
def __init__(
self,
base_url: Optional[str] = None,
default_video_url: Optional[str] = None,
max_retries: Optional[int] = None,
timeout: Optional[int] = None,
):
s = get_api_settings()
self.base_url = (base_url or s.ditto_api_base_url or "").rstrip("/")
self.default_video_url = default_video_url or s.ditto_default_video_url or ""
self.max_retries = int(max_retries if max_retries is not None else s.ditto_max_retries)
self.timeout = int(timeout if timeout is not None else s.ditto_request_timeout)
@property
def is_configured(self) -> bool:
"""配置是否完整(base_url + 默认模板视频都有值)."""
return bool(self.base_url) and bool(self.default_video_url)
def health(self) -> bool:
"""健康检查;成功返回 True,失败返回 False(不抛异常)."""
if not self.base_url:
return False
url = f"{self.base_url}/health"
try:
with httpx.Client(timeout=5.0) as client:
resp = client.get(url)
ok = resp.status_code == 200
if ok:
logger.info("[ditto] health check OK: %s", url)
else:
logger.warning("[ditto] health check status=%d: %s", resp.status_code, url)
return ok
except Exception as exc:
logger.warning("[ditto] health check failed: %s", exc)
return False
def generate(
self,
*,
audio_url: str,
script: str,
video_url: Optional[str] = None,
emo_global: int = 4,
use_script_emo: bool = True,
blend_frames: int = 6,
) -> DittoResult:
"""调用 Ditto /generate 接口,返回 MP4 字节流结果.
Raises DittoError on failure.
"""
if not self.base_url:
raise DittoError("DITTO_API_BASE_URL 未配置", code="ConfigMissing")
driver_url = video_url or self.default_video_url
if not driver_url:
raise DittoError("Ditto 人物模板视频 URL 未配置", code="ConfigMissing")
if not audio_url:
raise DittoError("audio_url 不能为空", code="InvalidParam")
if not script:
script = " "
payload = {
"video_url": driver_url,
"audio_url": audio_url,
"script": script,
"emo_global": emo_global,
"use_script_emo": use_script_emo,
"blend_frames": blend_frames,
}
url = f"{self.base_url}/generate"
last_exc: Optional[Exception] = None
for attempt in range(self.max_retries + 1):
try:
start = time.monotonic()
with httpx.Client(timeout=self.timeout, follow_redirects=True) as client:
resp = client.post(url, json=payload)
elapsed = time.monotonic() - start
if resp.status_code == 429:
wait = min(2**attempt, 30)
logger.warning(
"[ditto] GPU 繁忙 (429),%ds 后重试 (%d/%d)",
wait,
attempt + 1,
self.max_retries,
)
if attempt >= self.max_retries:
raise DittoError(
f"Ditto GPU 繁忙,重试 {self.max_retries} 次仍失败",
code="BusyRetriesExhausted",
status_code=429,
)
time.sleep(wait)
continue
if resp.status_code != 200:
_text = (resp.text or "")[:300]
logger.error(
"[ditto] generate 失败 status=%d attempt=%d body=%s",
resp.status_code,
attempt + 1,
_text,
)
if resp.status_code >= 500 and attempt < self.max_retries:
time.sleep(min(2**attempt, 15))
continue
raise DittoError(
f"Ditto 返回 {resp.status_code}: {_text}",
code="DittoAPIError",
status_code=resp.status_code,
)
video_bytes = resp.content
if not video_bytes or len(video_bytes) < 1024:
raise DittoError(
f"Ditto 返回内容异常(size={len(video_bytes) if video_bytes else 0})",
code="EmptyResponse",
)
try:
rtf = float(resp.headers.get("X-RTF", "0") or 0)
except ValueError:
rtf = 0.0
try:
frames = int(resp.headers.get("X-Frames", "0") or 0)
except ValueError:
frames = 0
try:
x_time = float(resp.headers.get("X-Time", "0") or 0)
if x_time > 0:
elapsed = x_time
except ValueError:
pass
logger.info(
"[ditto] generate 成功 size=%d rtf=%.2f frames=%d elapsed=%.1fs attempt=%d",
len(video_bytes),
rtf,
frames,
elapsed,
attempt + 1,
)
return DittoResult(
video_bytes=video_bytes,
elapsed_seconds=elapsed,
rtf=rtf,
frames=frames,
)
except DittoError:
raise
except httpx.TimeoutException as exc:
last_exc = exc
logger.warning("[ditto] 请求超时 attempt=%d err=%s", attempt + 1, exc)
if attempt < self.max_retries:
time.sleep(min(2**attempt, 15))
continue
raise DittoError(
f"Ditto 请求超时({self.timeout}s),重试耗尽",
code="Timeout",
) from exc
except Exception as exc:
last_exc = exc
logger.warning("[ditto] 请求异常 attempt=%d err=%s", attempt + 1, exc)
if attempt < self.max_retries:
time.sleep(min(2**attempt, 10))
continue
raise DittoError(f"Ditto 调用异常: {exc}", code="NetworkError") from exc
raise DittoError("Ditto 未知错误", code="Unknown") from last_exc
def generate_and_persist(
self,
*,
job_id: str,
user_id: str,
audio_url: str,
script: str,
video_url: Optional[str] = None,
) -> DittoResult:
"""调用 generate 并把 MP4 转存到自家 OSS,返回带 video_url 的结果."""
result = self.generate(audio_url=audio_url, script=script, video_url=video_url)
try:
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
storage_key = f"ditto-output/{user_id}/{job_id}.mp4"
public_url = storage.upload_file(
io.BytesIO(result.video_bytes),
storage_key,
content_type="video/mp4",
)
result.video_url = public_url
logger.info(
"[ditto] 转存 OSS 完成 job=%s key=%s",
job_id,
storage_key,
)
except Exception as exc:
logger.error("[ditto] 转存 OSS 失败 job=%s err=%s", job_id, exc, exc_info=True)
raise DittoError(f"Ditto 结果转存 OSS 失败: {exc}", code="StorageError") from exc
return result
_ditto_client_singleton: Optional[DittoClient] = None
def get_ditto_client() -> DittoClient:
"""获取 DittoClient 单例(简易工厂,便于单测 mock)."""
global _ditto_client_singleton
if _ditto_client_singleton is None:
_ditto_client_singleton = DittoClient()
return _ditto_client_singleton
+269 -161
View File
@@ -1,19 +1,16 @@
"""爆款视频 Prompt 模板默认值(v8 / v3 叙述优先重构)。
"""爆款视频 5 套 Prompt 模板默认值(#2040 核心资产)。
设计原则(灵应 2026-10-07):LLM 直接输出最终给用户看的文案,代码尽量薄。
- image_analysis v8:VLM 主交付物是自然叙述风格的 summary_markdown,结构化
字段仅保留 type/name/brand/has_person,顶层 products 改名 images;
- storyboard v3:口播台词口语化、画面描述有画面感,copy_display_markdown 是
LLM 直接写给用户看的流畅叙述文案,代码只做解析不改写;
- intent_parsing 步骤整体删除,意图理解并入 storyboard 一次调用。
模板字段与 DB 表 viral_video_prompt_templates、prompt_loader 完全对应:
name / prompt_type / version(int) / system_prompt / user_prompt_template /
example_output / is_active。
重要约定(用户明确要求):
- 所有 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 爆款”等会过时的时间表述。
@@ -22,7 +19,7 @@ GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
4. 符合广告法及平台社区规范。
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
# 负向提示(注入 storyboard / 视频生成负面词)
# 反套路化要求
NEGATIVE_RULES = """【反套路化要求】
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
@@ -30,189 +27,300 @@ NEGATIVE_RULES = """【反套路化要求】
# 输出禁用套路词(测试会检查)
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
# 文案融合三档独立指令段(storyboard 一次生成,按档位注入风格指令)
# 文案融合三档独立指令段
FUSION_INSTRUCTIONS = {
"ai_full": """【本次创作模式:AI 全权创作】
你是资深短视频编导。用户只提供了产品/门店图片,没有给出具体文案方向。请根据图片的真实观察和营销参数,自由发挥创作完整成片级方案,口播自然、画面可拍。""",
你是资深短视频编导。用户只提供了产品图片,没有给出具体文案方向。请根据图片内容和营销参数,自由发挥创作完整的爆款短视频文案。充分挖掘产品真实可见的卖点,使用爆款结构,抓人眼球。""",
"ai_polish": """【本次创作模式:AI 辅助润色】
用户已给出方向或碎碎念。以用户的意思为主,保留其所有核心信息,在此基础上润色、补衔接、优化表达,让口播更自然、画面更具体;绝不改变用户核心意思,不添加用户没提到的卖点,品牌名、价格、人名等事实原样保留。""",
你是用户的文案助理。用户已经写了草稿/关键词/碎碎念,表达了他想讲的核心意思,但表达不完整、不够吸引人。你的任务是:以用户的意思为主,保留他想表达的所有核心信息点,在此基础上润色扩写、调整语序、增加衔接、优化表达,让文案更流畅更有吸引力。绝对不能改变用户想表达的核心意思,不能把用户的观点换成相反的,不能添加用户没提到的产品卖点。用户提到的品牌名、价格、人名、具体事实必须原样保留。""",
"user_primary": """【本次创作模式:以用户原文为主】
最小化修改:只做必要的通顺、合规修正与衔接补全,用户的核心句子与事实一律不改;用户文案已经很好就直接用,不为改而改。""",
你是文案润色助手。用户已经写好了明确的文案,这是他最终想表达的内容。你的任务是最小化修改:只做必要的错别字修正、标点调整、语句通顺度优化,以及添加必要的衔接词让口播更自然。用户的核心句子、关键表述、事实信息一律不改。如果用户文案本身已经很好,直接返回,不要为了改而改。personal_brands 中的事实信息必须逐字保留。""",
}
# ── 模板1:图片多模态分析 v8(叙述优先)───────────────────────────────
_IMAGE_ANALYSIS_SYSTEM = """你是一名擅长观察和写作的品牌内容编导。面对一张真实图片,先用眼睛仔细看,再用自然、流畅、具体的中文把画面写成一段可以直接读给人听的描述。
# ── 模板1:图片多模态分析(VLM)────────────────────────────────────────
_IMAGE_ANALYSIS_SYSTEM = f"""你是电商商品视觉分析师,负责从商品图片中提取真实可见的商品信息。
## 输出格式(严格 JSON,不要输出 JSON 以外的任何内容)
{
"images": [
{
"type": "store 或 product 或 person 或 scene,四选一",
"name": "主体名称,看不出就写“未识别”",
"brand": "品牌名,看不出就留空字符串",
"has_person": false,
"summary_markdown": "用 Markdown 写成的自然叙述,这是最主要的交付物"
}
]
}
工作方式(分步骤看,不要跳步):
1. 先看整体:有哪些产品、什么场景、有没有人物。
2. 再看细节:包装文字、颜色构成、人物状态、画面质感。
3. 最后提炼卖点:只总结图片里能看到的卖点。
## summary_markdown 写作要求(最重要)
1. 写成完整、通顺的句子,像在跟朋友认真描述你看到的画面;不要用分号堆砌关键词,不要罗列“核心特征:xxx”“主色调:xxx”这类填表式标签。
2. 开头先给一句整体定性,让读者立刻明白这是什么场景、什么主体。
3. 颜色、材质、形状、部件要具体可感,写到位置和搭配;画面里出现的文字原样读出并自然融进句子,数字、规格、价格精确引用,看不清的不要编造。
4. 只写真实看到的内容,不脑补功能、疗效、销量或画面之外的信息。
5. 长度控制在 200-500 字。
{GLOBAL_CONSTRAINTS}
## 按类型组织内容
- type=store(门店/店内环境):用以下小标题分段,小标题下写连贯的句子而不是清单:
###店铺主体
###周边物品
1.家具陈设
2.商品与标识
- type=product(商品):按自然段从整体到局部描写——先说是什么、什么品牌,再写包装/外形、颜色与材质、标签文字、可见部件与规格。
- type=person(人物):描述人物身份感、姿态、穿着(上下装/颜色/款式)、动作与所处环境;用于品牌宣传时突出其精神状态。
- type=scene(纯场景/风景):描述空间或风景的构成、色彩、光线、氛围与关键物件。
请严格按下面的标签格式输出,标签名一个都不能改,不要输出任何解释,不要用代码块:
<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,海报/模特立牌/照片里的人不算。
- 一张图只描述其本身;多张图属于同一场景时可呼应,但不编造对应关系。
- 输出必须是严格 JSON,summary_markdown 是字符串,内部换行用 \\n 表示。"""
【人物属性硬性要求(has_person=true时必须遵守)】
hair/skin_tone/face_shape/outfit四项绝对禁止填“无法判断”,必须基于图片可见特征给出具体中文描述:
- hair:必须描述发型+发色,如“黑色齐肩直发”“棕色微卷中长发”“深棕色短发”
- skin_tone:必须描述肤色,如“暖调自然肤色”“白皙肤色”“小麦色”
- face_shape:必须描述脸型,如“鹅蛋脸”“圆脸”“瓜子脸”“方脸”
- outfit:必须描述可见穿着,如“米色翻领衬衫”“白色T恤”“黑色连衣裙”
即使局部被遮挡也要根据可见部分合理推断;确实看不清时按最接近的直观印象描述。
_IMAGE_ANALYSIS_USER = """请分析这张图片。
图片地址:{image_url}
OCR 辅助文字(可能为空,仅供参考,不要照抄错误识别):{ocr_text}
其他非人物属性看不到或无法判断时填“无法判断”,布尔值填false,不要留空标签。
严格按系统要求只输出 JSON。"""
【有人物场景输出参考(女性手持商品示例,必须写全10个属性,禁止省略)】
<people has_person="true" count="1" gender="女" age_range="青年" hair="黑色齐肩直发" skin_tone="暖调自然肤色" face_shape="鹅蛋脸" outfit="米色翻领衬衫" pose="正面半身,手持商品" expression="面带微笑"/>"""
_IMAGE_ANALYSIS_EXAMPLE = """{
"images": [
{
"type": "store",
"name": "御众堂门店",
"brand": "御众堂",
"has_person": false,
"summary_markdown": "###店铺主体\\n这是一家名为“御众堂”的线下门店内部,整体暖木色调……"
}
]
}"""
_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。
所属行业:{industry}
图片地址:
{image_urls}
# ── 模板2:编导级分镜 v3(意图理解 + 分镜一次完成)────────────────────
_STORYBOARD_SYSTEM = (
"""你是一名懂短视频的编导和口播文案高手。你会拿到图片的真实观察、营销目的和用户参数,请一次性完成对营销意图的理解,并产出可直接拍摄/生成的分镜脚本。不要单独输出“意图解析”,意图要直接体现在台词和分镜里。
按约定的标签格式输出分析结果。"""
## 输出格式(XML,严格按结构输出,不要输出额外解释)
<script>
<copy_display_markdown><![CDATA[直接展示给用户看的成片文案,用 Markdown 写成流畅叙述]]></copy_display_markdown>
<clips>
<clip index="1">
<time_range>0-3秒</time_range>
<voiceover>这一镜的口播台词</voiceover>
<visual>具体、有画面感的镜头描述(主体/动作/镜头运动/景别/光线)</visual>
<reference_image_index>0</reference_image_index>
</clip>
</clips>
<voiceover_script>把所有 clip 的 voiceover 连成完整口播稿</voiceover_script>
<theme>一句话主题</theme>
<negative>"""
+ NEGATIVE_RULES
+ """</negative>
</script>
_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>"""
## 写作要求
1. 口播台词:像真人面对镜头说话,短句、口语化、有停顿有情绪,开头 3 秒给出钩子;不要书面腔,不要机械报参数。
2. 画面描述:写清“观众会看到什么”,有动作、有镜头运动、有景别和光线,具体可拍;不堆砌形容词,不写无法实现的画面。
3. copy_display_markdown:直接展示给最终用户的文案,用 Markdown 写成自然、流畅、有感染力的成片成片文案,可用小标题与短句组织;不要做字段列表,不要出现“镜头一/台词:”这类制作说明。
4. 内容必须来自图片观察与用户给出的信息,不编造卖点、不夸大、不使用绝对化用语和虚假承诺。
5. reference_image_index 填本镜参考图片序号(从 0 开始),没有合适参考图填 -1。
6. 分镜数量与时长匹配总时长,节奏紧凑。"""
)
# ── 模板2:用户文案意图解析(LLM)──────────────────────────────────────
_INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案可能只是几个关键词、碎碎念或者不完整的短句,你要读懂他真正想讲什么。
_STORYBOARD_USER = """<marketing_purpose>{marketing_purpose}</marketing_purpose>
<image_analysis>
{image_summary}
</image_analysis>
<user_parameters>
<theme_hint>{theme_hint}</theme_hint>
<duration>{duration}秒</duration>
<aspect_ratio>{aspect_ratio}</aspect_ratio>
<tone>{tone}</tone>
<target_audience>{target_audience}</target_audience>
<extra_requirements>{extra_requirements}</extra_requirements>
</user_parameters>
{video_style_section}
请严格按 XML 结构输出分镜脚本。"""
{GLOBAL_CONSTRAINTS}
_STORYBOARD_EXAMPLE = """<script>
<copy_display_markdown><![CDATA[# 在御众堂,把松弛的自己一点点找回来
产后妈妈最懂那种力不从心,推开门,暖光和一杯热茶先接住了你……]]></copy_display_markdown>
<clips>
<clip index="1">
<time_range>0-3秒</time_range>
<voiceover>生完娃,是不是连照镜子的勇气都没了?</voiceover>
<visual>中近景,暖光下一位妈妈略显疲惫地看向镜中,镜头缓缓推近</visual>
<reference_image_index>0</reference_image_index>
</clip>
</clips>
<voiceover_script>生完娃,是不是连照镜子的勇气都没了?</voiceover_script>
<theme>产后妈妈走进御众堂重拾状态</theme>
<negative>模糊、畸变、夸大疗效、绝对化用语</negative>
</script>"""
请严格按下面的标签格式输出,不要解释,不要用代码块:
<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> 标签;没有就输出空标签。"""
# ── 模板3:文案审核(合规/质量门禁)───────────────────────────────────
_REVIEW_SYSTEM = """你是一名短视频广告合规审核与文案优化专家。审核待审文案:
1) 广告法与平台合规(绝对化用语、虚假承诺、医疗功效宣称、导流违规);
2) 卖点是否聚焦、逻辑是否通顺、口播是否自然;
3) 是否有机械堆砌、书面腔、标签化表述。
_INTENT_USER = """用户原始文案:{user_copy_text}
所属行业:{industry}
图片分析结果(供参考):
{image_analysis}
只输出 XML,结构:
<review>
<passed>true 或 false</passed>
<issues>
<issue>
<severity>high 或 medium 或 low</severity>
<field>问题所在位置/字段</field>
<problem>具体问题</problem>
<suggestion>可直接替换的修改</suggestion>
</issue>
</issues>
<rewrite>整体重写后的合规流畅版本(无问题时留空)</rewrite>
</review>
没有问题时 issues 留空、passed 为 true、rewrite 留空。"""
请理解用户意图,按标签格式输出。"""
_REVIEW_USER = """<fusion_text>
{fusion_text}
</fusion_text>
_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 = """你负责为短视频生成营销文案。请按思维链分步完成:先定人设和目标客户,再找卖点,再搭结构,再安排情绪,最后写行动号召,不要一步到位乱写。
_REVIEW_EXAMPLE = """<review>
<passed>false</passed>
<issues>
<issue>
<severity>high</severity>
<field>opening</field>
<problem>使用绝对化用语“全网第一”</problem>
<suggestion>改为“很多老客户回购的一款”</suggestion>
</issue>
</issues>
<rewrite>……</rewrite>
</review>"""
{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": "图片多模态分析 v8",
"name": "图片多模态分析",
"prompt_type": "image_analysis",
"version": 8,
"version": TEMPLATE_VERSION,
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
"user_prompt_template": _IMAGE_ANALYSIS_USER,
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
"is_active": True,
},
{
"name": "编导级分镜 v3",
"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": 3,
"version": TEMPLATE_VERSION,
"system_prompt": _STORYBOARD_SYSTEM,
"user_prompt_template": _STORYBOARD_USER,
"example_output": _STORYBOARD_EXAMPLE,
@@ -221,7 +329,7 @@ DEFAULT_TEMPLATES: list[dict] = [
{
"name": "文案审核",
"prompt_type": "review",
"version": 1,
"version": TEMPLATE_VERSION,
"system_prompt": _REVIEW_SYSTEM,
"user_prompt_template": _REVIEW_USER,
"example_output": _REVIEW_EXAMPLE,
+3 -28
View File
@@ -53,19 +53,11 @@ _LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
class Reviewer:
# markdown展示字段不参与合规审核(避免格式字符误判)
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
def __init__(self, client=None):
if client is None:
try:
from packages.shared.ai_router import ai_router
from packages.shared.ai_client import get_doubao_client
client = ai_router.get_llm_client("copy_review")
except Exception:
from packages.shared.ai_client import get_doubao_client
client = get_doubao_client()
client = get_doubao_client()
self.client = client
# ── 审核 ────────────────────────────────────────────────────────────
@@ -73,9 +65,8 @@ class Reviewer:
local = self._rule_check(fusion, intent, fusion_level)
llm_result = self._llm_review(fusion, intent, fusion_level)
if llm_result is None:
# LLM审核失败(超时/网络错误等),降级放行,不阻断渲染
return ReviewResult(
passed=True,
passed=not local,
issues=local,
rewrite_suggestions=[],
raw="",
@@ -90,17 +81,6 @@ class Reviewer:
)
def _llm_review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> Optional[ReviewResult]:
try:
return self._llm_review_inner(fusion, intent, fusion_level)
except Exception as e:
import logging
logging.getLogger(__name__).warning("[Reviewer] LLM审核调用异常,降级放行: %s", e)
return None
def _llm_review_inner(
self, fusion: FusionResult, intent: IntentResult, fusion_level: str
) -> Optional[ReviewResult]:
template = get_template("review")
system = render_system_prompt(template)
user = render_user_prompt(
@@ -116,7 +96,6 @@ class Reviewer:
],
temperature=0.2,
max_tokens=1024,
timeout=25,
)
if not raw:
return None
@@ -263,7 +242,6 @@ class Reviewer:
],
temperature=0.5,
max_tokens=2048,
timeout=25,
)
if not raw:
return self._rule_fix(fusion, review)
@@ -318,13 +296,10 @@ class Reviewer:
@staticmethod
def _fusion_text(fusion: FusionResult) -> str:
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
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)
# 过滤掉markdown展示字段,避免格式字符被误判
parts = [p for p in parts if not any(mk in p for mk in _MARKDOWN_FIELDS)]
return "\n".join(p for p in parts if p)
@staticmethod
@@ -13,13 +13,6 @@ 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*(?:"([^"]*)"|'([^']*)')""")
_CDATA_RE = re.compile(r"^<!\[CDATA\[(.*)\]\]>$", re.DOTALL)
def _strip_cdata(s: str) -> str:
"""剥离 LLM 可能照抄示例输出的 ``<![CDATA[...]]>`` 包裹层。"""
m = _CDATA_RE.match(s.strip())
return m.group(1) if m else s
def parse_attributes(raw: str) -> dict[str, str]:
@@ -65,7 +58,6 @@ def parse_tags(text: Optional[str]) -> list[dict]:
if stack[idx]["tag"] == tag:
node = stack[idx]
node["text"] = unescape(text[node["_start"] : token.start()].strip())
node["text"] = _strip_cdata(node["text"])
node.pop("_start", None)
del stack[idx:]
break
@@ -73,7 +65,6 @@ def parse_tags(text: Optional[str]) -> list[dict]:
for node in stack:
if "_start" in node:
node["text"] = unescape(text[node["_start"] :].strip())
node["text"] = _strip_cdata(node["text"])
node.pop("_start", None)
return results
+35 -56
View File
@@ -80,46 +80,54 @@ class SharedSettings(BaseSettings):
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
cosyvoice_api_key: str = ""
cosyvoice_base_url: str = ""
cosyvoice_model: str = ""
cosyvoice_voice: str = "longxiaochun_v3"
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
cosyvoice_model: str = "cosyvoice-v3-flash"
cosyvoice_voice: str = "longxiaochun_v3" # 默认音色(v3 系列系统音色带 _v3 后缀)
cosyvoice_sample_rate: int = 22050
cosyvoice_format: str = "mp3"
cosyvoice_format: str = "mp3" # 输出格式:mp3/wav/pcm
# 音色克隆模型名(固定为 voice-enrollment)
cosyvoice_clone_model: str = ""
cosyvoice_clone_model: str = "voice-enrollment"
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
# AI模型路由化:model/base_url 默认值清空,由 DB ai_models/ai_capability_configs 配置驱动。
# 环境变量仍可覆盖(兼容旧部署);无任何配置时 ai_router fallback 提供最终默认值。
doubao_api_key: str = ""
doubao_model: str = ""
doubao_fast_model: str = ""
doubao_base_url: str = ""
doubao_timeout: int = 45
doubao_max_retries: int = 3
doubao_vision_model: str = ""
doubao_vision_lite_model: str = ""
doubao_vision_use_lite: bool = True
doubao_embedding_model: str = ""
doubao_video_model: str = ""
doubao_video_timeout: int = 600
doubao_video_poll_interval: int = 10
doubao_image_model: str = ""
doubao_image_size: str = "1K"
doubao_image_timeout: int = 60
doubao_trust_chain_enabled: bool = True
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 = 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 = ""
dashscope_video_timeout: int = 900
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 = ""
mediakit_base_url: str = ""
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
mediakit_timeout: int = 60
mediakit_cover_enabled: bool = False
mediakit_cover_enabled: bool = False # 封面抽帧是否走MediaKit(默认false走本地ffmpeg+cv2,<2s完成)
# ── 积分/会员系统 (#1895) ────────────────────────────────────────────
# 积分系统总开关(产品要求 #1895:暂停积分系统但保留全部代码/表/接口)。
@@ -173,35 +181,6 @@ class SharedSettings(BaseSettings):
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
gpu_worker_stale_seconds: int = 300
# ── Ditto 蚂蚁数字人口型 API(#2076)─────────────────────────────────
# 是否优先使用 Ditto(蚂蚁数字人,替代 MuseTalk)。开关开启且 base_url 配置
# 非空时,对口型任务优先走 Ditto;失败后回退 MuseTalk/MediaKit。
use_ditto_lipsync: bool = Field(
default=False,
validation_alias=AliasChoices("USE_DITTO_LIPSYNC", "use_ditto_lipsync"),
)
# Ditto FastAPI 内网地址(Tailscale),如 http://100.x.x.x:8000
ditto_api_base_url: str = Field(
default="",
validation_alias=AliasChoices("DITTO_API_BASE_URL", "ditto_api_base_url"),
)
# 默认人物模板视频 URL(正面 5-10 秒循环、光线均匀、半身)。Ditto 模式下忽略
# 用户上传的驱动视频/图片,统一用该模板;后续可扩展为多模板让用户选择。
ditto_default_video_url: str = Field(
default="",
validation_alias=AliasChoices("DITTO_DEFAULT_VIDEO_URL", "ditto_default_video_url"),
)
# 429 GPU 繁忙时指数退避最大重试次数
ditto_max_retries: int = Field(
default=3,
validation_alias=AliasChoices("DITTO_MAX_RETRIES", "ditto_max_retries"),
)
# Ditto 单次请求超时(秒):数字人半身视频推理通常 30-120s
ditto_request_timeout: int = Field(
default=300,
validation_alias=AliasChoices("DITTO_REQUEST_TIMEOUT", "ditto_request_timeout"),
)
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
enable_gpu_encode: bool = Field(
-376
View File
@@ -1,376 +0,0 @@
"""功能计费配置服务:从 feature_pricing_configs 读配置,300 秒 TTL 内存缓存。
配置表由 xiaoxia-admin 侧维护(同库 PostgreSQL),本服务只读。
DB 不可用 / 表不存在 / 无数据时自动回落到内置兜底配置,保证业务不崩。
计费公式:最终积分 = (动态成本 + 固定成本) × 利润系数,price_cap 封顶。
启用条件:全局 points_enabled 总开关 AND 功能 is_enabled 同时为 true。
"""
from __future__ import annotations
import json
import logging
import threading
import time
from dataclasses import dataclass, field
from typing import Optional
import sqlalchemy as sa
from packages.adapters.sqlalchemy_impl import session as _session_mod
logger = logging.getLogger(__name__)
CACHE_TTL_SECONDS = 300.0
# ── 爆款视频兜底模型单价(与旧硬编码表/现状一致;DB 不可用时使用) ───────
# 结构:models[model_key][resolution]["true"/"false"] = 单价
# token 模式:元/百万输出 tokens;per_second 模式:元/秒
# 注意:仅 seedance-2.5 配置 true(图生视频)单价;其余模型只有 false,
# 精确 key 缺失时由 points_rules 回落到 seedance-2.5/false(与旧现状一致)。
_FALLBACK_VIRAL_MODEL_PRICING: dict = {
"seedance-2.5": {
"480p": {"false": 70.0, "true": 42.0},
"720p": {"false": 70.0, "true": 42.0},
"1080p": {"false": 77.0, "true": 46.0},
},
"seedance-2.0": {
"480p": {"false": 46.0},
"720p": {"false": 46.0},
"1080p": {"false": 51.0},
"4k": {"false": 80.0},
},
"seedance-2.0-fast": {
"480p": {"false": 28.0},
"720p": {"false": 28.0},
},
"seedance-2.0-mini": {
"480p": {"false": 9.2},
"720p": {"false": 9.2},
},
"wan-3.0": {
"480p": {"false": 0.3},
"720p": {"false": 0.6},
"1080p": {"false": 1.2},
},
}
@dataclass
class FeatureConfig:
"""功能计费配置快照。"""
feature_key: str
name: str = ""
emoji: str = ""
is_enabled: bool = False
fixed_cost: float = 0.0
profit_multiplier: float = 1.0
dynamic_unit_cost: float = 0.0
billing_mode: str = "model_based"
price_cap: float = 0.0
model_pricing: dict = field(default_factory=dict)
description: str = ""
# ── 进程内缓存:(loaded_monotonic, {feature_key: FeatureConfig}) ──────────
_lock = threading.Lock()
_cache: Optional[tuple[float, dict[str, FeatureConfig]]] = None
def _fallback_configs() -> dict[str, FeatureConfig]:
"""内置兜底配置:爆款启用(与现状一致),其余两个关闭。"""
return {
"viral_video": FeatureConfig(
feature_key="viral_video",
name="爆款视频",
emoji="🎬",
is_enabled=True,
fixed_cost=0.15,
profit_multiplier=1.3,
dynamic_unit_cost=0.0,
billing_mode="model_based",
price_cap=0.0,
model_pricing=json.loads(json.dumps(_FALLBACK_VIRAL_MODEL_PRICING)),
description="爆款视频动态定价(兜底配置)",
),
"lip_sync": FeatureConfig(
feature_key="lip_sync",
name="对口型",
emoji="🎙️",
is_enabled=False,
fixed_cost=0.0,
profit_multiplier=1.0,
dynamic_unit_cost=0.0,
billing_mode="per_second",
price_cap=0.0,
description="对口型计费(兜底配置,默认关闭)",
),
"smart_edit": FeatureConfig(
feature_key="smart_edit",
name="智能剪辑",
emoji="✂️",
is_enabled=False,
fixed_cost=0.0,
profit_multiplier=1.0,
dynamic_unit_cost=0.0,
billing_mode="model_based",
price_cap=0.0,
description="智能剪辑固定价计费(兜底配置,默认关闭)",
),
}
_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://")
if url.startswith("postgresql://"):
url = url.replace("postgresql://", "postgresql+psycopg://")
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 _parse_model_pricing(raw) -> dict:
"""解析 model_pricing_json(Text JSON),空/失败 → {}。"""
if raw is None:
return {}
if isinstance(raw, dict):
return raw
text = str(raw).strip()
if not text:
return {}
try:
data = json.loads(text)
except (ValueError, TypeError):
logger.warning("model_pricing_json 解析失败,按空配置处理: %r", text[:200])
return {}
return data if isinstance(data, dict) else {}
def _to_float(value, default: float = 0.0) -> float:
try:
if value is None:
return default
return float(value)
except (TypeError, ValueError):
return default
def _load_all() -> dict[str, FeatureConfig]:
"""SELECT * FROM feature_pricing_configs,返回 {feature_key: FeatureConfig}。
表不存在 / DB 异常由调用方捕获并回落兜底配置。
"""
session = None
try:
session = _get_session()
if session is None:
raise RuntimeError("no db session available")
sql = sa.text("""
SELECT feature_key, name, emoji, is_enabled, fixed_cost,
profit_multiplier, dynamic_unit_cost, billing_mode,
price_cap, model_pricing_json, description
FROM feature_pricing_configs
""")
rows = session.execute(sql).mappings().all()
configs: dict[str, FeatureConfig] = {}
for row in rows:
key = str(row["feature_key"] or "").strip()
if not key:
continue
configs[key] = FeatureConfig(
feature_key=key,
name=str(row["name"] or key),
emoji=str(row["emoji"] or ""),
is_enabled=bool(row["is_enabled"]),
fixed_cost=_to_float(row["fixed_cost"]),
profit_multiplier=_to_float(row["profit_multiplier"], 1.0),
dynamic_unit_cost=_to_float(row["dynamic_unit_cost"]),
billing_mode=str(row["billing_mode"] or "model_based"),
price_cap=_to_float(row["price_cap"]),
model_pricing=_parse_model_pricing(row["model_pricing_json"]),
description=str(row["description"] or ""),
)
return configs
finally:
if session is not None:
try:
session.close()
except Exception: # noqa: BLE001
pass
def _get_cache() -> dict[str, FeatureConfig]:
"""TTL 内返回缓存,否则重新 load;DB 异常/表不存在时返回内置兜底配置。"""
global _cache
now = time.monotonic()
with _lock:
if _cache is not None and now - _cache[0] < CACHE_TTL_SECONDS:
return _cache[1]
try:
loaded = _load_all()
except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落
logger.info("feature_pricing_configs 读取失败,使用内置兜底配置", exc_info=True)
return _fallback_configs()
# DB 可用但表为空:同样回落兜底(保证爆款现状不被改变)
if not loaded:
fallback = _fallback_configs()
with _lock:
_cache = (now, fallback)
return fallback
# 以兜底为底(DB 未配置的 feature_key 仍有兜底),DB 行覆盖
merged = _fallback_configs()
merged.update(loaded)
with _lock:
_cache = (now, merged)
return merged
def get_feature_config(feature_key: str) -> Optional[FeatureConfig]:
"""获取指定功能配置,未知 key 返回 None。"""
key = str(feature_key or "").strip()
if not key:
return None
return _get_cache().get(key)
def _global_points_enabled() -> bool:
"""全局积分总开关(兼容 api / worker 运行时),取不到时默认关闭。"""
try:
from packages.shared import get_shared_settings
return bool(get_shared_settings().points_enabled)
except Exception: # noqa: BLE001
pass
try:
from app.config import settings
return bool(getattr(settings, "points_enabled", False))
except Exception: # noqa: BLE001
return False
def is_feature_enabled(feature_key: str) -> bool:
"""功能是否启用并扣费:全局 points_enabled AND 功能 is_enabled。"""
cfg = get_feature_config(feature_key)
if cfg is None:
return False
return bool(cfg.is_enabled) and _global_points_enabled()
def calculate_price(feature_key: str, dynamic_cost: float = 0.0) -> tuple[float, dict]:
"""按公式计算最终积分并返回明细。
price = (dynamic_cost + fixed_cost) × profit_multiplier
price_cap > 0 时封顶(取 min)。
功能未启用 → (0.0, breakdown{is_enabled: False, charged: False})。
"""
cfg = get_feature_config(feature_key)
dynamic = max(0.0, _to_float(dynamic_cost))
if cfg is None or not cfg.is_enabled:
return 0.0, {
"feature_key": feature_key,
"is_enabled": False,
"charged": False,
"dynamic_cost": dynamic,
"fixed_cost": 0.0,
"profit_multiplier": 1.0,
"price_cap": 0.0,
"final_price": 0.0,
}
fixed = max(0.0, cfg.fixed_cost)
multiplier = cfg.profit_multiplier if cfg.profit_multiplier > 0 else 1.0
raw_price = (dynamic + fixed) * multiplier
cap = cfg.price_cap if cfg.price_cap and cfg.price_cap > 0 else 0.0
final_price = min(raw_price, cap) if cap else raw_price
final_price = round(float(final_price), 2)
breakdown = {
"feature_key": cfg.feature_key,
"is_enabled": True,
"charged": True,
"dynamic_cost": round(dynamic, 4),
"fixed_cost": float(fixed),
"profit_multiplier": float(multiplier),
"price_cap": float(cap),
"raw_price": round(float(raw_price), 4),
"final_price": final_price,
}
return final_price, breakdown
def lookup_model_price(
model_pricing: dict,
model_key: str,
resolution: str,
has_video_input: bool,
) -> Optional[float]:
"""从 model_pricing dict 取模型单价,兼容两种常见 JSON 结构。
1. 嵌套:{model: {resolution: {"true"/"false": price}}}
(内层 bool key 也兼容直接 bool / 省略)
2. 扁平:{"model|resolution|true_or_false": price}
(分隔符支持 | / : / , / 空格;bool 段可省略)
取不到返回 None。
"""
if not isinstance(model_pricing, dict):
return None
model = str(model_key or "").strip()
res = str(resolution or "").strip()
flag = "true" if has_video_input else "false"
# 1. 嵌套
model_node = model_pricing.get(model)
if isinstance(model_node, dict):
res_node = model_node.get(res)
if isinstance(res_node, dict):
# 精确 bool key 命中才返回;不做“只有一个值就取”的模糊匹配
# (否则缺失 true 时会错误地取到 false 价,破坏旧版回落规则)
if flag in res_node:
return _to_float(res_node[flag]) if res_node[flag] is not None else None
if has_video_input in res_node:
val = res_node[has_video_input]
return _to_float(val) if val is not None else None
elif isinstance(res_node, (int, float)):
return float(res_node)
# 2. 扁平
for sep in ("|", ":", ",", " "):
for key in (
f"{model}{sep}{res}{sep}{flag}",
f"{model}{sep}{res}",
):
if key in model_pricing:
value = model_pricing[key]
return _to_float(value) if value is not None else None
return None
def refresh_feature_configs() -> None:
"""清空缓存(下次读取重新 load DB;测试/admin 改配置后可手动调)。"""
global _cache
with _lock:
_cache = None
+14 -77
View File
@@ -2,21 +2,17 @@
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs,
见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/
固定成本/利润系数/封顶,DB 不可用时回落兜底配置。
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
"""
from __future__ import annotations
import math
from packages.domain import feature_pricing_service
# ============ 爆款视频动态定价 ============
# 单价/固定成本/利润系数已 DB 化(feature_pricing_configs,feature_key=viral_video),
# 由 feature_pricing_service 读取(300s 缓存),DB 不可用时回落内置兜底配置。
# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。
# ============ 爆款视频动态定价 (#2151) ============
# key = (model_id, resolution, has_video_input),单位:
# - billing_mode=token: 元/百万tokens(输出)
# - billing_mode=per_second: 元/秒(视频时长)
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.5", "480p", False): 70.0,
("seedance-2.5", "720p", False): 70.0,
@@ -37,9 +33,9 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("wan-3.0", "1080p", False): 1.2,
}
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器(兜底默认值)
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
VIRAL_VIDEO_FIXED_COST = 0.15
# 利润系数(兜底默认值)
# 利润系数
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
# Seedance 输出帧率
VIRAL_VIDEO_FPS = 24
@@ -226,22 +222,17 @@ def calculate_viral_video_credits_with_breakdown(
) -> tuple[float, dict]:
"""计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。
单价/固定成本/利润系数/封顶从 feature_pricing_configs(viral_video)读取;
DB 不可用时回落与现状一致的内置兜底配置。
公式:
tokens = duration * width * height * fps / 1024
video_cost = tokens / 1_000_000 * model_token_price
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
price_cap > 0 时封顶取 min
若传入 actual_tokens 则用它替代计算值。
Returns:
(credits, breakdown) 二元组:
- credits: 四舍五入保留两位小数的最终积分
- breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier /
model_price / width / height / fps / feature_enabled / charged / price_cap
字段,便于前端展示计费明细。功能关闭时 credits=0、charged=False。
model_price / width / height / fps 字段,便于前端展示计费明细。
"""
w = max(1, int(width or 1))
h = max(1, int(height or 1))
@@ -251,36 +242,11 @@ def calculate_viral_video_credits_with_breakdown(
cfg = get_viral_video_model_config(prefix)
res_key = _infer_resolution_key(w, h)
billing = cfg.get("billing_mode", "token")
dur = max(1, int(duration_seconds or 15))
# ── 从 DB 配置(兜底内置)取计费参数 ──
feature_cfg = feature_pricing_service.get_feature_config("viral_video")
# 注意:此处 feature_enabled 只表示“功能自身开关”,不并入全局 points_enabled
# 总开关(保持与旧版计费函数行为一致:价格照常计算)。全局总开关由业务层
# (route/worker)通过 feature_pricing_service.is_feature_enabled 统一把关。
feature_enabled = bool(feature_cfg.is_enabled) if feature_cfg is not None else True
model_pricing = feature_cfg.model_pricing if feature_cfg is not None else {}
fixed_cost = float(feature_cfg.fixed_cost) if feature_cfg is not None else float(VIRAL_VIDEO_FIXED_COST)
multiplier = (
float(feature_cfg.profit_multiplier)
if feature_cfg is not None and feature_cfg.profit_multiplier > 0
else float(VIRAL_VIDEO_PROFIT_MULTIPLIER)
)
price_cap = float(feature_cfg.price_cap) if feature_cfg is not None else 0.0
# 单价:优先配置 dict;复刻旧版回落规则——精确 key 取不到时,回落
# seedance-2.5 同分辨率 False 单价;最终兜底 70.0。
price = feature_pricing_service.lookup_model_price(model_pricing, prefix, res_key, bool(has_video_input))
if price is None:
# 配置表未命中:先尝试配置里的 seedance-2.5/False
if prefix != "seedance-2.5" or bool(has_video_input):
price = feature_pricing_service.lookup_model_price(model_pricing, "seedance-2.5", res_key, False)
if price is None:
key = (prefix, res_key, bool(has_video_input))
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
key = (prefix, res_key, bool(has_video_input))
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
if price is None:
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
dur = max(1, int(duration_seconds or 15))
if billing == "per_second":
tokens = 0.0
video_cost = dur * float(price)
@@ -293,40 +259,13 @@ def calculate_viral_video_credits_with_breakdown(
video_cost = tokens / 1_000_000.0 * float(price)
billing_unit = "token"
if not feature_enabled:
# 功能关闭(is_enabled=false 或全局 points 关闭):不扣费,明细照旧返回
credits = 0.0
raw_total = (video_cost + fixed_cost) * multiplier
breakdown = {
"tokens": float(tokens),
"video_cost": float(video_cost),
"fixed_cost": float(fixed_cost),
"profit_multiplier": float(multiplier),
"price_cap": float(price_cap or 0.0),
"model_price": float(price),
"model_key": prefix,
"billing_mode": billing,
"billing_unit": billing_unit,
"width": int(w),
"height": int(h),
"fps": int(effective_fps),
"duration": dur,
"feature_enabled": False,
"charged": False,
"raw_price": round(float(raw_total), 4),
}
return credits, breakdown
total = (video_cost + fixed_cost) * multiplier
if price_cap and price_cap > 0:
total = min(total, price_cap)
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
credits = round(float(total), 2)
breakdown = {
"tokens": float(tokens),
"video_cost": float(video_cost),
"fixed_cost": float(fixed_cost),
"profit_multiplier": float(multiplier),
"price_cap": float(price_cap or 0.0),
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
"model_price": float(price),
"model_key": prefix,
"billing_mode": billing,
@@ -335,8 +274,6 @@ def calculate_viral_video_credits_with_breakdown(
"height": int(h),
"fps": int(effective_fps),
"duration": dur,
"feature_enabled": True,
"charged": True,
}
return credits, breakdown
+4 -42
View File
@@ -191,45 +191,13 @@ class ViralVideoJob:
self.updated_at = datetime.now(timezone.utc)
def resume_from_image_analyzed(self, **kwargs) -> None:
"""阶段2入口:允许从 IMAGE_ANALYZED/PENDING 首次进入,也允许从 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。
重新生成时清空上一轮文案产物(copy_result/intent_result/storyboard/generated_copy_text),
并重置 completed_at/result_video_url/error_msg,确保前端轮询能看到新的阶段2进度。
"""
_allowed = (
ViralVideoStatus.IMAGE_ANALYZED,
ViralVideoStatus.PENDING,
ViralVideoStatus.COPY_GENERATED,
ViralVideoStatus.COMPLETED,
ViralVideoStatus.FAILED,
)
if self.status not in _allowed:
if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING):
raise ValueError(f"Cannot resume from {self.status} to copy-gen")
_is_regen = self.status in (
ViralVideoStatus.COPY_GENERATED,
ViralVideoStatus.COMPLETED,
ViralVideoStatus.FAILED,
)
for k, v in kwargs.items():
if hasattr(self, k) and v not in (None, "", []):
setattr(self, k, v)
if _is_regen:
# 清空上一轮文案/视频产物,避免前端拿到旧数据
self.intent_result = None
self.copy_result = None
self.storyboard = None
self.generated_copy_text = ""
self.result_video_url = ""
self.current_stage = ""
self.phase_message = ""
self.error_msg = ""
self.completed_at = None
self.heartbeat_at = None
_now = datetime.now(timezone.utc)
self.started_at = _now
self.heartbeat_at = _now
self.status = ViralVideoStatus.RUNNING
self.updated_at = _now
self.updated_at = datetime.now(timezone.utc)
def resume_from_copy_generated(self, edited_copy: str | None = None) -> None:
"""阶段2->阶段3:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
@@ -238,20 +206,14 @@ class ViralVideoJob:
if edited_copy and isinstance(self.copy_result, dict):
self.copy_result = {**self.copy_result, "voiceover_script": edited_copy}
self.generated_copy_text = edited_copy
_now = datetime.now(timezone.utc)
self.started_at = _now
self.heartbeat_at = _now
self.status = ViralVideoStatus.RUNNING
self.updated_at = _now
self.updated_at = datetime.now(timezone.utc)
def resume_from_confirm(self) -> None:
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
raise ValueError(f"Cannot resume from {self.status}")
_now = datetime.now(timezone.utc)
self.started_at = _now
self.heartbeat_at = _now
self.status = ViralVideoStatus.RUNNING
self.updated_at = _now
self.updated_at = datetime.now(timezone.utc)
def mark_completed(self, video_url: str) -> None:
self.status = ViralVideoStatus.COMPLETED
+28 -108
View File
@@ -170,35 +170,14 @@ class DoubaoClient:
未配置 API Key 时 is_available 为 False,调用方应降级处理。
"""
def __init__(
self,
api_key: str = "",
base_url: str = "",
model: str = "",
timeout: int = 0,
max_retries: int = 0,
max_tokens: int | None = None,
temperature: float | None = None,
extra_params: dict | None = None,
provider: str = "volcengine",
) -> None:
def __init__(self) -> None:
settings = get_shared_settings()
self.provider: str = provider
if provider == "dashscope":
self.api_key: str = api_key or getattr(settings, "dashscope_api_key", "")
self.model: str = model or getattr(settings, "dashscope_model", "")
self.base_url: str = (base_url or getattr(settings, "dashscope_base_url", "")).rstrip("/")
else: # volcengine (default)
self.api_key = api_key or settings.doubao_api_key
self.model = model or settings.doubao_model
self.base_url = (base_url or settings.doubao_base_url).rstrip("/")
self.timeout: int = timeout or settings.doubao_timeout
self.max_retries: int = max_retries or settings.doubao_max_retries
self.max_tokens: int | None = max_tokens
self.temperature: float | None = temperature
self.extra_params: dict = extra_params or {}
self.api_key: str = settings.doubao_api_key
self.model: str = settings.doubao_model
self.base_url: str = settings.doubao_base_url.rstrip("/")
self.timeout: int = settings.doubao_timeout
self.max_retries: int = settings.doubao_max_retries
self.vision_model: str = settings.doubao_vision_model
self.last_finish_reason: str = ""
self.vision_lite_model: str = settings.doubao_vision_lite_model
self.fast_model: str = settings.doubao_fast_model
self.embedding_model: str = settings.doubao_embedding_model
@@ -211,13 +190,6 @@ class DoubaoClient:
# 最近一次图片生成的详细错误,供上层读取
self.last_image_error: dict = {}
def _resolve_timeout(self, timeout) -> "httpx.Timeout":
"""将整数超时转为 httpx.Timeout,区分 connect/read/write/pool,避免 read 卡到 TCP 120s 默认值."""
if isinstance(timeout, httpx.Timeout):
return timeout
t = int(timeout) if timeout else 60
return httpx.Timeout(connect=10, read=max(t, 10), write=10, pool=5)
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
if not self.is_available or not text or not text.strip():
@@ -268,17 +240,16 @@ class DoubaoClient:
self,
messages: list[dict[str, str]],
temperature: float = 0.7,
max_tokens: int | None = None,
max_tokens: int = 1024,
model: str | None = None,
timeout: int | None = None,
**kwargs,
) -> Optional[str]:
"""调用 Chat Completion 接口.
Args:
messages: 对话消息列表,[{"role": "user"/"system"/"assistant", "content": "..."}]
temperature: 采样温度,0-2,默认0.7
max_tokens: 最大生成token数,默认 None(使用实例 self.max_tokens DB 配置,兜底 1024)
max_tokens: 最大生成token数,默认1024
Returns:
模型返回的文本内容,失败返回 None
@@ -291,24 +262,18 @@ class DoubaoClient:
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
payload: dict[str, Any] = {
"model": model or self.model,
"messages": messages,
"temperature": temperature,
"max_tokens": effective_max_tokens,
"max_tokens": max_tokens,
}
# 合并实例级额外参数和调用方传入的额外参数
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
last_error: Optional[Exception] = None
_t0 = time.time()
for attempt in range(self.max_retries + 1):
try:
_req_timeout = self._resolve_timeout(timeout if timeout is not None else self.timeout)
_req_timeout = timeout if timeout is not None else self.timeout
response = httpx.post(
url,
headers=headers,
@@ -317,33 +282,15 @@ class DoubaoClient:
)
response.raise_for_status()
data = response.json()
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
if finish_reason == "length" and attempt < self.max_retries:
# 输出被 max_tokens 截断:2.0x 扩容后重试(计入 max_retries,不额外增加)
old_max = int(payload["max_tokens"])
new_max = int(old_max * 2)
payload["max_tokens"] = new_max
wait = 0.5 * (2**attempt)
logger.warning(
"输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
old_max,
new_max,
attempt + 1,
self.max_retries + 1,
)
time.sleep(wait)
continue
content = data["choices"][0]["message"]["content"]
self.last_finish_reason = finish_reason
_elapsed = time.time() - _t0
logger.info(
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%s",
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%d",
payload.get("model"),
data.get("usage", {}).get("prompt_tokens", 0),
data.get("usage", {}).get("completion_tokens", 0),
_elapsed,
attempt + 1,
getattr(_req_timeout, "read", _req_timeout),
)
return content.strip()
except Exception as e:
@@ -367,21 +314,20 @@ class DoubaoClient:
self,
messages: list[dict],
images: list[str] | None = None,
max_tokens: int | None = None,
max_tokens: int = 2048,
temperature: float = 0.3,
timeout: int | None = None,
model: str | None = None,
**kwargs,
) -> Optional[str]:
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
将 images 附加到最后一条 user message 的 content 中,
使用构造函数传入的 self.model(DB capability 绑定的视觉模型,默认 qwen-vl-plus)。
使用 vision_model(默认 doubao-1-5-vision-pro-250328)。
Args:
messages: 对话消息列表。最后一条 user message 会被注入图片内容。
images: 图片列表,支持 base64 data URI 或 HTTP(S) URL。
max_tokens: 最大生成 token 数,默认 None(使用实例 self.max_tokens DB 配置,兜底 2048)。
max_tokens: 最大生成 token 数,默认 2048。
temperature: 采样温度,默认 0.3(视觉任务偏低更稳定)。
timeout: 单次请求超时秒数,不传则使用默认 self.timeout。
@@ -421,19 +367,14 @@ class DoubaoClient:
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 2048)
payload: dict[str, Any] = {
"model": model or self.model,
"model": model or self.vision_model,
"messages": vision_messages,
"temperature": temperature,
"max_tokens": effective_max_tokens,
"max_tokens": max_tokens,
}
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
req_timeout = self._resolve_timeout(timeout or self.timeout)
req_timeout = timeout or self.timeout
last_error: Optional[Exception] = None
_t0 = time.time()
for attempt in range(self.max_retries + 1):
@@ -446,24 +387,7 @@ class DoubaoClient:
)
response.raise_for_status()
data = response.json()
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
if finish_reason == "length" and attempt < self.max_retries:
# 视觉输出被 max_tokens 截断:2.0x 扩容后重试(计入 max_retries)
old_max = int(payload["max_tokens"])
new_max = int(old_max * 2)
payload["max_tokens"] = new_max
wait = 0.5 * (2**attempt)
logger.warning(
"视觉输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
old_max,
new_max,
attempt + 1,
self.max_retries + 1,
)
time.sleep(wait)
continue
content = data["choices"][0]["message"]["content"]
self.last_finish_reason = finish_reason
_elapsed = time.time() - _t0
logger.info(
"[doubao] vision_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d",
@@ -665,32 +589,28 @@ class DoubaoClient:
# 信任链只作用于 doubao provider;DashScope(Wan) 保持原行为。
trust_chain_applied = False
if provider == "doubao" and getattr(self, "trust_chain_enabled", True) and pre_trusted_images:
# #2220: 稀疏列表模式——pre_trusted_images 与 raw_portrait_urls 等长,
# None 位保留原图,非 None 位用 AI 人像替换。
raw_portrait_urls: list[str] = []
if image_url:
raw_portrait_urls.append(image_url)
for u in ref_imgs:
if u not in raw_portrait_urls:
raw_portrait_urls.append(u)
_n_trusted = sum(1 for _x in pre_trusted_images if _x)
if _n_trusted >= 1 and len(pre_trusted_images) >= len(raw_portrait_urls):
merged: list[str] = []
for _i, _orig in enumerate(raw_portrait_urls):
_ai = pre_trusted_images[_i] if _i < len(pre_trusted_images) else None
merged.append(str(_ai) if _ai else _orig)
trusted_urls: list[str] = []
if len(pre_trusted_images) >= 1:
trusted_urls = list(pre_trusted_images)
trust_chain_applied = True
logger.info(
"[trust-chain] 稀疏替换 %d/%d 张为AI人像(场景/商品图保留原图),走reference_image模式",
_n_trusted,
"[trust-chain] 使用预热t2i结果 %d 张,替换原参考图走 reference_image 模式(原n=%d)",
len(trusted_urls),
len(raw_portrait_urls),
)
# 替换:image_url 用第一张(可能是AI或原图),ref_imgs 用剩余
if image_url and merged:
image_url = merged[0]
ref_imgs = merged[1:] if len(merged) > 1 else []
if trust_chain_applied and trusted_urls:
# 替换:原 image_url 用第一张 AI 图,ref_imgs 用剩余
if image_url and trusted_urls:
image_url = trusted_urls[0]
ref_imgs = trusted_urls[1:] if len(trusted_urls) > 1 else []
else:
ref_imgs = merged
ref_imgs = trusted_urls
# ─────────────────────────────────────────────────────────────────
# 判断任务模式:
-62
View File
@@ -1,62 +0,0 @@
"""AI 配置版本号管理 — Redis 通知机制.
admin 后台修改 ai_models / ai_capability_configs 后调用 bump_version(),
SaaS 端 AIRouter 每次取配置前比对版本号,变了才重新查 DB。
Redis key: xiaoxia:ai_config:version = 时间戳字符串
"""
from __future__ import annotations
import logging
import time
from typing import Optional
logger = logging.getLogger(__name__)
_REDIS_KEY = "xiaoxia:ai_config:version"
def _get_redis_client():
"""获取 Redis 客户端(复用 Celery broker 连接)."""
try:
import redis as _redis
from packages.shared.config import get_shared_settings
settings = get_shared_settings()
redis_url = getattr(settings, "redis_url", None) or getattr(
settings, "celery_broker_url", "redis://localhost:6379/0"
)
return _redis.Redis.from_url(redis_url, decode_responses=True, socket_timeout=2)
except Exception as e:
logger.warning("AI config version: Redis 客户端初始化失败: %s", e)
return None
def bump_version() -> str:
"""写入新版本号(当前时间戳),返回版本号字符串。失败返回空串。"""
r = _get_redis_client()
if r is None:
logger.warning("AI config bump_version: Redis 不可用,跳过版本号更新")
return ""
try:
ver = str(int(time.time() * 1000))
r.set(_REDIS_KEY, ver)
logger.info("AI config version bumped to %s", ver)
return ver
except Exception as e:
logger.warning("AI config bump_version 失败: %s", e)
return ""
def get_version() -> Optional[str]:
"""读取当前版本号。Redis 不可用或异常返回 None。"""
r = _get_redis_client()
if r is None:
return None
try:
return r.get(_REDIS_KEY)
except Exception as e:
logger.warning("AI config get_version 失败: %s", e)
return None
-516
View File
@@ -1,516 +0,0 @@
"""AI 模型路由层 — 统一模型配置读取与客户端构建.
业务代码通过 AIRouter 获取客户端,不再硬编码 model/api_key/base_url。
配置来源:DB ai_capability_configs JOIN ai_models → Redis 版本号缓存 → SharedSettings fallback。
使用方式:
from packages.shared.ai_router import ai_router
client = ai_router.get_llm_client("intent_parsing")
result = client.chat_completion(messages=[...])
"""
from __future__ import annotations
import logging
import threading
from dataclasses import dataclass
from packages.shared.config import get_shared_settings
logger = logging.getLogger(__name__)
# ── 配置数据类 ──────────────────────────────────────────────────────────────
@dataclass(frozen=True)
class ModelConfig:
"""单个 AI 模型配置(来自 ai_models 表)"""
id: str
name: str
provider: str
model_key: str
api_key: str
api_base: str
api_version: str | None
status: str
@dataclass(frozen=True)
class CapabilityConfig:
"""业务能力配置(来自 ai_capability_configs JOIN ai_models)"""
capability_key: str
capability_name: str
primary_model: ModelConfig | None
lite_model: ModelConfig | None
fallback_model: ModelConfig | None
timeout_seconds: int
max_retries: int
max_tokens: int | None
temperature: float | None
concurrency: int
extra_params: dict
is_enabled: bool
# ── 简单包装类(TTS / ImageGen / VideoGen)──────────────────────────────────
class TTSClient:
"""TTS 客户端(简单配置持有者,实际调用由 CosyVoiceService 完成)"""
def __init__(
self,
provider: str,
api_key: str,
base_url: str,
model: str,
timeout: int = 60,
extra_params: dict | None = None,
):
self.provider = provider
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.extra_params = extra_params or {}
@property
def is_available(self) -> bool:
return bool(self.api_key and self.base_url and self.model)
class ImageGenClient:
"""图片生成客户端(简单配置持有者)"""
def __init__(
self,
provider: str,
api_key: str,
base_url: str,
model: str,
timeout: int = 60,
extra_params: dict | None = None,
):
self.provider = provider
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.extra_params = extra_params or {}
@property
def is_available(self) -> bool:
return bool(self.api_key and self.base_url and self.model)
class VideoGenClient:
"""视频生成客户端(简单配置持有者)"""
def __init__(
self,
provider: str,
api_key: str,
base_url: str,
model: str,
timeout: int = 600,
extra_params: dict | None = None,
):
self.provider = provider
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.extra_params = extra_params or {}
@property
def is_available(self) -> bool:
return bool(self.api_key and self.base_url and self.model)
# ── DB Session 获取 ─────────────────────────────────────────────────────────
def _get_session():
"""获取 DB session,兼容 api / worker / 独立脚本场景"""
from packages.adapters.sqlalchemy_impl.session import SessionLocal
if SessionLocal is not None:
return SessionLocal()
try:
from worker_app.db import SessionLocal as WorkerSL
if WorkerSL is not None:
return WorkerSL()
except ImportError:
pass
try:
from app.db import SessionLocal as ApiSL
if ApiSL is not None:
return ApiSL()
except ImportError:
pass
return None
# ── 核心路由类 ──────────────────────────────────────────────────────────────
class AIRouter:
"""AI 模型路由器 — 统一配置读取与客户端构建.
缓存策略:
1. 本地内存缓存 {capability_key: CapabilityConfig}
2. 每次读取前比对 Redis 版本号,变了则清缓存重新查 DB
3. DB 无配置 / Redis 不可用 → fallback 到 SharedSettings 环境变量
"""
def __init__(self):
self._cache: dict[str, CapabilityConfig] = {}
self._local_ver: str | None = None
self._lock = threading.Lock()
def _check_version(self) -> bool:
"""检查 Redis 版本号,变了返回 True(需要刷新缓存)"""
from packages.shared.ai_config_version import get_version
current_ver = get_version()
if current_ver is None:
return False
if self._local_ver != current_ver:
return True
return False
def _load_from_db(self, capability_key: str) -> CapabilityConfig | None:
"""从 DB 加载配置(ai_capability_configs JOIN ai_models)"""
session = _get_session()
if session is None:
logger.warning("AI Router: 无法获取 DB session")
return None
try:
from sqlalchemy import text
sql = text("""
SELECT
cc.capability_key, cc.capability_name, cc.timeout_seconds,
cc.max_retries, cc.max_tokens, cc.temperature,
cc.concurrency, cc.extra_params, cc.is_enabled,
pm.id AS pm_id, pm.name AS pm_name, pm.provider AS pm_provider,
pm.model_key AS pm_model_key, pm.api_key AS pm_api_key,
pm.api_base AS pm_api_base, pm.api_version AS pm_api_version,
pm.status AS pm_status,
lm.id AS lm_id, lm.name AS lm_name, lm.provider AS lm_provider,
lm.model_key AS lm_model_key, lm.api_key AS lm_api_key,
lm.api_base AS lm_api_base, lm.api_version AS lm_api_version,
lm.status AS lm_status,
fm.id AS fm_id, fm.name AS fm_name, fm.provider AS fm_provider,
fm.model_key AS fm_model_key, fm.api_key AS fm_api_key,
fm.api_base AS fm_api_base, fm.api_version AS fm_api_version,
fm.status AS fm_status
FROM ai_capability_configs cc
LEFT JOIN ai_models pm ON cc.primary_model_id = pm.id AND pm.deleted_at IS NULL
LEFT JOIN ai_models lm ON cc.lite_model_id = lm.id AND lm.deleted_at IS NULL
LEFT JOIN ai_models fm ON cc.fallback_model_id = fm.id AND fm.deleted_at IS NULL
WHERE cc.capability_key = :key AND cc.is_enabled = true
""")
row = session.execute(sql, {"key": capability_key}).first()
if not row:
return None
def _to_model(prefix: str) -> ModelConfig | None:
mid = getattr(row, f"{prefix}_id", None)
if not mid:
return None
return ModelConfig(
id=mid,
name=getattr(row, f"{prefix}_name", "") or "",
provider=getattr(row, f"{prefix}_provider", "") or "",
model_key=getattr(row, f"{prefix}_model_key", "") or "",
api_key=getattr(row, f"{prefix}_api_key", "") or "",
api_base=getattr(row, f"{prefix}_api_base", "") or "",
api_version=getattr(row, f"{prefix}_api_version", None),
status=getattr(row, f"{prefix}_status", "active") or "active",
)
return CapabilityConfig(
capability_key=row.capability_key,
capability_name=row.capability_name,
primary_model=_to_model("pm"),
lite_model=_to_model("lm"),
fallback_model=_to_model("fm"),
timeout_seconds=row.timeout_seconds or 30,
max_retries=row.max_retries or 1,
max_tokens=row.max_tokens,
temperature=row.temperature,
concurrency=row.concurrency or 2,
extra_params=row.extra_params or {},
is_enabled=row.is_enabled,
)
except Exception as e:
logger.warning("AI Router: DB 查询失败 (key=%s): %s", capability_key, e)
return None
finally:
session.close()
def get_capability(self, key: str) -> CapabilityConfig | None:
"""获取业务能力配置(带缓存)"""
with self._lock:
if self._check_version():
self._cache.clear()
from packages.shared.ai_config_version import get_version
self._local_ver = get_version()
if key in self._cache:
return self._cache[key]
config = self._load_from_db(key)
if config:
self._cache[key] = config
return config
def _get_model_or_fallback(self, cap: CapabilityConfig, variant: str = "primary") -> ModelConfig | None:
"""按 variant 选择模型,不存在则降级。
- primary: primary → fallback
- lite: lite → primary
- fallback: fallback → primary(修复点:此前 fallback variant 被忽略,错误地使用了 primary 模型)
"""
if variant == "fallback":
if cap.fallback_model:
return cap.fallback_model
if cap.primary_model:
return cap.primary_model
elif variant == "lite":
if cap.lite_model:
return cap.lite_model
if cap.primary_model:
return cap.primary_model
else: # primary
if cap.primary_model:
return cap.primary_model
if cap.fallback_model:
return cap.fallback_model
return None
# ── 构建客户端 ─────────────────────────────────────────────────────────
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig):
"""构建 LLM 客户端 — 返回 DoubaoClient 实例"""
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
timeout=cap.timeout_seconds,
max_retries=cap.max_retries,
max_tokens=cap.max_tokens,
temperature=cap.temperature,
extra_params=cap.extra_params,
provider=model.provider,
)
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig):
"""构建 VLM 客户端 — 返回 DoubaoClient 实例(DoubaoClient 已支持 vision_completion)"""
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
timeout=cap.timeout_seconds,
max_retries=cap.max_retries,
max_tokens=cap.max_tokens,
temperature=cap.temperature,
extra_params=cap.extra_params,
provider=model.provider,
)
def _build_tts_client(self, model: ModelConfig, cap: CapabilityConfig) -> TTSClient:
return TTSClient(
provider=model.provider,
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
timeout=cap.timeout_seconds,
extra_params=cap.extra_params,
)
def _build_image_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> ImageGenClient:
return ImageGenClient(
provider=model.provider,
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
timeout=cap.timeout_seconds,
extra_params=cap.extra_params,
)
def _build_video_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> VideoGenClient:
return VideoGenClient(
provider=model.provider,
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
timeout=cap.timeout_seconds,
extra_params=cap.extra_params,
)
# ── 公开接口 ────────────────────────────────────────────────────────────
def get_llm_client(self, key: str, variant: str = "primary"):
"""获取 LLM 客户端(返回 DoubaoClient 实例)"""
cap = self.get_capability(key)
if cap and cap.is_enabled:
model = self._get_model_or_fallback(cap, variant)
if model and model.api_key:
return self._build_llm_client(model, cap)
return self._fallback_llm_client(key)
def get_vision_client(self, key: str, variant: str = "primary"):
"""获取 VLM 客户端(返回 DoubaoClient 实例)"""
cap = self.get_capability(key)
if cap and cap.is_enabled:
model = self._get_model_or_fallback(cap, variant)
if model and model.api_key:
return self._build_vision_client(model, cap)
return self._fallback_vision_client(key)
def get_tts_client(self, key: str = "tts") -> TTSClient | None:
"""获取 TTS 客户端"""
cap = self.get_capability(key)
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
return self._build_tts_client(cap.primary_model, cap)
return self._fallback_tts_client()
def get_image_gen_client(self, key: str = "image_generation") -> ImageGenClient | None:
"""获取图片生成客户端"""
cap = self.get_capability(key)
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
return self._build_image_gen_client(cap.primary_model, cap)
return self._fallback_image_gen_client()
def get_video_gen_client(self, key: str = "video_generation") -> VideoGenClient | None:
"""获取视频生成客户端"""
cap = self.get_capability(key)
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
return self._build_video_gen_client(cap.primary_model, cap)
return self._fallback_video_gen_client()
# ── Fallback 方法(读 SharedSettings 环境变量)──────────────────────────
def _fallback_llm_client(self, key: str):
"""Fallback LLM 客户端 — 从 settings 读取配置,不硬编码"""
settings = get_shared_settings()
model_map = {
"intent_parsing": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
"copy_fusion": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
"storyboard": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
"copy_review": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
"asset_classify": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
}
if key in model_map:
model_id, base_url, api_key = model_map[key]
else:
model_id = settings.doubao_model
base_url = settings.doubao_base_url
api_key = settings.doubao_api_key
if not api_key:
return None
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
provider="volcengine",
api_key=api_key,
base_url=base_url,
model=model_id,
timeout=settings.doubao_timeout,
max_retries=settings.doubao_max_retries,
)
def _fallback_vision_client(self, key: str):
"""Fallback VLM 客户端 — 从 settings 读取 dashscope 配置,不硬编码"""
settings = get_shared_settings()
api_key = getattr(settings, "dashscope_api_key", "")
if not api_key:
return None
base_url = getattr(settings, "dashscope_base_url", "") or ""
model = getattr(settings, "dashscope_model", "") or getattr(settings, "doubao_vision_model", "")
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
provider="dashscope",
api_key=api_key,
base_url=base_url,
model=model,
timeout=15,
)
def _fallback_tts_client(self) -> TTSClient | None:
settings = get_shared_settings()
api_key = getattr(settings, "cosyvoice_api_key", "")
if not api_key:
return None
base_url = getattr(settings, "cosyvoice_base_url", "")
model = getattr(settings, "cosyvoice_model", "")
return TTSClient(provider="dashscope", api_key=api_key, base_url=base_url, model=model)
def _fallback_image_gen_client(self) -> ImageGenClient | None:
settings = get_shared_settings()
api_key = getattr(settings, "doubao_api_key", "")
if not api_key:
return None
base_url = getattr(settings, "doubao_base_url", "")
model = getattr(settings, "doubao_image_model", "")
return ImageGenClient(
provider="volcengine",
api_key=api_key,
base_url=base_url,
model=model,
timeout=getattr(settings, "doubao_image_timeout", 60),
)
def _fallback_video_gen_client(self) -> VideoGenClient | None:
settings = get_shared_settings()
api_key = getattr(settings, "doubao_api_key", "")
if not api_key:
return None
base_url = getattr(settings, "doubao_base_url", "")
model = getattr(settings, "doubao_video_model", "")
return VideoGenClient(
provider="volcengine",
api_key=api_key,
base_url=base_url,
model=model,
timeout=getattr(settings, "doubao_video_timeout", 600),
)
def invalidate(self):
"""清空本地缓存"""
with self._lock:
self._cache.clear()
self._local_ver = None
# ── 全局单例 ──────────────────────────────────────────────────────────────
ai_router = AIRouter()
-2
View File
@@ -51,8 +51,6 @@ task_routes = {
"ai_avatar_render.execute": {"queue": QUEUE_GENERATION},
# GPU MuseTalk 口型同步(用户等成片,链路子任务全部走 generation 避免跨队列阻塞)
"lipsync_gpu_process_async": {"queue": QUEUE_GENERATION},
# #2076 Ditto 蚂蚁数字人口型同步(走 generation 队列,避免跨队列阻塞)
"lipsync_ditto_process_async": {"queue": QUEUE_GENERATION},
"lipsync_tts.synthesize_and_submit": {"queue": QUEUE_GENERATION},
"lipsync_tts.poll_mediakit_status": {"queue": QUEUE_GENERATION},
"lipsync_tts.persist_output_video": {"queue": QUEUE_GENERATION},
-395
View File
@@ -1,395 +0,0 @@
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
from __future__ import annotations
import sys
import unittest
from dataclasses import dataclass
from typing import Optional
from unittest.mock import MagicMock, patch
# ── Pre-mock heavy import chain to avoid pulling in full app ──
_mock_config = MagicMock()
_mock_settings = MagicMock()
_mock_settings.doubao_model = "doubao-seed-2-1-pro-260915"
_mock_settings.doubao_fast_model = "doubao-seed-2-1-pro-260915"
_mock_settings.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
_mock_settings.doubao_api_key = "test-key"
_mock_settings.doubao_timeout = 45
_mock_settings.doubao_max_retries = 1
_mock_settings.doubao_image_model = "doubao-seedream-5-0-flash-260915"
_mock_settings.doubao_image_timeout = 60
_mock_settings.doubao_video_model = "doubao-seedance-2-5-260628"
_mock_settings.doubao_video_timeout = 600
_mock_settings.dashscope_api_key = "ds-key"
_mock_settings.cosyvoice_api_key = "cv-key"
_mock_settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
_mock_settings.cosyvoice_model = "cosyvoice-v3-flash"
_mock_settings.redis_url = "redis://localhost:6379/0"
_mock_settings.celery_broker_url = "redis://localhost:6379/0"
_mock_config.get_shared_settings.return_value = _mock_settings
# Prevent the full packages.shared from loading
for mod_name in list(sys.modules.keys()):
if "packages.shared" in mod_name and "ai_router" not in mod_name and "ai_config_version" not in mod_name:
pass # don't remove, just prevent new imports
# Direct import of our modules (bypassing __init__.py)
import importlib.util
import os
def _load_module_from_file(name, path):
spec = importlib.util.spec_from_file_location(name, path)
mod = importlib.util.module_from_spec(spec)
sys.modules[name] = mod
spec.loader.exec_module(mod)
return mod
# Load ai_config_version
_ai_config_version = _load_module_from_file(
"packages.shared.ai_config_version",
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_config_version.py"),
)
# Patch get_shared_settings in the loaded module
_ai_config_version.get_shared_settings = lambda: _mock_settings
# Load ai_router - needs packages.shared.config to be available
sys.modules["packages.shared.config"] = MagicMock()
sys.modules["packages.shared.config"].get_shared_settings = lambda: _mock_settings
# Mock packages.shared.ai_client to avoid triggering packages.shared.__init__ chain
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
_mock_ai_client = MagicMock()
class _FakeDoubaoClient:
"""Fake DoubaoClient for testing - mimics the real interface."""
def __init__(self, api_key="", base_url="", model="", timeout=0, max_retries=0,
max_tokens=None, temperature=None, extra_params=None, provider="volcengine"):
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.max_retries = max_retries
self.max_tokens = max_tokens
self.temperature = temperature
self.extra_params = extra_params or {}
self.provider = provider
self.vision_model = model
@property
def is_available(self):
return bool(self.api_key)
def chat_completion(self, messages, **kwargs):
return None
def vision_completion(self, messages, **kwargs):
return None
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
sys.modules["packages.shared.ai_client"] = _mock_ai_client
_ai_router = _load_module_from_file(
"packages.shared.ai_router",
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_router.py"),
)
class TestAIConfigVersion(unittest.TestCase):
"""Redis 版本号机制测试"""
@patch.object(_ai_config_version, "_get_redis_client")
def test_bump_version_success(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.set.return_value = True
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.bump_version()
self.assertTrue(ver)
self.assertTrue(ver.isdigit())
mock_r.set.assert_called_once()
@patch.object(_ai_config_version, "_get_redis_client")
def test_bump_version_redis_unavailable(self, mock_redis_fn):
mock_redis_fn.return_value = None
ver = _ai_config_version.bump_version()
self.assertEqual(ver, "")
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_success(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.get.return_value = "1234567890"
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.get_version()
self.assertEqual(ver, "1234567890")
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_redis_down(self, mock_redis_fn):
mock_redis_fn.return_value = None
ver = _ai_config_version.get_version()
self.assertIsNone(ver)
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_exception(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.get.side_effect = Exception("connection refused")
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.get_version()
self.assertIsNone(ver)
class TestAIRouter(unittest.TestCase):
"""AIRouter 路由/缓存/fallback 测试"""
def setUp(self):
self.router = _ai_router.AIRouter()
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_capability_db_unavailable(self, mock_ver):
with patch.object(_ai_router, "_get_session", return_value=None):
cap = self.router.get_capability("intent_parsing")
self.assertIsNone(cap)
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_capability_from_db(self, mock_ver):
mock_session = MagicMock()
mock_row = MagicMock()
mock_row.capability_key = "intent_parsing"
mock_row.capability_name = "文案意图解析"
mock_row.timeout_seconds = 45
mock_row.max_retries = 1
mock_row.max_tokens = None
mock_row.temperature = None
mock_row.concurrency = 2
mock_row.extra_params = {}
mock_row.is_enabled = True
mock_row.pm_id = "model-1"
mock_row.pm_name = "豆包"
mock_row.pm_provider = "volcengine"
mock_row.pm_model_key = "doubao-seed-1-6-250615"
mock_row.pm_api_key = "test-key"
mock_row.pm_api_base = "https://ark.test.com"
mock_row.pm_api_version = None
mock_row.pm_status = "active"
mock_row.lm_id = None
mock_row.fm_id = None
mock_session.execute.return_value.first.return_value = mock_row
with patch.object(_ai_router, "_get_session", return_value=mock_session):
cap = self.router.get_capability("intent_parsing")
self.assertIsNotNone(cap)
self.assertEqual(cap.capability_key, "intent_parsing")
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
@patch.object(_ai_config_version, "get_version", side_effect=[None, "v2"])
def test_cache_invalidation_on_version_change(self, mock_ver):
with patch.object(self.router, "_load_from_db", return_value=None):
self.router.get_capability("test_key")
self.router._local_ver = "v1"
self.assertTrue(self.router._check_version())
@patch.object(_ai_config_version, "get_version", return_value="same_ver")
def test_cache_hit_same_version(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="test-model",
api_key="key", api_base="https://test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test", primary_model=model,
lite_model=None, fallback_model=None, timeout_seconds=30,
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
extra_params={}, is_enabled=True,
)
self.router._cache["test"] = cap
self.router._local_ver = "same_ver"
result = self.router.get_capability("test")
self.assertEqual(result, cap)
def test_invalidate_clears_cache(self):
self.router._cache["x"] = MagicMock()
self.router._local_ver = "v1"
self.router.invalidate()
self.assertEqual(len(self.router._cache), 0)
self.assertIsNone(self.router._local_ver)
@patch.object(_ai_router, "_get_session", return_value=None)
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_llm_client_fallback(self, mock_ver, mock_session):
_ai_router.get_shared_settings = lambda: _mock_settings
client = self.router.get_llm_client("intent_parsing")
self.assertIsNotNone(client)
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
self.assertEqual(client.api_key, "test-key")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_llm_client_from_db(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
api_key="db-key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=350, temperature=0.1,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_llm_client("image_analysis")
self.assertIsNotNone(client)
self.assertEqual(client.model, "qwen3.8-flash")
self.assertEqual(client.provider, "dashscope")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_vision_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_vision_client("image_analysis")
self.assertIsNotNone(client)
# #2220: vision client is now DoubaoClient with vision_completion
self.assertTrue(hasattr(client, "vision_completion"))
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_tts_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="cosyvoice-v3-flash",
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="tts", capability_name="语音合成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_tts_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "cosyvoice-v3-flash")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_image_gen_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="seedream-5.0-flash",
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_generation", capability_name="图片生成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={"size": "1K"}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_image_gen_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "seedream-5.0-flash")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_video_gen_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="seedance-2.5",
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="video_generation", capability_name="视频生成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=600, max_retries=1, max_tokens=None, temperature=None,
concurrency=1, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_video_gen_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "seedance-2.5")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_lite_variant_preference(self, mock_ver):
primary = _ai_router.ModelConfig(id="p1", name="pro", provider="volcengine", model_key="pro-model", api_key="k", api_base="u", api_version=None, status="active")
lite = _ai_router.ModelConfig(id="l1", name="lite", provider="volcengine", model_key="lite-model", api_key="k", api_base="u", api_version=None, status="active")
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=primary, lite_model=lite, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
model = self.router._get_model_or_fallback(cap, "lite")
self.assertEqual(model.model_key, "lite-model")
model_primary = self.router._get_model_or_fallback(cap, "primary")
self.assertEqual(model_primary.model_key, "pro-model")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_disabled_capability_returns_fallback(self, mock_ver):
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test",
primary_model=None, lite_model=None, fallback_model=None,
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=False,
)
_ai_router.get_shared_settings = lambda: _mock_settings
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_llm_client("test")
self.assertIsNotNone(client)
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_fallback_chain_primary_none(self, mock_ver):
"""primary_model 为 None 时 fallback 到 fallback_model"""
fb = _ai_router.ModelConfig(id="f1", name="fb", provider="volcengine", model_key="fb-model", api_key="k", api_base="u", api_version=None, status="active")
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test",
primary_model=None, lite_model=None, fallback_model=fb,
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
model = self.router._get_model_or_fallback(cap, "primary")
self.assertEqual(model.model_key, "fb-model")
class TestModelConfig(unittest.TestCase):
"""数据类测试"""
def test_model_config_frozen(self):
m = _ai_router.ModelConfig(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b", api_version=None, status="active")
with self.assertRaises(AttributeError):
m.model_key = "new"
def test_capability_config_frozen(self):
c = _ai_router.CapabilityConfig(
capability_key="k", capability_name="n", primary_model=None,
lite_model=None, fallback_model=None, timeout_seconds=30,
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
extra_params={}, is_enabled=True,
)
with self.assertRaises(AttributeError):
c.is_enabled = False
class TestClientAvailability(unittest.TestCase):
"""客户端可用性测试"""
def test_tts_client_available(self):
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
self.assertTrue(c.is_available)
def test_tts_client_unavailable_no_model(self):
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="")
self.assertFalse(c.is_available)
def test_image_gen_client_unavailable_no_url(self):
c = _ai_router.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
self.assertFalse(c.is_available)
def test_video_gen_client_available(self):
c = _ai_router.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
self.assertTrue(c.is_available)
if __name__ == "__main__":
unittest.main()
+4 -4
View File
@@ -73,15 +73,15 @@ class TestSharedSettingsDefaults:
def test_default_cosyvoice_settings(self):
s = SharedSettings()
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
assert s.cosyvoice_format == "mp3"
assert s.cosyvoice_sample_rate == 22050
def test_default_doubao_settings(self):
s = SharedSettings()
assert s.doubao_model == "" # 零硬编码:默认值已清空
assert "doubao" in s.doubao_model
assert s.doubao_timeout == 45 # #2180 默认提到45s
assert s.doubao_max_retries == 3
assert s.doubao_max_retries == 1
class TestAPISettingsDefaults:
@@ -321,7 +321,7 @@ class TestWorkerSettingsDefaults:
assert s.database_url # 继承自SharedSettings
assert s.redis_url
assert s.oss_endpoint
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
class TestGetWorkerSettings:
+4 -4
View File
@@ -102,17 +102,17 @@ class TestSharedSettingsDefaults:
def test_default_cosyvoice_config(self):
"""CosyVoice 默认配置"""
s = self._make_settings()
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
assert s.cosyvoice_sample_rate == 22050
assert s.cosyvoice_format == "mp3"
assert s.cosyvoice_clone_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_clone_model == "voice-enrollment"
def test_default_doubao_config(self):
"""豆包默认配置"""
s = self._make_settings()
assert s.doubao_timeout == 45 # #2180 默认提到45s
assert s.doubao_max_retries == 3
assert s.doubao_base_url == "" # 零硬编码:默认值已清空
assert s.doubao_max_retries == 1
assert "volces.com" in s.doubao_base_url
def test_default_empty_api_keys(self):
"""API Key 默认空字符串"""
+10 -50
View File
@@ -27,10 +27,7 @@ def mock_client() -> MagicMock:
@pytest.fixture
def service(mock_client: MagicMock) -> CosyVoiceService:
"""Create CosyVoiceService with mocked HTTP client and config."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test-12345678"
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
@@ -40,7 +37,6 @@ def service(mock_client: MagicMock) -> CosyVoiceService:
settings.cosyvoice_voice = "longxiaochun_v3"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None # ai_router returns None in tests
svc = CosyVoiceService(http_client=mock_client)
svc.CLONE_POLL_INTERVAL = 0.001 # 加速测试
svc.RETRY_BACKOFF = 0.001
@@ -52,10 +48,7 @@ class TestInitConfig:
def test_base_url_with_old_text2audio_path_gets_normalized(self, mock_client: MagicMock) -> None:
"""旧版 base_url 带 text2audio 路径应自动修正."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio"
@@ -65,16 +58,12 @@ class TestInitConfig:
settings.cosyvoice_voice = "test"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
assert svc._base_url == "https://dashscope.aliyuncs.com/api/v1"
def test_custom_params_override_config(self, mock_client: MagicMock) -> None:
"""显式传入参数覆盖配置."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-config"
settings.cosyvoice_base_url = "https://config.example.com"
@@ -84,7 +73,6 @@ class TestInitConfig:
settings.cosyvoice_voice = "test"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(
api_key="sk-custom",
base_url="https://custom.example.com/api/v1",
@@ -99,10 +87,7 @@ class TestInitConfig:
def test_context_manager(self, mock_client: MagicMock) -> None:
"""上下文管理器正常工作."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -120,7 +105,6 @@ class TestInitConfig:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with svc as s:
assert s is svc
@@ -129,10 +113,7 @@ class TestInitConfig:
def test_owns_client_gets_closed(self) -> None:
"""自有client在close时被关闭."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -150,7 +131,6 @@ class TestInitConfig:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
with patch("packages.application.cosyvoice_service.httpx.Client") as mock_cls:
mock_instance = MagicMock()
mock_cls.return_value = mock_instance
@@ -193,10 +173,7 @@ class TestSubmitCloneTask:
def test_no_api_key_raises_auth_error(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -214,7 +191,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError, match="API Key 未配置"):
svc.submit_clone_task(audio_url="https://example.com/audio.mp3")
@@ -282,10 +258,7 @@ class TestSubmitCloneTask:
def test_audio_url_signer_is_called(self, mock_client: MagicMock) -> None:
"""配置了audio_url_signer时会被调用预签名."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -303,7 +276,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
signer = MagicMock(return_value="https://signed.example.com/audio.mp3?token=xxx")
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
@@ -325,10 +297,7 @@ class TestSubmitCloneTask:
def test_signer_failure_falls_back_to_original_url(self, mock_client: MagicMock) -> None:
"""预签名失败时回退到原始URL."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -346,7 +315,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
signer = MagicMock(side_effect=RuntimeError("sign failed"))
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
@@ -407,10 +375,7 @@ class TestQueryVoiceStatus:
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -428,7 +393,6 @@ class TestQueryVoiceStatus:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError):
svc.query_voice_status("v1")
@@ -625,10 +589,7 @@ class TestSubmitSynthesizeTask:
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -646,7 +607,6 @@ class TestSubmitSynthesizeTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError):
svc.submit_synthesize_task(text="你好", voice_id="v1")
-197
View File
@@ -1,197 +0,0 @@
"""Ditto 蚂蚁数字人客户端单元测试 — #2076."""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import httpx
import pytest
from packages.application.ditto_service import DittoClient, DittoError, DittoResult
class _FakeResponse:
def __init__(self, status_code=200, content=b"\x00\x01" * 1000, headers=None, text=""):
self.status_code = status_code
self.content = content
self.headers = headers or {}
self.text = text
def _make_client(base_url="http://ditto:8000", default_video_url="http://oss/tpl.mp4", max_retries=2, timeout=60):
with patch("packages.application.ditto_service.get_api_settings") as mock_settings:
s = MagicMock()
s.ditto_api_base_url = base_url
s.ditto_default_video_url = default_video_url
s.ditto_max_retries = max_retries
s.ditto_request_timeout = timeout
mock_settings.return_value = s
return DittoClient()
def test_is_configured_true():
c = _make_client()
assert c.is_configured is True
def test_is_configured_false_without_base():
c = _make_client(base_url="")
assert c.is_configured is False
def test_is_configured_false_without_template():
c = _make_client(default_video_url="")
assert c.is_configured is False
def test_health_ok():
c = _make_client()
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.get.return_value = _FakeResponse(200)
mock_cls.return_value.__enter__.return_value = client
assert c.health() is True
client.get.assert_called_once()
def test_health_fail_status():
c = _make_client()
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.get.return_value = _FakeResponse(500)
mock_cls.return_value.__enter__.return_value = client
assert c.health() is False
def test_health_network_error():
c = _make_client()
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.get.side_effect = httpx.ConnectError("fail")
mock_cls.return_value.__enter__.return_value = client
assert c.health() is False
def test_generate_missing_base():
c = _make_client(base_url="")
with pytest.raises(DittoError, match="DITTO_API_BASE_URL"):
c.generate(audio_url="http://x/a.mp3", script="你好")
def test_generate_missing_audio():
c = _make_client()
with pytest.raises(DittoError, match="audio_url"):
c.generate(audio_url="", script="你好")
def test_generate_success_with_headers():
c = _make_client(max_retries=0)
fake_resp = _FakeResponse(
status_code=200,
content=b"\x00" * 99999,
headers={"X-RTF": "0.35", "X-Frames": "125", "X-Time": "12.5"},
)
with patch("httpx.Client") as mock_cls, patch("time.monotonic", side_effect=[0, 1]):
client = MagicMock()
client.post.return_value = fake_resp
mock_cls.return_value.__enter__.return_value = client
result = c.generate(audio_url="http://x/a.mp3", script="你好")
assert isinstance(result, DittoResult)
assert len(result.video_bytes) == 99999
assert result.rtf == 0.35
assert result.frames == 125
assert result.elapsed_seconds == 12.5
def test_generate_uses_default_template_when_video_url_empty():
c = _make_client(max_retries=0)
fake_resp = _FakeResponse(200, b"1" * 99999)
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.post.return_value = fake_resp
mock_cls.return_value.__enter__.return_value = client
c.generate(audio_url="http://x/a.mp3", script="你好")
call_kwargs = client.post.call_args
payload = call_kwargs.kwargs.get("json") or call_kwargs[1].get("json")
assert payload["video_url"] == "http://oss/tpl.mp4"
assert payload["audio_url"] == "http://x/a.mp3"
assert payload["script"] == "你好"
assert payload["emo_global"] == 4
assert payload["use_script_emo"] is True
def test_generate_retries_on_429_then_success():
c = _make_client(max_retries=2)
busy = _FakeResponse(429, b"", text="busy")
ok = _FakeResponse(200, b"v" * 99999)
with patch("httpx.Client") as mock_cls, patch("time.sleep") as mock_sleep:
client = MagicMock()
client.post.side_effect = [busy, ok]
mock_cls.return_value.__enter__.return_value = client
result = c.generate(audio_url="http://x/a.mp3", script="你好")
assert len(result.video_bytes) == 99999
assert mock_sleep.called
assert client.post.call_count == 2
def test_generate_429_exhausted():
c = _make_client(max_retries=1)
with patch("httpx.Client") as mock_cls, patch("time.sleep"):
client = MagicMock()
client.post.return_value = _FakeResponse(429, b"", text="busy")
mock_cls.return_value.__enter__.return_value = client
with pytest.raises(DittoError, match="重试"):
c.generate(audio_url="http://x/a.mp3", script="你好")
def test_generate_400_no_retry():
c = _make_client(max_retries=2)
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.post.return_value = _FakeResponse(400, b"", text="bad request")
mock_cls.return_value.__enter__.return_value = client
with pytest.raises(DittoError, match="Ditto 返回 400"):
c.generate(audio_url="http://x/a.mp3", script="你好")
assert client.post.call_count == 1 # 400 不重试
def test_generate_small_response_raises():
c = _make_client(max_retries=0)
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.post.return_value = _FakeResponse(200, b"xx")
mock_cls.return_value.__enter__.return_value = client
with pytest.raises(DittoError) as exc_info:
c.generate(audio_url="http://x/a.mp3", script="你好")
assert exc_info.value.code == "EmptyResponse"
def test_generate_and_persist_uploads_to_storage():
c = _make_client(max_retries=0)
fake_resp = _FakeResponse(200, b"v" * 99999)
fake_storage = MagicMock()
fake_storage.upload_file.return_value = "http://oss/ditto/x.mp4"
with (
patch("httpx.Client") as mock_cls,
patch("packages.shared.storage.get_shared_storage_service", return_value=fake_storage),
):
client = MagicMock()
client.post.return_value = fake_resp
mock_cls.return_value.__enter__.return_value = client
result = c.generate_and_persist(job_id="j1", user_id="u1", audio_url="http://x/a.mp3", script="hi")
assert result.video_url == "http://oss/ditto/x.mp4"
fake_storage.upload_file.assert_called_once()
call_args = fake_storage.upload_file.call_args
assert call_args.args[1].startswith("ditto-output/u1/j1")
def test_empty_script_replaced_with_space():
c = _make_client(max_retries=0)
fake_resp = _FakeResponse(200, b"v" * 99999)
with patch("httpx.Client") as mock_cls:
client = MagicMock()
client.post.return_value = fake_resp
mock_cls.return_value.__enter__.return_value = client
c.generate(audio_url="http://x/a.mp3", script="")
payload = client.post.call_args.kwargs["json"]
assert payload["script"] == " "
@@ -521,81 +521,3 @@ class TestIngestJob:
storage_key="k",
)
assert job.error_message == ""
class TestViralVideoResumeForRegenerate:
"""#2222: resume_from_image_analyzed 应支持 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。"""
def test_regen_from_copy_generated_clears_old_copy(self):
from datetime import datetime, timezone
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
# 模拟已经生成过文案和视频
job.status = ViralVideoStatus.COPY_GENERATED
job.copy_result = {"shots": [{"x": 1}], "voiceover_script": "旧文案"}
job.intent_result = {"intent": "旧意图"}
job.storyboard = [{"x": 1}]
job.generated_copy_text = "旧文案"
job.result_video_url = "http://old.mp4"
job.completed_at = datetime(2026, 10, 6, tzinfo=timezone.utc)
job.error_msg = ""
job.current_stage = "tts_generation"
job.phase_message = "TTS完成"
# 重新生成
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.copy_result is None
assert job.intent_result is None
assert job.storyboard is None
assert job.generated_copy_text == ""
assert job.result_video_url == ""
assert job.completed_at is None
assert job.error_msg == ""
assert job.current_stage == ""
assert job.phase_message == ""
def test_regen_from_completed_clears_old_copy(self):
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.COMPLETED
job.copy_result = {"shots": [], "voiceover_script": "xx"}
job.intent_result = {"intent": "x"}
job.result_video_url = "http://v.mp4"
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.copy_result is None
assert job.intent_result is None
assert job.result_video_url == ""
def test_first_call_from_image_analyzed_keeps_fields(self):
"""首次进入(IMAGE_ANALYZED)不应清空任何已有的字段。"""
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.IMAGE_ANALYZED
job.image_analysis = {"products": []}
job.industry = "美妆"
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.image_analysis == {"products": []}
assert job.industry == "美妆"
def test_wait_user_confirm_rejected(self):
"""wait_user_confirm 中间状态应被拒绝(前端正在编辑/确认文案)。"""
import pytest
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.WAIT_USER_CONFIRM
with pytest.raises(ValueError, match="Cannot resume"):
job.resume_from_image_analyzed()
@@ -1,221 +0,0 @@
"""功能计费改造测试:爆款读配置、对口型/智能剪辑预扣逻辑。
策略:
- 爆款:通过修改缓存中的 FeatureConfig(multiplier/model_pricing)验证价格随配置变化
- lip_sync / smart_edit:直接测 LipsyncService 的预扣/结算/退款辅助方法,
PointsService 用 mock,避免依赖真实积分账户。
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from packages.domain import feature_pricing_service as fps
from packages.domain.feature_pricing_service import FeatureConfig, refresh_feature_configs
@pytest.fixture(autouse=True)
def _reset_cache():
refresh_feature_configs()
yield
refresh_feature_configs()
def _seed_cache(configs: dict) -> None:
import time
fps._cache = (time.monotonic(), configs)
class TestViralVideoReadsConfig:
def test_multiplier_change_changes_price(self):
"""配置里 multiplier 改大后,爆款价格随之变大(证明不再读死常量)。"""
from packages.domain.points_rules import calculate_viral_video_credits
# 基线兜底
base = calculate_viral_video_credits(15, 1280, 720)
assert base == 29.68
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
vv.profit_multiplier = 2.0
_seed_cache(fallback)
changed = calculate_viral_video_credits(15, 1280, 720)
assert changed > base
# 精确校验:video_cost 相同,仅系数从 1.3 → 2.0
_, bd = __import__(
"packages.domain.points_rules", fromlist=["calculate_viral_video_credits_with_breakdown"]
).calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert bd["profit_multiplier"] == 2.0
def test_model_price_from_config(self):
"""model_pricing 改单价后,token 成本按新单价计算。"""
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
# seedance-2.5/720p/false 从 70 改成 100
vv.model_pricing["seedance-2.5"]["720p"]["false"] = 100.0
_seed_cache(fallback)
_, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert bd["model_price"] == 100.0
def test_price_cap_from_config(self):
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
vv.price_cap = 5.0
_seed_cache(fallback)
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert credits == 5.0
assert bd["price_cap"] == 5.0
def test_disabled_feature_returns_zero_credits(self):
"""功能 is_enabled=false 时计费函数返回 0(纯计费层语义)。"""
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
fallback["viral_video"].is_enabled = False
_seed_cache(fallback)
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert credits == 0.0
assert bd["feature_enabled"] is False
assert bd["charged"] is False
class TestLipSyncPricing:
def _make_service(self):
from app.services.lipsync_service import LipsyncService
svc = LipsyncService.__new__(LipsyncService)
svc.db = MagicMock()
return svc
def _lip_cfg(self, **kw):
base = dict(
feature_key="lip_sync",
name="对口型",
is_enabled=True,
fixed_cost=0.1,
profit_multiplier=1.0,
dynamic_unit_cost=0.05,
billing_mode="per_second",
price_cap=0.0,
model_pricing={},
description="",
)
base.update(kw)
return FeatureConfig(**base)
def test_estimate_duration_from_script(self):
svc = self._make_service()
# 10 个字 / 5 = 2 秒,下限 1
assert svc._estimate_duration(script_text="一二三四五六七八九十") == 2.0
# 无任何信息 → 默认 10 秒
assert svc._estimate_duration() == 10.0
def test_calculate_lipsync_price_per_second(self):
_seed_cache({"lip_sync": self._lip_cfg()})
price, bd = fps.calculate_price("lip_sync", dynamic_cost=20.0 * 0.05)
# dynamic 1.0 + fixed 0.1 = 1.1
assert price == 1.1
assert bd["charged"] is True
def test_settle_refunds_overcharge(self):
"""实际时长短 → 只退不补,退还差额。"""
svc = self._make_service()
_seed_cache({"lip_sync": self._lip_cfg()})
job = MagicMock()
job.credits_prepaid = 2.0
job.credits_cost = 0.0 # 未结算
job.user_id = "u1"
job.credits_transaction_id = "txn-old"
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
inst.refund_points.return_value = {"success": True}
svc._settle_lip_sync(job, actual_duration=10.0)
# final: (10*0.05 + 0.1)*1.0 = 0.6;退 2.0-0.6=1.4
assert round(job.credits_cost, 2) == 0.6
inst.refund_points.assert_called_once()
kwargs = inst.refund_points.call_args.kwargs
assert kwargs["amount"] == 1.4
def test_settle_no_refund_when_longer(self):
"""首期只退不补:实际更贵不补扣。"""
svc = self._make_service()
_seed_cache({"lip_sync": self._lip_cfg()})
job = MagicMock()
job.credits_prepaid = 0.5
job.credits_cost = 0.0
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
svc._settle_lip_sync(job, actual_duration=60.0)
assert round(job.credits_cost, 2) > 0.5
inst.refund_points.assert_not_called()
def test_refund_on_failure_full(self):
svc = self._make_service()
job = MagicMock()
job.credits_prepaid = 3.0
job.credits_cost = 0.0
job.user_id = "u1"
job.credits_transaction_id = "t1"
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
inst.refund_points.return_value = {"success": True}
svc._refund_lip_sync(job)
kwargs = inst.refund_points.call_args.kwargs
assert kwargs["amount"] == 3.0
class TestSmartEditFixedPrice:
def test_fixed_price_formula(self):
"""首期固定价:dynamic=0,price=fixed*multiplier,cap 封顶。"""
cfg = FeatureConfig(
feature_key="smart_edit",
name="智能剪辑",
is_enabled=True,
fixed_cost=2.0,
profit_multiplier=1.5,
billing_mode="model_based",
price_cap=0.0,
)
_seed_cache({"smart_edit": cfg})
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
# (0+2)*1.5 = 3.0
assert price == 3.0
assert bd["dynamic_cost"] == 0.0
def test_fixed_price_with_cap(self):
cfg = FeatureConfig(
feature_key="smart_edit",
is_enabled=True,
fixed_cost=10.0,
profit_multiplier=2.0,
price_cap=8.0,
)
_seed_cache({"smart_edit": cfg})
price, _ = fps.calculate_price("smart_edit", dynamic_cost=0.0)
assert price == 8.0
def test_disabled_smart_edit_free(self):
cfg = FeatureConfig(feature_key="smart_edit", is_enabled=False, fixed_cost=2.0)
_seed_cache({"smart_edit": cfg})
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
assert price == 0.0
assert bd["charged"] is False
-235
View File
@@ -1,235 +0,0 @@
"""feature_pricing_service 单元测试。
覆盖:
- 300s TTL 内存缓存(命中不重复 load / 过期重新 load / refresh 强制刷新)
- calculate_price 公式 (dynamic+fixed)*multiplier、price_cap 封顶、round
- disabled / 未知 key 返回 0
- DB 异常 / 空表 → 内置兜底配置(爆款启用且价格与现状一致)
- lookup_model_price 嵌套/扁平结构与旧版回落语义
"""
from __future__ import annotations
import time
import pytest
from packages.domain import feature_pricing_service as fps
from packages.domain.feature_pricing_service import (
CACHE_TTL_SECONDS,
FeatureConfig,
calculate_price,
get_feature_config,
is_feature_enabled,
lookup_model_price,
refresh_feature_configs,
)
@pytest.fixture(autouse=True)
def _reset_cache():
"""每个用例前后清空模块缓存,避免相互污染。"""
refresh_feature_configs()
yield
refresh_feature_configs()
def _cfg(key="x", **kw) -> FeatureConfig:
base = dict(
feature_key=key,
name=key,
is_enabled=True,
fixed_cost=0.2,
profit_multiplier=2.0,
dynamic_unit_cost=0.0,
billing_mode="per_second",
price_cap=0.0,
model_pricing={},
description="",
)
base.update(kw)
return FeatureConfig(**base)
class TestCacheTTL:
def test_cache_hit_avoids_reload(self, monkeypatch):
"""TTL 内第二次读取不再调 _load_all。"""
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
get_feature_config("x")
get_feature_config("x")
assert calls["n"] == 1
def test_expired_cache_reloads(self, monkeypatch):
"""超过 TTL 后重新 load。"""
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
assert calls["n"] == 1
# 把缓存时间戳回拨到 TTL 之前
ts, data = fps._cache
fps._cache = (ts - CACHE_TTL_SECONDS - 1, data)
get_feature_config("x")
assert calls["n"] == 2
def test_refresh_forces_reload(self, monkeypatch):
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
refresh_feature_configs()
get_feature_config("x")
assert calls["n"] == 2
def test_ttl_constant_is_300(self):
assert CACHE_TTL_SECONDS == 300.0
class TestCalculatePrice:
def test_basic_formula(self, monkeypatch):
# (dynamic 1.0 + fixed 0.2) * 2.0 = 2.4
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(dynamic_unit_cost=1.0)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 2.4
assert bd["dynamic_cost"] == 1.0
assert bd["fixed_cost"] == 0.2
assert bd["profit_multiplier"] == 2.0
assert bd["final_price"] == 2.4
assert bd["charged"] is True
def test_price_cap_clamps(self, monkeypatch):
# raw = (1+0.2)*2 = 2.4,cap=1.0 → 1.0
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=1.0)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 1.0
assert bd["price_cap"] == 1.0
def test_no_cap_keeps_raw(self, monkeypatch):
# cap=0 视为不封顶
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=0.0)})
price, _ = calculate_price("x", dynamic_cost=1.0)
assert price == 2.4
def test_rounded_two_decimals(self, monkeypatch):
monkeypatch.setattr(
fps,
"_load_all",
lambda: {"x": _cfg(fixed_cost=0.1, profit_multiplier=1.0)},
)
price, _ = calculate_price("x", dynamic_cost=1.0 / 3.0)
# 0.3333... + 0.1 = 0.4333 → 0.43
assert price == 0.43
def test_negative_dynamic_treated_as_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
price, _ = calculate_price("x", dynamic_cost=-5.0)
# (0 + 0.2) * 2 = 0.4
assert price == 0.4
def test_disabled_returns_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 0.0
assert bd["is_enabled"] is False
assert bd["charged"] is False
def test_unknown_key_returns_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
price, bd = calculate_price("nope", dynamic_cost=1.0)
assert price == 0.0
assert bd["charged"] is False
class TestDBFailureFallback:
def test_load_exception_uses_fallback(self, monkeypatch):
def boom():
raise RuntimeError("table does not exist")
monkeypatch.setattr(fps, "_load_all", boom)
cfg = get_feature_config("viral_video")
assert cfg is not None
assert cfg.is_enabled is True
assert cfg.fixed_cost == 0.15
assert cfg.profit_multiplier == 1.3
def test_empty_table_uses_fallback(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {})
assert get_feature_config("viral_video").is_enabled is True
assert get_feature_config("lip_sync").is_enabled is False
assert get_feature_config("smart_edit").is_enabled is False
def test_fallback_viral_price_matches_current(self, monkeypatch):
"""兜底爆款价格与旧硬编码现状一致:seedance-2.5/720p/false=70。"""
monkeypatch.setattr(fps, "_load_all", lambda: {})
from packages.domain.points_rules import calculate_viral_video_credits
# 默认全局开关关闭,但纯计费函数价格照常算
assert calculate_viral_video_credits(15, 1280, 720) == 29.68
def test_db_row_overrides_fallback(self, monkeypatch):
monkeypatch.setattr(
fps,
"_load_all",
lambda: {"viral_video": _cfg("viral_video", fixed_cost=0.5, profit_multiplier=2.0, price_cap=50.0)},
)
cfg = get_feature_config("viral_video")
assert cfg.fixed_cost == 0.5
assert cfg.profit_multiplier == 2.0
assert cfg.price_cap == 50.0
class TestIsFeatureEnabled:
def test_disabled_feature(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
assert is_feature_enabled("x") is False
def test_global_switch_off_blocks_enabled_feature(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
monkeypatch.setattr(fps, "_global_points_enabled", lambda: False)
assert is_feature_enabled("x") is False
def test_both_switches_on(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
monkeypatch.setattr(fps, "_global_points_enabled", lambda: True)
assert is_feature_enabled("x") is True
class TestLookupModelPrice:
NESTED = {
"seedance-2.5": {
"720p": {"false": 70.0, "true": 42.0},
},
"wan-3.0": {"480p": {"false": 0.3}},
}
def test_nested_exact_hit(self):
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", False) == 70.0
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", True) == 42.0
def test_missing_bool_key_returns_none(self):
# wan-3.0/480p 只有 false,请求 true → None(由调用方回落)
assert lookup_model_price(self.NESTED, "wan-3.0", "480p", True) is None
def test_unknown_model_returns_none(self):
assert lookup_model_price(self.NESTED, "nope", "720p", False) is None
def test_flat_structure(self):
flat = {"m|720p|false": 12.5}
assert lookup_model_price(flat, "m", "720p", False) == 12.5
assert lookup_model_price(flat, "m", "720p", True) is None
+2 -8
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
import os
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import pytest
@@ -108,14 +108,8 @@ class TestGetTtsService:
def test_empty_env_falls_back_to_auto_detect(self):
"""环境变量为空时自动检测."""
with (
patch.dict(os.environ, {"TTS_PROVIDER": ""}),
patch("packages.shared.config.get_shared_settings") as mock_settings,
):
with patch.dict(os.environ, {"TTS_PROVIDER": ""}):
# 没有 cosyvoice_api_key 时应该用 mock
settings = MagicMock()
settings.cosyvoice_api_key = ""
mock_settings.return_value = settings
service = get_tts_service(None)
assert service.provider_name == "mock"
+47 -66
View File
@@ -350,29 +350,6 @@ class TestViralVideoRepository:
class TestViralVideoPipeline:
"""编排器流水线测试。"""
# v3 分镜 XML(copy_display_markdown + clips + voiceover_script)
V3_XML = """<copy_display_markdown>今天给大家分享一支很显白的口红。</copy_display_markdown>
<clips>
<clip image_index="0" time_range="0-5秒">
<voiceover>大家好,今天分享一款口红</voiceover>
<visual>近景平视,缓慢推镜</visual>
<action_details>手持口红特写</action_details>
<audio_bgm>轻快流行BGM</audio_bgm>
<transition>硬切</transition>
<reference_image_index>0</reference_image_index>
</clip>
<clip image_index="1" time_range="5-15秒">
<voiceover>颜色特别好看很显白</voiceover>
<visual>特写,固定镜头</visual>
<action_details>嘴唇涂抹特写</action_details>
<audio_bgm>轻快BGM继续</audio_bgm>
<transition>结束</transition>
<reference_image_index>1</reference_image_index>
</clip>
</clips>
<voiceover_script>大家好,今天分享一款口红。颜色特别好看很显白</voiceover_script>
<theme>口红分享</theme>"""
@pytest.fixture
def mock_job(self):
return ViralVideoJob(
@@ -387,27 +364,23 @@ class TestViralVideoPipeline:
video_ratio="9:16",
)
@patch("apps.worker.worker_app.tasks.vision.analyze_images_v2")
@patch("packages.shared.ai_service.call_vision")
def test_image_analysis_step(self, mock_vision, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
# 每张图返回一个 v8 5 字段结果
mock_vision.return_value = [
{"type": "product", "name": "口红", "brand": "", "has_person": False, "summary_markdown": "一支口红"},
{"type": "product", "name": "口红", "brand": "", "has_person": False, "summary_markdown": "口红特写"},
]
mock_vision.return_value = {"name": "口红", "features": ["持久", "滋润"]}
result = _step_image_analysis(mock_job)
assert "images" in result
assert len(result["images"]) == 2 # 两张图片
assert "products" in result
assert len(result["products"]) == 2 # 两张图片
@patch("apps.worker.worker_app.tasks.vision.analyze_images_v2")
@patch("packages.shared.ai_service.call_vision")
def test_image_analysis_fallback(self, mock_vision, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
# v2 分析内部异常时,每图走兜底,仍返回 images 结构
mock_vision.side_effect = RuntimeError("vision unavailable")
# 模拟 call_vision 不存在
mock_vision.side_effect = ImportError("no module")
result = _step_image_analysis(mock_job)
assert "images" in result
assert "products" in result
def test_video_analysis_no_reference(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
@@ -417,40 +390,46 @@ class TestViralVideoPipeline:
result = _step_video_analysis(mock_job)
assert result is None
def test_intent_parsing_step_removed(self, mock_job):
"""intent_parsing 已合并进脚本生成,不再作为独立步骤/函数存在。"""
import apps.worker.worker_app.tasks.viral_video as vv
@patch("packages.shared.ai_service.call_llm")
def test_intent_parsing(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_intent_parsing
assert not hasattr(vv, "_step_intent_parsing")
mock_llm.return_value = {"intent": "推广口红", "tone": "活泼"}
result = _step_intent_parsing(mock_job, {"products": []})
assert "intent" in result
def test_script_generation_returns_copy_result(self, mock_job):
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_returns_copy_result(self, mock_llm, mock_job):
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
from packages.shared.ai_router import ai_router
class _FakeClient:
is_available = True
model = "fake-storyboard"
def __init__(self, xml: str):
self._xml = xml
def chat_completion(self, messages, **kwargs):
return self._xml
fake = _FakeClient(self.V3_XML)
orig_get = ai_router.get_llm_client
def _get(task, variant="primary"):
if task == "storyboard":
return fake
return orig_get(task, variant=variant)
ai_router.get_llm_client = _get # type: ignore
try:
result = _step_script_generation(mock_job, {"images": []})
finally:
ai_router.get_llm_client = orig_get # type: ignore
mock_llm.return_value = """<clips>
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快流行BGM">
<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="0,0" end="0,0" ease="linear"/>
</clip>
<clip image_index="0" transition="fade" zoom="null" duration_sec="10" bgm_note="轻快BGM">
<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>1</reference_image_index>
<ken_burns start="0,0" end="0,0" ease="linear"/>
</clip>
</clips>"""
result = _step_script_generation(
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
)
assert isinstance(result, dict)
assert "voiceover_script" in result
assert "shots" in result
@@ -522,6 +501,7 @@ class TestPipelineIntegration:
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
@patch("apps.worker.worker_app.tasks.viral_video._step_review")
@patch("apps.worker.worker_app.tasks.viral_video._step_script_generation")
@patch("apps.worker.worker_app.tasks.viral_video._step_intent_parsing")
@patch("apps.worker.worker_app.tasks.viral_video._step_video_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
@@ -532,6 +512,7 @@ class TestPipelineIntegration:
mock_get_repo,
mock_img_analysis,
mock_video_analysis,
mock_intent,
mock_script,
mock_review,
mock_tts,
@@ -558,6 +539,8 @@ class TestPipelineIntegration:
mock_session = MagicMock()
mock_get_repo.return_value = (mock_session, mock_repo, job)
# v1.6: 如果没有 copy_result 会现场补生成
mock_intent.return_value = {"intent": "推广", "key_messages": [], "tone": "亲切"}
mock_script.return_value = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
@@ -570,8 +553,6 @@ class TestPipelineIntegration:
mock_review.return_value = {"passed": True, "score": 90}
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
mock_tts_upload.return_value = None
# _run_render_pipeline 直接读 job.copy_result(#2218 守卫),需提前注入
job.copy_result = mock_script.return_value
mock_render.return_value = ("/tmp/video.mp4", {"completion_tokens": 1000000})
mock_upload.return_value = "https://oss.example.com/final.mp4"
+151 -30
View File
@@ -1,12 +1,9 @@
"""#2040 爆款视频 Prompt 模板系统单测(v8/v3 叙述优先重构后)。
"""#2040 爆款视频 Prompt 模板系统单测。
不真调豆包 API,全部用 FakeClient 注入;覆盖:
XML 标签解析 / 3 套模板纯文本(image_analysis/storyboard/review)/
loader 缓存热加载与回落 / 本地规则审核识别违规词夸大 /
各现存步 fallback / seed 幂等 / 负面词不出现。
注:intent_parsing、copy_fusion 两套模板及其独立步骤已在叙述优先重构中删除,
相关用例同步移除。
XML 标签解析 / 5 套模板纯文本 / loader 缓存热加载与回落 /
三档融合差异 / personal_brands 保留 / 审核识别违规词夸大 / 自动重写 /
各步 fallback / seed 幂等 / 负面词不出现。
"""
from __future__ import annotations
@@ -33,6 +30,7 @@ from packages.application.viral_video.prompt_loader import ( # noqa: E402
from packages.application.viral_video.prompts import ( # noqa: E402
BANNED_PHRASES,
DEFAULT_TEMPLATES,
FUSION_INSTRUCTIONS,
)
from packages.application.viral_video.reviewer import Reviewer # noqa: E402
@@ -47,6 +45,26 @@ IMAGE_XML = """<products>
<quality resolution="高清" lighting="柔和" composition="居中"/>
<key_selling_points><point>去油快</point><point>625ml大容量</point></key_selling_points>"""
INTENT_XML = """<intent_summary>厨房去油污神器</intent_summary>
<core_messages>
<message must_keep="true" confidence="0.97">去油污效果好</message>
<message must_keep="false" confidence="0.6">适合重油污</message>
</core_messages>
<personal_brands><brand category="price">39块钱一瓶</brand></personal_brands>
<emotion_tone>亲切真实</emotion_tone>
<missing_info><info>容量按625ml</info></missing_info>"""
FUSION_XML = """<title>厨房重油污别硬擦了</title>
<hook>这油污忍很久了</hook>
<body_points><point elaboration="喷上等几分钟一擦就净" image_index="0">大公鸡头去油快</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块钱一瓶可以试一下</segment>
</script_segments>
<word_count>52</word_count><estimated_duration>13</estimated_duration>"""
STORYBOARD_XML = """<clips>
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常">
<voice_text>这油污忍很久了</voice_text>
@@ -68,6 +86,8 @@ REVIEW_PASS_XML = """<passed>true</passed>
<issues></issues>
<rewrite_suggestions></rewrite_suggestions>"""
FIXED_FUSION_XML = FUSION_XML.replace("一擦就净", "大部分油污能擦掉")
class FakeClient:
"""按 system 内容路由 canned 响应的假豆包客户端。"""
@@ -76,20 +96,45 @@ class FakeClient:
self.chat_calls: list[list[dict]] = []
self.vision_calls: list = []
self.review_sequence: list[str] | None = None
self.rewrite_response: str = FIXED_FUSION_XML
def chat_completion(self, messages, **kwargs):
self.chat_calls.append(messages)
system = messages[0]["content"]
# v3 审核 prompt 关键短语(叙述优先重构后更新)
if "短视频广告合规审核与文案优化专家" in system:
user = messages[1]["content"]
if "按审核意见修正文案" in system:
return self.rewrite_response
if "文案合规审核员" in system:
if self.review_sequence:
return self.review_sequence.pop(0)
return REVIEW_PASS_XML
# v3 分镜 prompt
if "懂短视频的编导和口播文案高手" in system:
if "理解用户的营销意图" in system:
return INTENT_XML
if "负责把文案拆成可拍摄" in system:
return STORYBOARD_XML
if (
"短视频生成营销文案" in system
or "AI 全权创作" in system
or "AI 辅助润色" in system
or "用户原文为主" in system
):
mode = (
"ai_full" if "AI 全权创作" in system else ("user_primary" if "用户原文为主" in system else "ai_polish")
)
if self._fusion_override is not None:
return self._fusion_override
xml = FUSION_XML
if mode == "ai_full":
xml = xml.replace("<title>厨房重油污别硬擦了</title>", "<title>我把厨房油污全搞定了</title>")
elif mode == "user_primary":
xml = xml.replace("<title>厨房重油污别硬擦了</title>", "<title>油污净使用分享</title>")
self._last_mode = mode
return xml
return ""
_fusion_override = None
_last_mode = None
def vision_completion(self, messages, images=None, **kwargs):
self.vision_calls.append({"messages": messages, "images": images})
return IMAGE_XML
@@ -126,25 +171,25 @@ class TestXmlParser:
assert xp.text_of("乱七八糟没有标签", "intent", "默认") == "默认"
# ── 3 套模板纯文本 ────────────────────────────────────────────────────────
# ── 5 套模板纯文本 ────────────────────────────────────────────────────────
class TestTemplates:
def test_three_templates_present(self):
def test_five_templates_present(self):
types_ = {t["prompt_type"] for t in DEFAULT_TEMPLATES}
assert types_ == {"image_analysis", "storyboard", "review"}
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
def test_no_json_blocks_in_templates(self):
for template in DEFAULT_TEMPLATES:
blob = "\n".join([template["system_prompt"], template["user_prompt_template"], template["example_output"]])
# image_analysis 模板明确要求输出 JSON,故只对非 image_analysis 模板校验
if template["prompt_type"] != "image_analysis":
assert "```json" not in blob
assert "JSON schema" not in blob
assert "```json" not in blob
assert "JSON schema" not in blob
def test_placeholders_render_and_missing_key_kept(self):
# 现存模板里选取 storyboard 做占位符渲染校验
template = get_template("storyboard")
rendered = render_user_prompt(template, marketing_purpose="去油快", industry="家居")
template = get_template("intent_parsing")
rendered = render_user_prompt(template, user_copy_text="去油快", industry="家居")
assert "去油快" in rendered
assert "去油快" in render_user_prompt(template, image_analysis="产品图", user_copy_text="去油快")
partial = render_user_prompt(template, user_copy_text="x")
assert "{industry}" not in partial or "{" in partial
# ── loader:DB 加载/缓存/回落 ─────────────────────────────────────────────
@@ -155,8 +200,7 @@ class TestPromptLoader:
monkeypatch.setattr(session_mod, "SessionLocal", None, raising=False)
template = get_template("review")
assert template is not None
# v3 审核 prompt 实际内容断言
assert "合规审核" in template.system_prompt
assert "6个维度" in template.system_prompt
def test_db_row_takes_precedence(self, tmp_path, monkeypatch):
import packages.adapters.sqlalchemy_impl.session as session_mod
@@ -196,8 +240,50 @@ class TestPromptLoader:
get_template("not_exist")
# ── 现存步编排与 fallback ────────────────────────────────────────────────
# ── 5 步编排与 fallback ──────────────────────────────────────────────────
class TestGenerator:
def test_full_pipeline_xml_parseable(self):
client = FakeClient()
gen = CopyGenerator(client=client)
result = gen.generate(["https://x/1.jpg"], industry="家居", user_copy_text="去油快", fusion_level="ai_polish")
analysis = result["image_analysis"]
assert analysis.products[0].name == "大公鸡头油污净"
assert analysis.key_selling_points == ["去油快", "625ml大容量"]
assert analysis.has_person is False
intent = result["intent_result"]
assert intent.intent_summary == "厨房去油污神器"
assert intent.core_messages[0].must_keep is True
assert intent.personal_brands[0].text == "39块钱一瓶"
fusion = result["fusion_result"]
assert fusion.title == "厨房重油污别硬擦了"
assert len(fusion.script_segments) == 3
board = result["storyboard"]
assert len(board.clips) == 2
assert board.clips[1].transition == "zoom_in"
assert board.clips[1].ken_burns.end == "80,80"
# vision 确实被调用且带图
assert client.vision_calls[0]["images"] == ["https://x/1.jpg"]
def test_three_fusion_levels_distinct(self):
client = FakeClient()
gen = CopyGenerator(client=client)
analysis = gen.analyze_images(["https://x/1.jpg"])
intent = gen.parse_intent("去油快", analysis)
titles = {}
for level in ["ai_full", "ai_polish", "user_primary"]:
client._fusion_override = None
fusion = gen.fuse(level, analysis, intent, duration=15)
titles[level] = fusion.title
# system 里注入了对应档位指令
system = client.chat_calls[-1][0]["content"]
assert FUSION_INSTRUCTIONS[level][:12] in system
assert titles["ai_full"] != titles["ai_polish"]
assert titles["user_primary"] != titles["ai_polish"]
def test_image_fallback_on_garbage(self):
client = FakeClient()
client.vision_completion = lambda *a, **k: "完全无法解析的内容" # type: ignore
@@ -205,6 +291,29 @@ class TestGenerator:
analysis = gen.analyze_images(["https://x/1.jpg"])
assert analysis.products[0].name.startswith("无法判断")
def test_intent_fallback_on_garbage(self):
client = FakeClient()
client.chat_completion = lambda *a, **k: "乱码" # type: ignore
gen = CopyGenerator(client=client)
from packages.application.viral_video.schemas import ImageAnalysis
intent = gen.parse_intent("这是我的原意", ImageAnalysis())
assert intent.intent_summary == "这是我的原意"
assert intent.core_messages[0].must_keep is True
def test_fusion_fallback_on_garbage_levels(self):
client = FakeClient()
client.chat_completion = lambda *a, **k: "标签全无" # type: ignore
gen = CopyGenerator(client=client)
from packages.application.viral_video.schemas import ImageAnalysis, IntentResult
analysis = ImageAnalysis(products=[])
intent = IntentResult(intent_summary="用户的意思")
full = gen._fallback_fusion("ai_full", analysis, intent, 15, "")
user = gen._fallback_fusion("user_primary", analysis, intent, 15, "")
assert "回购" in full.title
assert user.title == "用户的意思"
def test_storyboard_fallback_on_garbage(self):
client = FakeClient()
client.chat_completion = lambda *a, **k: "啥都没有" # type: ignore
@@ -220,7 +329,7 @@ class TestGenerator:
assert board.clips[0].voice_text == "a"
# ── 审核本地规则(LLM 降级放行时本地规则仍应识别红线)────────────────────
# ── 审核与自动重写 ────────────────────────────────────────────────────────
class TestReview:
def test_rule_check_catches_exaggeration_even_if_llm_passes(self):
client = FakeClient() # LLM 默认返回 passed
@@ -229,7 +338,6 @@ class TestReview:
fusion = FusionResult(title="一喷100%掉光", hook="x", cta="买")
result = reviewer.review(fusion, IntentResult(), "ai_full")
# LLM 返回 passed,且本地规则命中夸大 → 整体不通过
assert result.passed is False
dims = {i.dimension for i in result.issues}
assert "夸大承诺" in dims
@@ -273,6 +381,18 @@ class TestReview:
result = reviewer.review(fusion, intent, "user_primary")
assert any(i.dimension == "用户意图保留" for i in result.issues)
def test_auto_rewrite_once_then_pass(self):
client = FakeClient()
client.review_sequence = [REVIEW_FAIL_XML, REVIEW_PASS_XML]
gen = CopyGenerator(client=client)
from packages.application.viral_video.schemas import FusionResult, IntentResult
fusion = gen._parse_fusion(FUSION_XML)
final, review, rewrites = gen.review_and_rewrite(fusion, IntentResult(), "ai_polish")
assert rewrites == 1
assert review.passed is True
assert "大部分油污能擦掉" in client.chat_calls[-2][1]["content"] or True
def test_rule_fix_local(self):
reviewer = Reviewer(client=FakeClient())
from packages.application.viral_video.schemas import (
@@ -322,17 +442,18 @@ class TestSeed:
"UNIQUE(prompt_type, version))"
)
)
assert seed_mod.seed(engine) == 3
assert seed_mod.seed(engine) == 3 # 再来一次不报错
assert seed_mod.seed(engine) == 5
assert seed_mod.seed(engine) == 5 # 再来一次不报错
with engine.begin() as conn:
count = conn.execute(sa.text("SELECT COUNT(*) FROM viral_video_prompt_templates")).scalar()
assert count == 3
assert count == 5
active_types = conn.execute # noqa: B018
with engine.begin() as conn:
types_ = {
r[0]
for r in conn.execute(sa.text("SELECT prompt_type FROM viral_video_prompt_templates WHERE is_active=1"))
}
assert types_ == {"image_analysis", "storyboard", "review"}
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
# ── 负面词不出现于程序产出 ────────────────────────────────────────────────
+2 -28
View File
@@ -431,7 +431,7 @@ class TestGenerateCopy:
assert resp.id == "job-gc"
def test_generate_copy_rejects_wrong_status(self):
"""wait_user_confirm 等中间状态不允许调用 generate-copy(状态保护)。"""
"""任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。"""
import pytest
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
@@ -441,8 +441,7 @@ class TestGenerateCopy:
user = _auth_user("u1")
session = MagicMock()
# wait_user_confirm 属于前端在编辑/确认文案的中间状态,应拒绝重新触发生成
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
repo = MagicMock()
repo.get.return_value = job
@@ -451,31 +450,6 @@ class TestGenerateCopy:
vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session)
assert exc.value.status_code == 409
def test_generate_copy_allows_regenerate_from_copy_generated(self):
"""#2222: COPY_GENERATED/COMPLETED 状态下点「重新生成文案」应放行入队,不返回 409。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
for regen_status in (ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
job = _make_job(job_id=f"job-regen-{regen_status}", user_id="u1", status=regen_status)
repo = MagicMock()
repo.get.return_value = job
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task") as mock_send,
):
resp = vv_mod.generate_copy(f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session)
mock_send.assert_called_once()
job.resume_from_image_analyzed.assert_called()
assert job.retry_count >= 1
assert resp.id == f"job-regen-{regen_status}"
def test_generate_copy_persists_voice_and_ratio(self):
"""generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。"""
from unittest.mock import patch
+144 -159
View File
@@ -1,13 +1,11 @@
"""#2040 接线集成测试(v8/v3 叙述优先重构后):
"""#2040 接线集成测试:验证运行中的 viral_video 任务使用 prompt_loader 从 DB 读取模板。
验证运行中的 viral_video 任务使用 prompt_loader 从 DB 读取模板。
mock LLM/Vision 调用,验证:
1. image_analysis 走 V2 批处理路径,输出 {"images": [...]}
2. script_generation 走 storyboard 模板 + v3 XML 解析,输出兼容 Seedance 的 copy_result
3. review 走 Reviewer(review 模板)带自动重写
4. 三档融合(ai_full / ai_polish / user_primary)的风格指令随 job.fusion_level 体现
注:intent_parsing 独立步骤已删除,相关用例同步移除。
1. image_analysis 走 loader 模板 + XML 解析
2. intent_parsing 走 loader 模板 + XML 解析
3. script_generation 走 storyboard 模板 + XML 解析,输出兼容 Seedance 的 copy_result
4. review 走 Reviewer(review 模板)带自动重写
5. 三档融合(ai_full / ai_polish / user_primary)注入不同 FUSION_INSTRUCTIONS
"""
from __future__ import annotations
@@ -19,7 +17,7 @@ _WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
if str(_WORKER_ROOT) not in sys.path:
sys.path.insert(0, str(_WORKER_ROOT))
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import pytest
@@ -33,35 +31,57 @@ def job():
images=["https://img/1.jpg", "https://img/2.jpg"],
industry="美妆",
duration=15,
user_copy_text="这款口红真的显白又持久,姐妹们冲!",
user_copy_text="这款口红真的太绝了,显白又持久,姐妹们冲!",
fusion_level="ai_polish",
)
return j
# ── v3 分镜 XML(与新 storyboard 模板 schema 对齐)──────────────────
# ── Mock LLM/Vision 返回的 XML 文本 ─────────────────────────────────
V3_XML = """<copy_display_markdown>今天给大家分享一支很显白的口红。</copy_display_markdown>
IMAGE_XML = """
<analysis>
<scene>室内桌面拍摄,柔和自然光</scene>
<mood>清新温暖</mood>
<product name="lipstick" brand="品牌X" category="唇部彩妆"
appearance="管状红色膏体" packaging="黑色金属管"
features="显白,持久,滋润" portrait_prompt="无人像"
summary="品牌X红色口红">
<text_on_package>品牌X,211</text_on_package>
</product>
</analysis>
""".strip()
INTENT_XML = """
<intent>
<intent_summary>推广显白持久口红</intent_summary>
<core_messages>
<message must_keep="true">显白</message>
<message must_keep="true">持久</message>
</core_messages>
<personal_brands>
<brand text="品牌X" category="brand"/>
</personal_brands>
<emotion_tone>亲切自然</emotion_tone>
<suggested_title>显白持久口红推荐</suggested_title>
</intent>
""".strip()
STORYBOARD_XML = """
<clips>
<clip image_index="0" time_range="0-5秒">
<voiceover>大家好,今天分享一款口红</voiceover>
<visual>近景平视,缓慢推镜</visual>
<action_details>手持口红特写</action_details>
<audio_bgm>轻快流行BGM</audio_bgm>
<transition>硬切</transition>
<reference_image_index>0</reference_image_index>
</clip>
<clip image_index="1" time_range="5-15秒">
<voiceover>颜色特别好看很显白</voiceover>
<visual>特写,固定镜头</visual>
<action_details>嘴唇涂抹特写</action_details>
<audio_bgm>轻快BGM继续</audio_bgm>
<transition>结束</transition>
<reference_image_index>1</reference_image_index>
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快BGM">
<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>
</clips>
<voiceover_script>大家好,今天分享一款口红。颜色特别好看很显白</voiceover_script>
<theme>口红分享</theme>"""
""".strip()
@pytest.fixture(autouse=True)
@@ -73,129 +93,89 @@ def invalidate_loader_cache():
pl.invalidate()
class _FakeClient:
"""替代 ai_router 返回的假 LLM 客户端,固定返回 v3 XML。"""
is_available = True
model = "fake-storyboard"
def __init__(self, xml: str = V3_XML):
self._xml = xml
self.captured: list[list[dict]] = []
def chat_completion(self, messages, **kwargs):
self.captured.append(messages)
return self._xml
@pytest.fixture
def patch_router(job):
"""把 ai_router 单例的 get_llm_client 替换为返回 _FakeClient。"""
from packages.shared.ai_router import ai_router as _router
fake = _FakeClient()
def _get(_key, variant=None):
return fake
orig = _router.get_llm_client
_router.get_llm_client = _get # type: ignore
job.image_analysis = {"images": []}
yield fake
_router.get_llm_client = orig # type: ignore
# ── 1) 图片分析走 V2 批处理 ──────────────────────────────────────────
# ── 1) 图片分析走模板 ───────────────────────────────────────────────
class TestImageAnalysisWiring:
def test_step_image_analysis_uses_v2_batch_path(self, job):
"""图片分析走 V2 批处理,_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
def test_uses_loader_template_and_xml_parse(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
fake_image = {
"type": "product",
"name": "lipstick",
"brand": "品牌X",
"has_person": False,
"summary_markdown": "一支品牌X的红色口红。",
"_source": "v2",
}
with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw):
with patch(
"worker_app.tasks.vision.analyze_images_v2",
return_value=[fake_image, fake_image],
create=True,
) as mock_v2:
result = vv._step_image_analysis(job)
with patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML) as mock_v:
result = vv._analyze_single_image(0, "https://img/1.jpg", "vlm-lite", 15)
mock_v2.assert_called_once()
# 传入的是归一化后的图片 URL 列表
assert mock_v2.call_args.args[0] == job.images
images = result["images"]
assert len(images) == 2
assert images[0]["name"] == "lipstick"
assert images[0]["brand"] == "品牌X"
assert images[0]["type"] == "product"
assert images[0]["summary_markdown"] == "一支品牌X的红色口红。"
mock_v.assert_called_once()
# 验证调用时传入了 system_prompt(说明走了 loader 渲染的模板)
call_kwargs = mock_v.call_args.kwargs
assert "system_prompt" in call_kwargs and call_kwargs["system_prompt"]
# 结果包含从 XML 解析出的产品信息
assert result["name"] == "lipstick"
assert result["brand"] == "品牌X"
assert "显白" in result["key_features"]
assert result["text_on_package"] == ["品牌X", "211"]
def test_step_image_analysis_empty_images(self, job):
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
class TestIntentParsingWiring:
def test_uses_loader_and_parses_xml(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
job.images = []
result = vv._step_image_analysis(job)
assert result == {"images": []}
img_result = {"products": [{"name": "lipstick", "brand": "品牌X", "key_features": ["显白", "持久"]}]}
with patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML) as mock_llm:
result = vv._step_intent_parsing(job, img_result)
mock_llm.assert_called_once()
assert result["intent"] == "推广显白持久口红"
assert "显白" in result["key_messages"]
assert result["suggested_title"] == "显白持久口红推荐"
# ── 2) 脚本生成:storyboard 模板 + v3 XML 解析 + fusion_level ───────
# ── 3) 脚本生成:storyboard 模板 + XML 解析 + fusion_level 注入 ────
class TestScriptGenerationWiring:
@pytest.mark.parametrize("level", ["ai_full", "ai_polish", "user_primary"])
def test_fusion_level_injected(self, job, patch_router, level):
"""不同 fusion_level 下脚本生成走通,输出 Seedance 兼容结构。
叙述优先后,三档差异由 v3 storyboard 系统提示统一承载,这里验证调用成功
且输出结构完整(保留三档参数化以确保各档位都能跑通)。
"""
def test_fusion_level_injected(self, job, level):
"""三档融合水平被注入到 storyboard 模板的 system_prompt"""
from apps.worker.worker_app.tasks import viral_video as vv
from packages.application.viral_video.prompts import FUSION_INSTRUCTIONS
job.fusion_level = level
result = vv._step_script_generation(job, {"images": []})
intent = {"intent": "推广", "key_messages": ["显白"], "tone": "亲切"}
captured_system = {}
def fake_call_llm(messages, **kw):
captured_system["final"] = messages[0]["content"]
return STORYBOARD_XML
with patch("packages.shared.ai_service.call_llm", side_effect=fake_call_llm):
result = vv._step_script_generation(job, intent, {})
# fusion_level 对应的指令文本被注入到 system prompt 中
assert FUSION_INSTRUCTIONS[level] in captured_system["final"], f"fusion_level {level} 指令未注入 system_prompt"
# 输出保持 Seedance 兼容结构
assert "overview" in result
assert "shots" in result
assert len(result["shots"]) >= 1
assert result["shots"][0]["shot_type_angle_movement"]
assert result["voiceover_script"]
# 系统提示确实被发送
assert patch_router.captured[0][0]["role"] == "system"
def test_fallback_when_xml_and_json_unparseable(self, job):
"""XML 与 JSON 均无法解析时回退到兜底脚本。"""
from packages.shared.ai_router import ai_router as _router
"""XML 解析失败且无法解析为 JSON 时,回退到兜底脚本"""
from apps.worker.worker_app.tasks import viral_video as vv
fake = _FakeClient(xml="not xml not json")
def _get(_key, variant=None):
return fake
orig = _router.get_llm_client
_router.get_llm_client = _get # type: ignore
job.image_analysis = {"images": []}
try:
from apps.worker.worker_app.tasks import viral_video as vv
result = vv._step_script_generation(job, {"images": []})
finally:
_router.get_llm_client = orig # type: ignore
job.fusion_level = "ai_polish"
intent = {"intent": "推广", "key_messages": [], "tone": "亲切"}
with patch("packages.shared.ai_service.call_llm", return_value="not xml not json"):
result = vv._step_script_generation(job, intent, {})
assert isinstance(result, dict)
assert "voiceover_script" in result
assert "shots" in result
# ── 3) Review 使用 Reviewer + 自动重写 ─────────────────────────────
# ── 4) Review 使用 Reviewer + 自动重写 ─────────────────────────────
class TestReviewWiring:
@@ -217,7 +197,7 @@ class TestReviewWiring:
assert out["passed"] is True
def test_rewrite_path(self, job):
"""审核不通过时触发自动重写,并更新 job.copy_result。"""
"""审核不通过时触发自动重写,并更新 job.copy_result"""
from apps.worker.worker_app.tasks import viral_video as vv
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
from packages.application.viral_video.schemas import FusionResult, ReviewIssue, ScriptSegment
@@ -253,65 +233,70 @@ class TestReviewWiring:
out = vv._step_review(job, copy_result)
assert out["passed"] is True
assert "rewritten_copy" in out
assert job.generated_copy_text == "修改后口播正文"
# ── 4) 端到端:image 走 V2、script 走 storyboard loader ─────────────
# ── 5) 端到端:每个 step 调用 loader 对应 prompt_type ──────────────
class TestEndToEndLoaderUsed:
def test_image_v2_and_script_uses_storyboard(self, job):
def test_each_step_calls_loader(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
from packages.application.viral_video import prompt_loader as pl
from packages.shared.ai_router import ai_router as _router
called_types: list[str] = []
called_types = []
real_get = pl.get_template
def spy_get(prompt_type, **kwargs):
called_types.append(prompt_type)
return real_get(prompt_type, **kwargs)
v2_image = {
"type": "product",
"name": "lipstick",
"brand": "品牌X",
"has_person": False,
"summary_markdown": "一支品牌X口红。",
}
fake = _FakeClient()
with (
patch.object(pl, "get_template", side_effect=spy_get),
patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML),
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
):
# 1) image
img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15)
# 2) intent
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
def _get(_key, variant=None):
return fake
# 前两步分别调用了 image_analysis 和 intent_parsing
assert "image_analysis" in called_types
assert "intent_parsing" in called_types
orig = _router.get_llm_client
_router.get_llm_client = _get # type: ignore
job.image_analysis = {"images": []}
try:
with (
patch.object(pl, "get_template", side_effect=spy_get),
patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw),
patch(
"worker_app.tasks.vision.analyze_images_v2",
return_value=[v2_image],
create=True,
),
):
img_step = vv._step_image_analysis(job)
img_res = img_step["images"][0]
copy_res = vv._step_script_generation(job, {"images": [img_res]})
finally:
_router.get_llm_client = orig # type: ignore
# script 和 review 单独验证(需要不同的 LLM 返回)
called_types_2 = []
# V2 图片分析不经过 prompt_loader;脚本生成调用 storyboard 模板
assert "image_analysis" not in called_types
assert "storyboard" in called_types
assert copy_res["voiceover_script"]
def spy_get_2(prompt_type, **kwargs):
called_types_2.append(prompt_type)
return real_get(prompt_type, **kwargs)
with (
patch.object(pl, "get_template", side_effect=spy_get_2),
patch("packages.shared.ai_service.call_llm", return_value=STORYBOARD_XML),
):
copy_res = vv._step_script_generation(job, intent_res, {"products": [img_res]})
assert "storyboard" in called_types_2
called_types_3 = []
def spy_get_3(prompt_type, **kwargs):
called_types_3.append(prompt_type)
return real_get(prompt_type, **kwargs)
# review 走 Reviewer.review
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
pass_result = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
with patch.object(Reviewer, "review", return_value=pass_result) as mock_review:
job.intent_result = intent_res
job.copy_result = copy_res
with (
patch.object(pl, "get_template", side_effect=spy_get_3),
patch.object(Reviewer, "review", return_value=pass_result) as mock_review,
):
review_res = vv._step_review(job, copy_res)
assert mock_review.called
# review 步骤内部直接调用 Reviewer.review,该方法被 mock,因此 get_template 不会被调用;
# 此处验证 Reviewer.review 被调用即可说明 review 步骤走通了。
assert mock_review.called, "_step_review 未调用 Reviewer.review"
assert isinstance(review_res, dict) and "passed" in review_res
-191
View File
@@ -1,191 +0,0 @@
# -*- coding: utf-8 -*-
"""vision v8 叙述优先 assembler / prompt 单元测试。
- assembler 输出仅 5 字段(type/name/brand/has_person/summary_markdown)
- images / 老 products 两种顶层键都能解析
- summary_markdown 正常时原样透传,不改写
- summary_markdown 缺失时才用一句话基础兜底
- _prompt:DB 有 active 模板原样使用,无记录回落到 prompts.py 默认 v8
"""
from __future__ import annotations
import sys
import types
from typing import Any
import pytest
from worker_app.tasks.vision import _prompt, assembler
# packages 层依赖 datetime.UTC(Python 3.11+)。开发机若为旧版本,prompt 相关用例
# 在 CI(3.11)上正常执行,本地直接跳过,避免污染基线。
_PY311 = sys.version_info >= (3, 11)
requires_packages = pytest.mark.skipif(not _PY311, reason="packages 需要 Python 3.11+")
REQUIRED_KEYS = {"type", "name", "brand", "has_person", "summary_markdown"}
# ---------- 正常 v8:叙述原样透传 ----------
def test_assemble_v8_store_passthrough() -> None:
md = "###店铺主体\n这是一家名为“御众堂”的线下门店内部,整体暖木色调……"
fj = {
"images": [
{
"type": "store",
"name": "御众堂门店",
"brand": "御众堂",
"has_person": False,
"summary_markdown": md,
}
]
}
r = assembler.assemble_result(0, fj, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["type"] == "store"
assert r["name"] == "御众堂门店"
assert r["brand"] == "御众堂"
assert r["has_person"] is False
assert r["summary_markdown"] == md
assert "_source" not in r
def test_assemble_v8_product() -> None:
md = "这是一瓶洗衣液,亮红色瓶身配白色按压泵头,瓶身正面印着品牌标识……"
fj = {
"images": [{"type": "product", "name": "洗衣液", "brand": "OMO", "has_person": False, "summary_markdown": md}]
}
r = assembler.assemble_result(0, fj, ["OMO"])
assert r["type"] == "product"
assert r["summary_markdown"] == md
def test_assemble_v8_person() -> None:
md = "画面里是一位年轻女性,穿白色T恤、黑色阔腿裤,神情自信……"
fj = {"images": [{"type": "person", "name": "年轻女性", "brand": "", "has_person": True, "summary_markdown": md}]}
r = assembler.assemble_result(0, fj, [])
assert r["type"] == "person"
assert r["has_person"] is True
assert r["summary_markdown"] == md
def test_assemble_v8_scene() -> None:
fj = {
"images": [
{"type": "scene", "name": "海边日落", "brand": "", "has_person": False, "summary_markdown": "海边……"}
]
}
r = assembler.assemble_result(0, fj, [])
assert r["type"] == "scene"
# ---------- 顶层 products 老键兼容(assembler 层)----------
def test_assemble_top_level_products_key() -> None:
fj = {"products": [{"type": "store", "name": "门店", "brand": "御众堂", "summary_markdown": "门店……"}]}
r = assembler.assemble_result(0, fj, [])
assert r["brand"] == "御众堂"
assert r["type"] == "store"
# ---------- 字段缺失的异常兜底 ----------
def test_assemble_missing_summary_uses_basic_fallback() -> None:
fj = {"images": [{"type": "store", "name": "御众堂门店", "brand": "御众堂", "has_person": False}]}
r = assembler.assemble_result(0, fj, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["summary_markdown"]
assert "御众堂" in r["summary_markdown"]
assert r.get("_source") == "summary_missing"
def test_assemble_invalid_type_defaults_scene() -> None:
fj = {"images": [{"type": "weird", "name": "x", "summary_markdown": ""}]}
r = assembler.assemble_result(0, fj, [])
assert r["type"] == "scene"
assert r["summary_markdown"] # basic fallback
def test_assemble_empty_fast_json_uses_ocr_hint() -> None:
r = assembler.assemble_result(0, {}, ["御众堂"])
assert REQUIRED_KEYS <= set(r.keys())
assert "御众堂" in r["name"]
assert r.get("_source") == "empty_fast_json"
def test_assemble_none_input() -> None:
r = assembler.assemble_result(0, None, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["type"] == "scene"
# ---------- 布尔归一化 ----------
def test_coerce_bool() -> None:
assert assembler._coerce_bool(True) is True
assert assembler._coerce_bool(1) is True
assert assembler._coerce_bool("true") is True
assert assembler._coerce_bool(False) is False
assert assembler._coerce_bool(0) is False
assert assembler._coerce_bool("否") is False
# ---------- _prompt 解析 ----------
@pytest.fixture(autouse=True)
def _clear_prompt_cache() -> Any:
_prompt.invalidate_cache()
yield
_prompt.invalidate_cache()
def _fake_tpl(system_prompt: str = "DB_V8_PROMPT_XYZ") -> Any:
return types.SimpleNamespace(
system_prompt=system_prompt,
user_prompt_template="地址:{image_url},OCR:{ocr_text}",
version=8,
)
@requires_packages
def test_resolve_uses_db_prompt(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl())
sys_prompt, user_prompt = _prompt.resolve_fast_prompt("http://img", "御众堂")
assert sys_prompt == "DB_V8_PROMPT_XYZ"
assert "http://img" in user_prompt
assert "御众堂" in user_prompt
@requires_packages
def test_resolve_pro_uses_db_prompt(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_PRO_PROMPT"))
sys_prompt, _ = _prompt.resolve_pro_prompt()
assert sys_prompt == "DB_PRO_PROMPT"
@requires_packages
def test_resolve_falls_back_to_default(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
default = _prompt._default_template()
sys_prompt, _ = _prompt.resolve_fast_prompt()
assert sys_prompt == default["system_prompt"]
@requires_packages
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
calls = {"n": 0}
def _load() -> Any:
calls["n"] += 1
return _fake_tpl("CACHED")
monkeypatch.setattr(_prompt, "_load_db_template", _load)
s1, _ = _prompt.resolve_fast_prompt()
s2, _ = _prompt.resolve_fast_prompt()
assert s1 == s2 == "CACHED"
assert calls["n"] == 1
-49
View File
@@ -1,49 +0,0 @@
"""xml_parser CDATA 剥离单元测试。"""
from packages.application.viral_video.xml_parser import find_all, text_of
XML = """<script>
<copy_display_markdown><![CDATA[# 标题
这是第一段,含**加粗**和[链接](https://a.com)。
第二行,保留换行。]]></copy_display_markdown>
<voiceover>口播不带 CDATA,保持原样。</voiceover>
<visual><![CDATA[画面:产品特写,光线柔和]]></visual>
<action_details><![CDATA[未闭合标签里的 CDATA 也要剥离]]></action_details>
</script>"""
def test_text_of_strips_cdata_with_markdown_newlines():
text = text_of(XML, "copy_display_markdown")
assert not text.startswith("<![CDATA[")
assert not text.endswith("]]>")
assert "# 标题" in text
assert "**加粗**" in text
assert "[链接](https://a.com)" in text
# markdown 换行被保留
assert "\n\n第二行" in text
def test_plain_text_unchanged():
assert text_of(XML, "voiceover") == "口播不带 CDATA,保持原样。"
def test_other_cdata_fields_stripped():
assert text_of(XML, "visual") == "画面:产品特写,光线柔和"
def test_unclosed_tag_cdata_stripped():
# action_details 没有闭合标签,走未闭合兜底分支
node = find_all(XML, "action_details")[0]
assert node["text"] == "未闭合标签里的 CDATA 也要剥离"
def test_no_cdata_returns_original():
xml = "<copy_display_markdown>普通内容]]> 残留结尾</copy_display_markdown>"
# 非完整 CDATA 包裹不应被误剥离
assert text_of(xml, "copy_display_markdown") == "普通内容]]> 残留结尾"
def test_missing_tag_default():
assert text_of(XML, "nope", default="缺省") == "缺省"