Compare commits
42 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 1f6d9fc862 | |||
| 3acc309b51 | |||
| bee59b7f27 | |||
| 4bd4a6c390 | |||
| aa898ff0c8 | |||
| fb0d429bd7 | |||
| 0510d101aa | |||
| 606f6988b5 | |||
| 733d8bb75c | |||
| e83048b7ec | |||
| ade6593d66 | |||
| a3a8d2561d | |||
| 7b6dfdc29f | |||
| b5e225e62a | |||
| 2141ddb19b | |||
| 3503542ec2 | |||
| 1a6f04b258 | |||
| ece5f7d8a7 | |||
| b023988402 | |||
| 4ba5b126b6 | |||
| cc8a3caba5 | |||
| fe36a05c51 | |||
| 12a8701283 | |||
| 176a5cbfe3 | |||
| c5e69db50d | |||
| 44bbe0f0e5 | |||
| 92680c47f9 | |||
| b335fbbcce | |||
| cc542f27d9 | |||
| d42ab5ffa8 | |||
| 74896727bb | |||
| 1f6d10f861 | |||
| 3ee5a4042d | |||
| a979af1488 | |||
| 77e19a6b44 | |||
| e4c3f9a046 | |||
| 902effc1f9 | |||
| a408cfdc97 | |||
| 9464322710 | |||
| aa1f318308 | |||
| 3a59948f53 | |||
| 3e6f87a8b5 |
@@ -0,0 +1,242 @@
|
|||||||
|
# -*- 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"
|
||||||
|
)
|
||||||
|
)
|
||||||
Executable
+116
@@ -0,0 +1,116 @@
|
|||||||
|
# -*- 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"
|
||||||
|
)
|
||||||
|
)
|
||||||
@@ -352,6 +352,28 @@ def confirm_copy(
|
|||||||
if not isinstance(job.copy_result, dict) or not job.copy_result:
|
if not isinstance(job.copy_result, dict) or not job.copy_result:
|
||||||
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
|
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
|
||||||
|
|
||||||
|
# Bug1 fix: 用户 confirm 时允许修改 video_model/video_resolution/video_ratio/duration
|
||||||
|
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||||
|
old_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||||
|
old_ratio = getattr(job, "video_ratio", "9:16") or "9:16"
|
||||||
|
old_model = getattr(job, "video_model", None) or "seedance-2.5"
|
||||||
|
|
||||||
|
if request.duration is not None:
|
||||||
|
job.duration = max(5, min(30, int(request.duration)))
|
||||||
|
if request.video_resolution is not None:
|
||||||
|
job.video_resolution = request.video_resolution
|
||||||
|
if request.video_ratio is not None:
|
||||||
|
job.video_ratio = request.video_ratio
|
||||||
|
if request.video_model is not None:
|
||||||
|
job.video_model = request.video_model
|
||||||
|
|
||||||
|
param_changed = (
|
||||||
|
(request.duration is not None and int(request.duration) != old_duration)
|
||||||
|
or (request.video_resolution is not None and request.video_resolution != old_resolution)
|
||||||
|
or (request.video_ratio is not None and request.video_ratio != old_ratio)
|
||||||
|
or (request.video_model is not None and request.video_model != old_model)
|
||||||
|
)
|
||||||
|
|
||||||
# 积分预扣(已扣过/重试任务跳过)
|
# 积分预扣(已扣过/重试任务跳过)
|
||||||
from app.config import settings as _settings
|
from app.config import settings as _settings
|
||||||
|
|
||||||
@@ -359,7 +381,49 @@ def confirm_copy(
|
|||||||
already_paid = (float(getattr(job, "credits_prepaid", 0) or 0) > 0) or (
|
already_paid = (float(getattr(job, "credits_prepaid", 0) or 0) > 0) or (
|
||||||
float(getattr(job, "credits_cost", 0) or 0) > 0
|
float(getattr(job, "credits_cost", 0) or 0) > 0
|
||||||
)
|
)
|
||||||
if not already_paid:
|
if param_changed and already_paid:
|
||||||
|
# 参数变更:回退旧预扣,按新参数重新预扣
|
||||||
|
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||||
|
from packages.domain.points_service import PointsService
|
||||||
|
|
||||||
|
old_w, old_h = resolve_video_dimensions(old_resolution, old_ratio)
|
||||||
|
old_est = calculate_viral_video_credits(old_duration, old_w, old_h, old_model)
|
||||||
|
new_w, new_h = resolve_video_dimensions(
|
||||||
|
getattr(job, "video_resolution", "720p") or "720p",
|
||||||
|
job.video_ratio or "9:16",
|
||||||
|
)
|
||||||
|
new_est = calculate_viral_video_credits(
|
||||||
|
int(job.duration or 15), new_w, new_h, job.video_model or "seedance-2.5"
|
||||||
|
)
|
||||||
|
svc = PointsService()
|
||||||
|
# 退回旧预扣
|
||||||
|
if getattr(job, "credits_transaction_id", None):
|
||||||
|
svc.refund_points(
|
||||||
|
user_id=authenticated_user.user.id,
|
||||||
|
amount=float(job.credits_prepaid),
|
||||||
|
source="viral_video",
|
||||||
|
db=session,
|
||||||
|
ref_id=job.credits_transaction_id,
|
||||||
|
description="confirm-copy 参数变更退还旧预扣",
|
||||||
|
)
|
||||||
|
# 预扣新金额
|
||||||
|
if new_est > 0:
|
||||||
|
res = svc.deduct_viral_video(authenticated_user.user.id, new_est, job.id, session)
|
||||||
|
if not res.get("success"):
|
||||||
|
balance = res.get("balance", 0)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=402,
|
||||||
|
detail={
|
||||||
|
"code": "INSUFFICIENT_POINTS",
|
||||||
|
"message": f"积分不足,需要 {new_est} 积分,当前余额 {balance}",
|
||||||
|
"required": new_est,
|
||||||
|
"balance": balance,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
job.credits_prepaid = new_est
|
||||||
|
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||||
|
logger.info("[爆款视频][confirm-copy] 参数变更,积分重算: old=%d new=%d job_id=%s", old_est, new_est, job.id)
|
||||||
|
elif not already_paid:
|
||||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||||
from packages.domain.points_service import PointsService
|
from packages.domain.points_service import PointsService
|
||||||
|
|
||||||
@@ -890,13 +954,28 @@ async def viral_video_websocket(websocket: WebSocket, job_id: str) -> None:
|
|||||||
job = job_repo.get(job_id)
|
job = job_repo.get(job_id)
|
||||||
if job is not None:
|
if job is not None:
|
||||||
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
|
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||||
|
# #P0: 初始快照必须包含前端重连/刷新所需的业务字段,
|
||||||
|
# 结构对齐 worker 推送的 image_analyzed / copy_generated 事件。
|
||||||
|
data: dict = {"status": status_val}
|
||||||
|
ia = getattr(job, "image_analysis", None)
|
||||||
|
if isinstance(ia, dict) and ia:
|
||||||
|
data["image_analysis"] = ia
|
||||||
|
cr = _build_copy_result(job)
|
||||||
|
if isinstance(cr, dict) and cr:
|
||||||
|
data["copy_result"] = cr
|
||||||
|
gct = getattr(job, "generated_copy_text", "") or ""
|
||||||
|
if gct:
|
||||||
|
data["generated_copy_text"] = gct
|
||||||
|
sb = getattr(job, "storyboard", None) or []
|
||||||
|
if sb:
|
||||||
|
data["storyboard"] = sb
|
||||||
initial = {
|
initial = {
|
||||||
"type": "viral_video:progress",
|
"type": "viral_video:progress",
|
||||||
"job_id": job_id,
|
"job_id": job_id,
|
||||||
"stage": _stage_from_status(job),
|
"stage": _stage_from_status(job),
|
||||||
"progress": _estimate_progress(job),
|
"progress": _estimate_progress(job),
|
||||||
"message": _initial_message(job),
|
"message": _initial_message(job),
|
||||||
"data": {"status": status_val},
|
"data": data,
|
||||||
}
|
}
|
||||||
await websocket.send_json(initial)
|
await websocket.send_json(initial)
|
||||||
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
|
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
|
||||||
|
|||||||
@@ -167,6 +167,10 @@ class ConfirmCopyRequest(BaseModel):
|
|||||||
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
|
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
|
||||||
|
|
||||||
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
|
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
|
||||||
|
video_model: str | None = Field(default=None, description="用户选定的视频生成模型(confirm时可选)")
|
||||||
|
video_resolution: str | None = Field(default=None, description="用户选定的分辨率(confirm时可选)")
|
||||||
|
video_ratio: str | None = Field(default=None, description="用户选定的比例(confirm时可选)")
|
||||||
|
duration: int | None = Field(default=None, ge=5, le=30, description="用户选定的时长秒数(confirm时可选,5~30)")
|
||||||
|
|
||||||
|
|
||||||
class ConfirmIntentRequest(BaseModel):
|
class ConfirmIntentRequest(BaseModel):
|
||||||
|
|||||||
@@ -223,7 +223,47 @@ class LipsyncService:
|
|||||||
if timings:
|
if timings:
|
||||||
job.sentence_timings = timings
|
job.sentence_timings = timings
|
||||||
|
|
||||||
# 4. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
# 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
|
||||||
use_gpu = False
|
use_gpu = False
|
||||||
if self.settings.use_gpu_lipsync:
|
if self.settings.use_gpu_lipsync:
|
||||||
try:
|
try:
|
||||||
@@ -805,6 +845,30 @@ class LipsyncService:
|
|||||||
if job.status in (STATUS_COMPLETED, "failed"):
|
if job.status in (STATUS_COMPLETED, "failed"):
|
||||||
return job
|
return job
|
||||||
|
|
||||||
|
# Ditto 异步路径:mediakit_task_id 以 "ditto:" 开头,由 Celery 任务异步更新
|
||||||
|
# 不做 MediaKit 轮询,只检查是否卡住太久(>10 分钟)则标失败
|
||||||
|
if job.mediakit_task_id and job.mediakit_task_id.startswith("ditto:"):
|
||||||
|
if job.status in ("processing", "submitted"):
|
||||||
|
_now = datetime.now(UTC)
|
||||||
|
_upd = job.updated_at
|
||||||
|
if _upd is not None and _upd.tzinfo is None:
|
||||||
|
_upd = _upd.replace(tzinfo=UTC)
|
||||||
|
stale_minutes = 10
|
||||||
|
if _upd and (_now - _upd).total_seconds() > stale_minutes * 60:
|
||||||
|
logger.warning(
|
||||||
|
"Ditto 异步任务超时(>%d 分钟),标记失败: job_id=%s",
|
||||||
|
stale_minutes,
|
||||||
|
job_id,
|
||||||
|
)
|
||||||
|
job.status = "failed"
|
||||||
|
job.error_message = f"Ditto 处理超时(>{stale_minutes} 分钟)"
|
||||||
|
job.error_code = "DittoTimeout"
|
||||||
|
job.completed_at = _now
|
||||||
|
job.updated_at = _now
|
||||||
|
self.db.commit()
|
||||||
|
self._refund_lip_sync(job)
|
||||||
|
return job
|
||||||
|
|
||||||
# GPU 异步路径:mediakit_task_id 以 "gpu:" 开头,由 Celery 任务异步更新
|
# GPU 异步路径:mediakit_task_id 以 "gpu:" 开头,由 Celery 任务异步更新
|
||||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>30 分钟)则标失败
|
# 不做 MediaKit 轮询,只检查是否卡住太久(>30 分钟)则标失败
|
||||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"):
|
if job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"):
|
||||||
|
|||||||
@@ -0,0 +1,351 @@
|
|||||||
|
"""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_emotion_service import get_ditto_emotion_service
|
||||||
|
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),
|
||||||
|
)
|
||||||
|
# ── LLM 情绪分析(#2076 后续):生成 emo_timeline ──
|
||||||
|
emo_timeline = ""
|
||||||
|
try:
|
||||||
|
emo_svc = get_ditto_emotion_service()
|
||||||
|
if emo_svc.enabled and script:
|
||||||
|
# 探测音频时长用于时间对齐
|
||||||
|
try:
|
||||||
|
from packages.domain.sentence_timings import probe_audio_duration
|
||||||
|
from packages.shared.url_security import safe_download_bytes
|
||||||
|
|
||||||
|
audio_bytes = safe_download_bytes(
|
||||||
|
audio_url,
|
||||||
|
allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav", "audio/mp3"),
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
audio_duration = probe_audio_duration(audio_bytes)
|
||||||
|
except Exception as audio_exc:
|
||||||
|
logger.warning("[ditto_task] 音频时长探测失败,emo_timeline 降级空: %s", audio_exc)
|
||||||
|
audio_duration = 0.0
|
||||||
|
if audio_duration > 0:
|
||||||
|
sentence_timings = getattr(job, "sentence_timings", None)
|
||||||
|
emo_timeline = emo_svc.build_timeline(
|
||||||
|
text=script,
|
||||||
|
audio_duration=audio_duration,
|
||||||
|
sentence_timings=sentence_timings,
|
||||||
|
)
|
||||||
|
if emo_timeline:
|
||||||
|
logger.info("[ditto_task] 情绪时间线已生成: segments=%d", len(emo_timeline) // 50)
|
||||||
|
except Exception as emo_exc:
|
||||||
|
logger.warning("[ditto_task] 情绪分析异常(降级中性): %s", emo_exc)
|
||||||
|
emo_timeline = ""
|
||||||
|
client = get_ditto_client()
|
||||||
|
result = client.generate_and_persist(
|
||||||
|
job_id=job_id,
|
||||||
|
user_id=user_id,
|
||||||
|
audio_url=audio_url,
|
||||||
|
script=script,
|
||||||
|
emo_timeline=emo_timeline,
|
||||||
|
# video_url 不传则用默认模板
|
||||||
|
)
|
||||||
|
|
||||||
|
# Ditto 返回的 MP4 自带音频,签名 OSS URL(7天有效)后标记完成
|
||||||
|
job.output_video_url = _sign_media_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()
|
||||||
@@ -261,7 +261,33 @@ def tts_synthesize_and_submit(
|
|||||||
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
|
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
|
||||||
)
|
)
|
||||||
|
|
||||||
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
|
# 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 的耦合)
|
||||||
audio_url = _sign_media_url(job.audio_url)
|
audio_url = _sign_media_url(job.audio_url)
|
||||||
video_url = _sign_media_url(job.video_url)
|
video_url = _sign_media_url(job.video_url)
|
||||||
|
|
||||||
|
|||||||
Generated
+12
@@ -14,6 +14,7 @@
|
|||||||
"axios": "^1.7.2",
|
"axios": "^1.7.2",
|
||||||
"classnames": "^2.5.1",
|
"classnames": "^2.5.1",
|
||||||
"dayjs": "^1.11.23",
|
"dayjs": "^1.11.23",
|
||||||
|
"marked": "^12.0.2",
|
||||||
"mp4box": "^2.4.1",
|
"mp4box": "^2.4.1",
|
||||||
"react": "^18.3.1",
|
"react": "^18.3.1",
|
||||||
"react-dom": "^18.3.1",
|
"react-dom": "^18.3.1",
|
||||||
@@ -4502,6 +4503,17 @@
|
|||||||
"url": "https://github.com/sponsors/sindresorhus"
|
"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": {
|
"node_modules/math-intrinsics": {
|
||||||
"version": "1.1.0",
|
"version": "1.1.0",
|
||||||
"resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz",
|
"resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz",
|
||||||
|
|||||||
@@ -25,6 +25,7 @@
|
|||||||
"axios": "^1.7.2",
|
"axios": "^1.7.2",
|
||||||
"classnames": "^2.5.1",
|
"classnames": "^2.5.1",
|
||||||
"dayjs": "^1.11.23",
|
"dayjs": "^1.11.23",
|
||||||
|
"marked": "^12.0.2",
|
||||||
"mp4box": "^2.4.1",
|
"mp4box": "^2.4.1",
|
||||||
"react": "^18.3.1",
|
"react": "^18.3.1",
|
||||||
"react-dom": "^18.3.1",
|
"react-dom": "^18.3.1",
|
||||||
|
|||||||
@@ -65,27 +65,23 @@ export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
|
|||||||
return isImageAnalysisStage(stage) || isCopyStage(stage)
|
return isImageAnalysisStage(stage) || isCopyStage(stage)
|
||||||
}
|
}
|
||||||
|
|
||||||
/** 单张图片 VLM 识别出的商品信息 */
|
/** 单张图片 VLM 识别结果(v8 叙述优先,仅保留最少结构化字段) */
|
||||||
export interface ImageProductAnalysis {
|
export interface ImageProductAnalysis {
|
||||||
|
/** store / product / person / scene */
|
||||||
|
type?: string
|
||||||
name?: string
|
name?: string
|
||||||
category?: string
|
|
||||||
brand?: string
|
brand?: string
|
||||||
colors?: string[]
|
has_person?: boolean
|
||||||
material_or_texture?: string
|
/** v8: 用户端展示用的叙述 markdown(由提示词控制排版) */
|
||||||
key_features?: string[]
|
summary_markdown?: string
|
||||||
visual_style?: string
|
/** 标题行兼容字段 */
|
||||||
scene?: string
|
category?: 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 {
|
export interface ImageAnalysisResult {
|
||||||
|
/** v8 字段 */
|
||||||
|
images?: ImageProductAnalysis[]
|
||||||
|
/** 老数据兼容 */
|
||||||
products?: ImageProductAnalysis[]
|
products?: ImageProductAnalysis[]
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,6 +128,8 @@ export interface CopyResult {
|
|||||||
/** 向后兼容:= voiceover_script */
|
/** 向后兼容:= voiceover_script */
|
||||||
suggested_copy?: string
|
suggested_copy?: string
|
||||||
title?: string
|
title?: string
|
||||||
|
/** v3 storyboard: 用户端展示用的 markdown 文案(由提示词控制排版) */
|
||||||
|
copy_display_markdown?: string
|
||||||
/** v1.5 旧字段兼容(老数据降级时可能出现) */
|
/** v1.5 旧字段兼容(老数据降级时可能出现) */
|
||||||
scenes?: Array<{ shot: string; narration: string; duration?: number }>
|
scenes?: Array<{ shot: string; narration: string; duration?: number }>
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1973,3 +1973,99 @@
|
|||||||
padding-bottom: 6px;
|
padding-bottom: 6px;
|
||||||
border-bottom: 1px dashed #e5e7eb;
|
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,5 +1,6 @@
|
|||||||
import React, { useCallback, useEffect, useRef, useState } from "react"
|
import React, { useCallback, useEffect, useRef, useState } from "react"
|
||||||
import axios from "axios"
|
import axios from "axios"
|
||||||
|
import { marked } from "marked"
|
||||||
import {
|
import {
|
||||||
PlusOutlined,
|
PlusOutlined,
|
||||||
CloseOutlined,
|
CloseOutlined,
|
||||||
@@ -150,6 +151,16 @@ type TabTask = {
|
|||||||
audioInst: HTMLAudioElement | null
|
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, "&").replace(/</g, "<")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/* ─────────── 常量 ─────────── */
|
/* ─────────── 常量 ─────────── */
|
||||||
|
|
||||||
const LANGUAGES = ["中文(普通话)", "粤语", "英语", "日语", "韩语"]
|
const LANGUAGES = ["中文(普通话)", "粤语", "英语", "日语", "韩语"]
|
||||||
@@ -296,6 +307,8 @@ interface Storyboard {
|
|||||||
hard_constraints: string[]
|
hard_constraints: string[]
|
||||||
negative_prompts: string[]
|
negative_prompts: string[]
|
||||||
voiceover_script: string
|
voiceover_script: string
|
||||||
|
/** v3: 用户端展示用 markdown 文案(由提示词控制排版) */
|
||||||
|
copy_display_markdown: string
|
||||||
}
|
}
|
||||||
|
|
||||||
/** 兼容旧 copy_result(final_copy/title/scenes)→ 新 Storyboard 结构 */
|
/** 兼容旧 copy_result(final_copy/title/scenes)→ 新 Storyboard 结构 */
|
||||||
@@ -323,6 +336,7 @@ function copyResultToStoryboard(cr: CopyResult | null | undefined): Storyboard |
|
|||||||
hard_constraints: Array.isArray(cr.hard_constraints) ? cr.hard_constraints : [],
|
hard_constraints: Array.isArray(cr.hard_constraints) ? cr.hard_constraints : [],
|
||||||
negative_prompts: Array.isArray(cr.negative_prompts) ? cr.negative_prompts : [],
|
negative_prompts: Array.isArray(cr.negative_prompts) ? cr.negative_prompts : [],
|
||||||
voiceover_script: cr.voiceover_script || cr.final_copy || cr.suggested_copy || "",
|
voiceover_script: cr.voiceover_script || cr.final_copy || cr.suggested_copy || "",
|
||||||
|
copy_display_markdown: cr.copy_display_markdown || "",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// 兜底:旧结构转简单分镜
|
// 兜底:旧结构转简单分镜
|
||||||
@@ -358,6 +372,7 @@ function copyResultToStoryboard(cr: CopyResult | null | undefined): Storyboard |
|
|||||||
hard_constraints: [],
|
hard_constraints: [],
|
||||||
negative_prompts: [],
|
negative_prompts: [],
|
||||||
voiceover_script: finalCopy,
|
voiceover_script: finalCopy,
|
||||||
|
copy_display_markdown: cr.copy_display_markdown || "",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -412,6 +427,7 @@ const MOCK_STORYBOARD: Storyboard = {
|
|||||||
negative_prompts: ["冷色调", "模糊", "变形", "水印文字", "卡通风格", "空无一人"],
|
negative_prompts: ["冷色调", "模糊", "变形", "水印文字", "卡通风格", "空无一人"],
|
||||||
voiceover_script:
|
voiceover_script:
|
||||||
"还在为餐桌选不到好桌子发愁?这张北美黑胡桃木餐桌,一家人坐下来吃饭刚刚好。全实木、无贴皮,纹理好看又耐刮。点小黄车,给家里添一张好桌子。",
|
"还在为餐桌选不到好桌子发愁?这张北美黑胡桃木餐桌,一家人坐下来吃饭刚刚好。全实木、无贴皮,纹理好看又耐刮。点小黄车,给家里添一张好桌子。",
|
||||||
|
copy_display_markdown: "",
|
||||||
}
|
}
|
||||||
|
|
||||||
const fmtSize = (bytes: number | undefined) => {
|
const fmtSize = (bytes: number | undefined) => {
|
||||||
@@ -1192,8 +1208,10 @@ const ViralVideoPage: React.FC = () => {
|
|||||||
|
|
||||||
/* ── 识别描述汇览渲染 ── */
|
/* ── 识别描述汇览渲染 ── */
|
||||||
const renderRecognition = () => {
|
const renderRecognition = () => {
|
||||||
const products: ImageProductAnalysis[] =
|
const images: ImageProductAnalysis[] =
|
||||||
(task.imageAnalysis?.products as ImageProductAnalysis[] | undefined) || []
|
(task.imageAnalysis?.images as ImageProductAnalysis[] | undefined) ||
|
||||||
|
(task.imageAnalysis?.products as ImageProductAnalysis[] | undefined) ||
|
||||||
|
[]
|
||||||
if (task.uiStep === "step1_analyzing") {
|
if (task.uiStep === "step1_analyzing") {
|
||||||
return (
|
return (
|
||||||
<div className="vv-recog">
|
<div className="vv-recog">
|
||||||
@@ -1201,83 +1219,36 @@ const ViralVideoPage: React.FC = () => {
|
|||||||
<LoadingOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
|
<LoadingOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
|
||||||
识别描述汇览
|
识别描述汇览
|
||||||
</div>
|
</div>
|
||||||
<div className="vv-muted">AI 正在识别商品特征…</div>
|
<div className="vv-muted">AI 正在识别画面…</div>
|
||||||
</div>
|
</div>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
if (products.length === 0) return null
|
if (images.length === 0) return null
|
||||||
const featureText = (f: string[] | string | undefined) => {
|
|
||||||
if (!f) return ""
|
|
||||||
if (Array.isArray(f)) return f.join(";")
|
|
||||||
return f
|
|
||||||
}
|
|
||||||
return (
|
return (
|
||||||
<div className="vv-recog">
|
<div className="vv-recog">
|
||||||
<div className="vv-recog-title">
|
<div className="vv-recog-title">
|
||||||
<CheckCircleFilled style={{ color: "#10b981" }} />
|
<CheckCircleFilled style={{ color: "#10b981" }} />
|
||||||
识别描述汇览
|
识别描述汇览
|
||||||
</div>
|
</div>
|
||||||
{products.map((p, i) => (
|
{images.map((p, i) => {
|
||||||
<div key={i} className="vv-recog-item">
|
const meta = [p.name || "未识别", p.brand, p.category].filter(Boolean)
|
||||||
<div className="vv-recog-line">
|
return (
|
||||||
<span className="vv-recog-k">图片{i + 1}:</span>
|
<div key={i} className="vv-recog-item vv-recog-md">
|
||||||
<span>
|
<div className="vv-recog-line">
|
||||||
{p.name || "未识别"}
|
<span className="vv-recog-k">图片{i + 1}:</span>
|
||||||
{p.spec && <span className="vv-recog-meta">({p.spec})</span>}
|
<span>{meta.join(" · ")}</span>
|
||||||
{p.brand && <span className="vv-recog-meta"> · {p.brand}</span>}
|
</div>
|
||||||
{p.category && <span className="vv-recog-meta"> · {p.category}</span>}
|
{p.summary_markdown ? (
|
||||||
</span>
|
<div
|
||||||
|
className="vv-md-body"
|
||||||
|
dangerouslySetInnerHTML={{ __html: renderMarkdown(p.summary_markdown) }}
|
||||||
|
/>
|
||||||
|
) : (
|
||||||
|
<div className="vv-muted">(暂无叙述描述)</div>
|
||||||
|
)}
|
||||||
</div>
|
</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>
|
</div>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -1416,6 +1387,19 @@ const ViralVideoPage: React.FC = () => {
|
|||||||
return (
|
return (
|
||||||
<div className="vv-copy-box vv-storyboard">
|
<div className="vv-copy-box vv-storyboard">
|
||||||
<div className="vv-sb-doc">
|
<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>
|
<h4 className="vv-sb-h">视频总览</h4>
|
||||||
<p className="vv-sb-inline-row">
|
<p className="vv-sb-inline-row">
|
||||||
|
|||||||
@@ -53,6 +53,9 @@ celery_app.conf.imports = (
|
|||||||
# #1998 GPU MuseTalk 异步推理:wait_for_result→签名 URL→回写 lipsync_jobs
|
# #1998 GPU MuseTalk 异步推理:wait_for_result→签名 URL→回写 lipsync_jobs
|
||||||
# 必须在 Worker 侧注册,否则 apply_async 消息无人消费,job 永远卡在 processing
|
# 必须在 Worker 侧注册,否则 apply_async 消息无人消费,job 永远卡在 processing
|
||||||
"app.tasks.lipsync_gpu",
|
"app.tasks.lipsync_gpu",
|
||||||
|
# #2076 Ditto 蚂蚁数字人异步推理:同步 HTTP 调用 Ditto → MP4 流转存 OSS → 回写 lipsync_jobs
|
||||||
|
# 必须在 Worker 侧注册;失败回退 GPU MuseTalk → MediaKit
|
||||||
|
"app.tasks.lipsync_ditto",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Celery Beat 定时任务调度
|
# Celery Beat 定时任务调度
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,13 +1,13 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates 表(prompt_type='image_analysis'
|
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates(prompt_type='image_analysis'
|
||||||
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到纯硬编码 JSON schema prompt。
|
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到 prompts.py 的
|
||||||
|
image_analysis v8 默认 system/user。
|
||||||
|
|
||||||
规则(简单直接,不做字符串匹配判断):
|
规则:
|
||||||
- DB 有 is_active=true 的 image_analysis 记录(含种子版本和用户修改后的版本):
|
- DB 有 is_active=true 的 image_analysis 记录:system 原样用 DB.system_prompt
|
||||||
* system = DB.system_prompt(DB prompt 自带完整输出格式,不追加硬编码 schema,
|
(自带完整输出格式,不追加任何硬编码 schema),user 用 DB.user_prompt_template
|
||||||
避免 DB 写 XML、调用强制 json_object 造成的格式冲突)
|
渲染(填入 image_url / ocr_text);
|
||||||
* user = DB.user_prompt_template 渲染后使用;渲染后为空则用硬编码默认
|
- DB 无记录/异常:system/user 用 prompts.py 里的 v8 默认模板。
|
||||||
- DB 无记录/连接异常/返回空:system/user 全部用纯硬编码 JSON schema prompt
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -15,192 +15,79 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# ---- 纯硬编码 JSON schema(DB 无有效配置时全量使用) ----
|
|
||||||
|
|
||||||
_FAST_JSON_SCHEMA = (
|
def _default_template() -> dict:
|
||||||
"你是图片结构化识别器。严格按下方 JSON schema 返回一个对象,不要任何解释、"
|
# 延迟导入:避免模块加载时拉起整个 packages 依赖链(也便于旧 Python 收集测试)
|
||||||
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填 null 或空数组。\n"
|
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
|
||||||
"{\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'
|
|
||||||
"}\n\n"
|
|
||||||
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
|
|
||||||
)
|
|
||||||
DEFAULT_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
|
|
||||||
|
|
||||||
_PRO_JSON_SCHEMA = (
|
for item in DEFAULT_TEMPLATES:
|
||||||
"你是图片分析专家。严格按下方 JSON schema 返回一个对象,不要解释、不要markdown、不要代码块、不要XML标签。\n"
|
if item["prompt_type"] == "image_analysis":
|
||||||
"{\n"
|
return item
|
||||||
' "has_person": true/false,\n'
|
raise RuntimeError("image_analysis 默认模板缺失")
|
||||||
' "gender": "男"/"女"/null,\n'
|
|
||||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
|
||||||
' "outfit": "整体穿着描述(含颜色款式)",\n'
|
|
||||||
' "hair": "发型发色",\n'
|
|
||||||
' "pose": "姿势",\n'
|
|
||||||
' "expression": "表情",\n'
|
|
||||||
' "scene": "场景",\n'
|
|
||||||
' "mood": "氛围",\n'
|
|
||||||
' "has_product": true/false,\n'
|
|
||||||
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
|
||||||
' "product_name": "产品名,非产品图填null",\n'
|
|
||||||
' "brand": "品牌,无则null",\n'
|
|
||||||
' "key_features": ["核心特征数组,3-6个短语"]\n'
|
|
||||||
"}\n\n"
|
|
||||||
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
|
|
||||||
)
|
|
||||||
DEFAULT_PRO_USER = "分析这张图片,返回符合schema的JSON。"
|
|
||||||
|
|
||||||
# 保留旧 JSON schema 追加文本作为常量(DB prompt 完全控制输出格式后不再使用,
|
|
||||||
# 保留以便排查历史行为)。
|
|
||||||
_FAST_JSON_APPEND = (
|
|
||||||
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
|
|
||||||
"严格包含以下字段(字段值不确定时填null或空数组):\n"
|
|
||||||
"{\n"
|
|
||||||
' "has_person": true/false,\n'
|
|
||||||
' "gender": "男"/"女"/null,\n'
|
|
||||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
|
||||||
' "upper_wear": "上装款式字符串",\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": "图案",\n'
|
|
||||||
' "colors": ["主色数组"],\n'
|
|
||||||
' "mood": "整体氛围"\n'
|
|
||||||
"}\n"
|
|
||||||
"不要输出任何其他文字、解释、XML标签或markdown。"
|
|
||||||
)
|
|
||||||
|
|
||||||
_PRO_JSON_APPEND = (
|
|
||||||
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
|
|
||||||
"严格包含以下字段(字段值不确定时填null或空数组):\n"
|
|
||||||
"{\n"
|
|
||||||
' "has_person": true/false,\n'
|
|
||||||
' "gender": "男"/"女"/null,\n'
|
|
||||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
|
||||||
' "outfit": "整体穿着描述(含颜色款式)",\n'
|
|
||||||
' "hair": "发型发色",\n'
|
|
||||||
' "pose": "姿势",\n'
|
|
||||||
' "expression": "表情",\n'
|
|
||||||
' "scene": "场景",\n'
|
|
||||||
' "mood": "氛围",\n'
|
|
||||||
' "has_product": true/false,\n'
|
|
||||||
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
|
||||||
' "product_name": "产品名,非产品图填null",\n'
|
|
||||||
' "brand": "品牌,无则null",\n'
|
|
||||||
' "key_features": ["核心特征3-6个短语"]\n'
|
|
||||||
"}\n"
|
|
||||||
"不要输出任何其他文字、解释、XML标签或markdown。"
|
|
||||||
)
|
|
||||||
|
|
||||||
_cache_lock = threading.Lock()
|
_cache_lock = threading.Lock()
|
||||||
_cache: dict[str, tuple[float, Any]] = {}
|
_cache: dict[str, tuple[float, tuple[str, str]]] = {}
|
||||||
_CACHE_TTL = 30.0
|
_CACHE_TTL = 30.0
|
||||||
|
|
||||||
|
|
||||||
def _load_db_template() -> Any | None:
|
def _load_db_template():
|
||||||
"""直接查DB viral_video_prompt_templates 中 is_active=true 的 image_analysis 记录;
|
"""查 DB is_active=true 的 image_analysis 记录;不可达/无记录返回 None。"""
|
||||||
DB不可达/无记录/异常返回None。
|
|
||||||
复用 prompt_loader._load_from_db,它只查DB不做DEFAULT_TEMPLATES fallback,
|
|
||||||
返回None表示DB无记录或异常。"""
|
|
||||||
try:
|
try:
|
||||||
from packages.application.viral_video.prompt_loader import _load_from_db
|
from packages.application.viral_video.prompt_loader import _load_from_db
|
||||||
|
|
||||||
return _load_from_db("image_analysis")
|
return _load_from_db("image_analysis")
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
logger.warning("[vision.v2] 查询DB prompt配置失败: %s", e)
|
logger.warning("[vision.v2] 查询DB image_analysis prompt失败: %s", e)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _render_user(tpl: Any | None, default_user: str) -> str:
|
def _render_user(user_tpl: str, image_url: str, ocr_text: str) -> str:
|
||||||
if not tpl:
|
try:
|
||||||
return default_user
|
return user_tpl.format(image_url=image_url, ocr_text=ocr_text or "无")
|
||||||
tpl_str = getattr(tpl, "user_prompt_template", "") or ""
|
except Exception: # noqa: BLE001
|
||||||
if not tpl_str.strip():
|
return user_tpl
|
||||||
return default_user
|
|
||||||
rendered = tpl_str.replace("{image_count}", "1").replace("{industry}", "通用").replace("{image_urls}", "").strip()
|
|
||||||
return rendered or default_user
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_fast_prompt() -> tuple[str, str]:
|
def _resolve(kind: str, image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
|
||||||
return _resolve("fast")
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_pro_prompt() -> tuple[str, str]:
|
|
||||||
return _resolve("pro")
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve(kind: str) -> tuple[str, str]:
|
|
||||||
now = time.time()
|
now = time.time()
|
||||||
cache_key = f"prompt_{kind}"
|
cache_key = f"prompt_{kind}"
|
||||||
with _cache_lock:
|
with _cache_lock:
|
||||||
hit = _cache.get(cache_key)
|
hit = _cache.get(cache_key)
|
||||||
if hit and now - hit[0] < _CACHE_TTL:
|
if hit and now - hit[0] < _CACHE_TTL:
|
||||||
return hit[1]
|
sys_prompt, usr_prompt = hit[1]
|
||||||
|
return sys_prompt, _render_user(usr_prompt, image_url, ocr_text)
|
||||||
|
|
||||||
default_sys = _FAST_JSON_SCHEMA if kind == "fast" else _PRO_JSON_SCHEMA
|
default = _default_template()
|
||||||
default_user = DEFAULT_FAST_USER if kind == "fast" else DEFAULT_PRO_USER
|
sys_prompt = default["system_prompt"]
|
||||||
|
usr_prompt = default["user_prompt_template"]
|
||||||
|
|
||||||
sys_prompt = default_sys
|
tpl = _load_db_template()
|
||||||
usr_prompt = default_user
|
if tpl is not None:
|
||||||
try:
|
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
|
||||||
tpl = _load_db_template()
|
if db_sys:
|
||||||
if tpl is not None:
|
sys_prompt = db_sys
|
||||||
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
|
db_usr = getattr(tpl, "user_prompt_template", "") or usr_prompt
|
||||||
if db_sys:
|
usr_prompt = db_usr or usr_prompt
|
||||||
sys_prompt = db_sys # DB prompt自带完整输出格式,不追加硬编码schema避免冲突
|
logger.info(
|
||||||
usr_prompt = _render_user(tpl, default_user)
|
"[vision.v2] 使用DB image_analysis prompt version=%s",
|
||||||
logger.info(
|
getattr(tpl, "version", "?"),
|
||||||
"[vision.v2] 使用DB image_analysis prompt (kind=%s version=%s sys_len=%d)",
|
)
|
||||||
kind,
|
|
||||||
getattr(tpl, "version", "?"),
|
|
||||||
len(db_sys),
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.debug("[vision.v2] DB image_analysis system_prompt为空,使用默认JSON (kind=%s)", kind)
|
|
||||||
else:
|
|
||||||
logger.debug("[vision.v2] DB无image_analysis记录/不可达,使用默认JSON prompt (kind=%s)", kind)
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning("[vision.v2] 解析DB prompt异常,使用默认: %s", e)
|
|
||||||
|
|
||||||
with _cache_lock:
|
with _cache_lock:
|
||||||
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
|
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
|
||||||
return 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:
|
def invalidate_cache() -> None:
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,12 +1,11 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""V2 图片分析主路径:每图并行 OCR(火山MediaKit,未配置时自动跳过)+ qwen3.8-flash JSON VLM,
|
"""V2 图片分析主路径(v8 叙述优先):每图并行 OCR(火山 MediaKit,未配置自动跳过)
|
||||||
失败时单次 qwen3.7-plus 兜底。
|
+ fast VLM 强约束 JSON;失败时单次 pro VLM 兜底。
|
||||||
|
|
||||||
架构(灵应10-05确认):
|
架构:
|
||||||
- 唯一后端:阿里云百炼 DashScope,qwen3.8-flash 做快速路径、qwen3.7-plus 做兜底
|
- 单图 2 路并行(OCR + fast VLM),外层 N 图全并发(workers=8);
|
||||||
- 主力:单图2路并行(OCR + fast VLM),外层N图全并发(workers=8)
|
- 兜底单次 pro VLM,无竞速/复杂重试;
|
||||||
- 兜底:单次 pro VLM 调用,无竞速/重试/复杂超时
|
- 输出统一为 5 字段 image dict(type/name/brand/has_person/summary_markdown)。
|
||||||
- 输出 dict 格式与旧版完全一致,下游零改动
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -21,36 +20,24 @@ from . import assembler, ocr_volc, vlm_fallback, vlm_fast_json
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# 超时(可通过环境变量覆盖)
|
|
||||||
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
||||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "20"))
|
_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_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "20"))
|
||||||
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
|
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
|
||||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
|
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
|
||||||
|
|
||||||
_FALLBACK_RESULT = {
|
|
||||||
"name": "未识别",
|
def _is_usable(r: dict[str, Any] | None) -> bool:
|
||||||
"brand": "无法判断",
|
if not isinstance(r, dict):
|
||||||
"category": "非产品图",
|
return False
|
||||||
"appearance": "无法判断",
|
return bool((r.get("summary_markdown") or "").strip())
|
||||||
"packaging": "无法判断",
|
|
||||||
"text_on_package": [],
|
|
||||||
"key_features": ["无法判断"],
|
|
||||||
"scene": "通用",
|
|
||||||
"mood": "",
|
|
||||||
"portrait_prompt": "无法判断",
|
|
||||||
"summary": "未识别",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _is_usable(r: dict[str, Any]) -> bool:
|
def _basic_failure(ocr_result: list[str], fast_elapsed: float, source: str) -> dict[str, Any]:
|
||||||
pp = (r.get("portrait_prompt") or "").strip()
|
image = assembler.assemble_result(-1, {}, ocr_result)
|
||||||
if pp and pp not in ("无人像", "无法判断", "未识别"):
|
image["_source"] = source
|
||||||
return True
|
image["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||||
name = (r.get("name") or "").strip()
|
return image
|
||||||
if name and name not in ("未识别", "无法判断", "未知"):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
||||||
@@ -58,7 +45,6 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
|||||||
|
|
||||||
fj_result: dict[str, Any] | None = None
|
fj_result: dict[str, Any] | None = None
|
||||||
ocr_result: list[str] = []
|
ocr_result: list[str] = []
|
||||||
fast_elapsed = 0.0
|
|
||||||
pool = ThreadPoolExecutor(max_workers=2)
|
pool = ThreadPoolExecutor(max_workers=2)
|
||||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
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)
|
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||||
@@ -66,7 +52,7 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
|||||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||||
try:
|
try:
|
||||||
res = fut.result(timeout=1)
|
res = fut.result(timeout=1)
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||||
continue
|
continue
|
||||||
if fut is f_fj and isinstance(res, dict):
|
if fut is f_fj and isinstance(res, dict):
|
||||||
@@ -80,37 +66,24 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
|||||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||||
finally:
|
finally:
|
||||||
fast_elapsed = time.time() - t0
|
fast_elapsed = time.time() - t0
|
||||||
pool.shutdown(wait=False) # 不等待未完成的线程,避免计时膨胀
|
pool.shutdown(wait=False)
|
||||||
|
|
||||||
if fj_result:
|
if fj_result:
|
||||||
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
||||||
if _is_usable(assembled):
|
if _is_usable(assembled):
|
||||||
assembled["_fast_elapsed"] = round(fast_elapsed, 2)
|
assembled["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||||
logger.info(
|
logger.info("[vision.v2] 图片 #%d fast命中 elapsed=%.2fs", idx, fast_elapsed)
|
||||||
"[vision.v2] 图片 #%d fast命中 elapsed=%.2fs pp=%s",
|
|
||||||
idx,
|
|
||||||
fast_elapsed,
|
|
||||||
(assembled.get("portrait_prompt") or "")[:40],
|
|
||||||
)
|
|
||||||
return assembled
|
return assembled
|
||||||
|
|
||||||
pro_t0 = time.time()
|
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, ocr_hint=ocr_result, timeout=_PRO_TIMEOUT)
|
||||||
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, timeout=_PRO_TIMEOUT)
|
if _is_usable(pro_result):
|
||||||
if pro_result and _is_usable(pro_result):
|
|
||||||
pro_result["_fallback_used"] = True
|
pro_result["_fallback_used"] = True
|
||||||
pro_result["_fast_elapsed"] = round(fast_elapsed, 2)
|
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)
|
logger.info("[vision.v2] 图片 #%d pro兜底命中 total=%.2fs", idx, time.time() - t0)
|
||||||
return pro_result
|
return pro_result
|
||||||
|
|
||||||
logger.warning("[vision.v2] 图片 #%d 全路径失败 elapsed=%.2fs", idx, time.time() - t0)
|
logger.warning("[vision.v2] 图片 #%d 全路径失败 elapsed=%.2fs", idx, time.time() - t0)
|
||||||
out = dict(_FALLBACK_RESULT)
|
return _basic_failure(ocr_result, fast_elapsed, "v2_all_failed")
|
||||||
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]]:
|
def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
|
||||||
@@ -133,14 +106,19 @@ def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
|
|||||||
idx = future_to_idx[fut]
|
idx = future_to_idx[fut]
|
||||||
try:
|
try:
|
||||||
results[idx] = fut.result()
|
results[idx] = fut.result()
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
|
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
|
||||||
r = dict(_FALLBACK_RESULT)
|
results[idx] = assembler.assemble_result(idx, {}, [])
|
||||||
r["_source"] = "v2_future_exception"
|
results[idx]["_source"] = "v2_future_exception" # type: ignore[index]
|
||||||
results[idx] = r
|
|
||||||
|
|
||||||
elapsed = time.time() - t0
|
elapsed = time.time() - t0
|
||||||
succ = sum(1 for r in results if r and _is_usable(r))
|
succ = sum(1 for r in results if _is_usable(r))
|
||||||
fb = sum(1 for r in results if r and r.get("_fallback_used"))
|
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)
|
logger.info(
|
||||||
return [r for r in results if r is not None]
|
"[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]
|
||||||
|
|||||||
@@ -1,14 +1,8 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""V2 兜底路径:image_analysis(默认 qwen-vl-plus 视觉模型,fallback qwen3.7-plus / DashScope)单图调用。
|
"""V2 兜底路径:vision client(fallback 变体)单图调用,走 v8 叙述优先 prompt。
|
||||||
|
|
||||||
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
|
fast 超时/非 JSON/为空时单次调用;输出统一走 assembler.assemble_result 组装,
|
||||||
设计要点:
|
与 fast 路径同为 5 字段 image dict。
|
||||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
|
||||||
- enable_thinking=False + response_format=json_object
|
|
||||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
|
||||||
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
|
|
||||||
- timeout=30s
|
|
||||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -28,11 +22,12 @@ def call_pro_vlm(
|
|||||||
img_url: str,
|
img_url: str,
|
||||||
idx: int,
|
idx: int,
|
||||||
*,
|
*,
|
||||||
|
ocr_hint: list[str] | None = None,
|
||||||
timeout: int = _DEFAULT_TIMEOUT,
|
timeout: int = _DEFAULT_TIMEOUT,
|
||||||
max_tokens: int | None = None,
|
max_tokens: int | None = None,
|
||||||
) -> dict[str, Any] | None:
|
) -> dict[str, Any] | None:
|
||||||
"""max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置。"""
|
|
||||||
t0 = time.time()
|
t0 = time.time()
|
||||||
|
ocr_text = "、".join(t for t in (ocr_hint or []) if t)[:200]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from packages.shared.ai_router import ai_router
|
from packages.shared.ai_router import ai_router
|
||||||
@@ -41,13 +36,13 @@ def call_pro_vlm(
|
|||||||
if not client or not client.is_available:
|
if not client or not client.is_available:
|
||||||
logger.warning("[vision.v2] pro vision client 不可用,跳过")
|
logger.warning("[vision.v2] pro vision client 不可用,跳过")
|
||||||
return None
|
return None
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
|
system_prompt, user_prompt = _prompt.resolve_pro_prompt(img_url, ocr_text)
|
||||||
|
|
||||||
messages = [
|
messages: list[dict[str, Any]] = [
|
||||||
{"role": "system", "content": system_prompt},
|
{"role": "system", "content": system_prompt},
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
@@ -61,31 +56,31 @@ def call_pro_vlm(
|
|||||||
try:
|
try:
|
||||||
call_kwargs: dict[str, Any] = {
|
call_kwargs: dict[str, Any] = {
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"images": None, # 图片已在 messages 中
|
"images": None,
|
||||||
"temperature": 0.3,
|
"temperature": 0.3,
|
||||||
"timeout": timeout,
|
"timeout": timeout,
|
||||||
"enable_thinking": False,
|
"enable_thinking": False,
|
||||||
"response_format": {"type": "json_object"},
|
"response_format": {"type": "json_object"},
|
||||||
|
"max_tokens": max_tokens if max_tokens is not None else 4000,
|
||||||
}
|
}
|
||||||
# pro fallback:显式4000 tokens给复杂门店图留足空间
|
|
||||||
call_kwargs["max_tokens"] = max_tokens if max_tokens is not None else 4000
|
|
||||||
|
|
||||||
from .json_utils import extract_json_object
|
from .json_utils import extract_json_object
|
||||||
|
|
||||||
raw = None
|
|
||||||
obj = None
|
obj = None
|
||||||
for _outer in range(2):
|
for _outer in range(2):
|
||||||
kw = dict(call_kwargs)
|
kw = dict(call_kwargs)
|
||||||
if _outer == 1:
|
if _outer == 1:
|
||||||
kw.pop("response_format", None)
|
kw.pop("response_format", None)
|
||||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||||
cont = [dict(c) for c in list(msgs2[1]["content"])]
|
cont = [dict(c) for c in msgs2[1]["content"]]
|
||||||
cont[-1] = {"type": "text", "text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。"}
|
cont[-1] = {
|
||||||
|
"type": "text",
|
||||||
|
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
|
||||||
|
}
|
||||||
msgs2[1] = {"role": "user", "content": cont}
|
msgs2[1] = {"role": "user", "content": cont}
|
||||||
kw["messages"] = msgs2
|
kw["messages"] = msgs2
|
||||||
raw = client.vision_completion(**kw)
|
raw = client.vision_completion(**kw)
|
||||||
if not raw:
|
if not raw:
|
||||||
logger.warning("[vision.v2] pro 返回空 outer=%s", _outer)
|
|
||||||
continue
|
continue
|
||||||
obj = extract_json_object(raw)
|
obj = extract_json_object(raw)
|
||||||
if obj is not None:
|
if obj is not None:
|
||||||
@@ -97,20 +92,18 @@ def call_pro_vlm(
|
|||||||
logger.warning("[vision.v2] pro 两次均未得到JSON elapsed=%.1fs", elapsed)
|
logger.warning("[vision.v2] pro 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||||
return None
|
return None
|
||||||
if obj.get("_partial"):
|
if obj.get("_partial"):
|
||||||
logger.warning("[vision.v2] pro 返回截断JSON(partial) elapsed=%.1fs", elapsed)
|
logger.warning("[vision.v2] pro 截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||||
logger.info(
|
|
||||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs type=%s",
|
|
||||||
client.model,
|
|
||||||
elapsed,
|
|
||||||
obj.get("type"),
|
|
||||||
)
|
|
||||||
|
|
||||||
# 通过assembler统一组装,兼容v4嵌套schema和旧扁平schema
|
result = assembler.assemble_result(idx, obj, ocr_hint or [])
|
||||||
result = assembler.assemble_result(idx, obj, [])
|
|
||||||
result["_source"] = "vlm_pro"
|
result["_source"] = "vlm_pro"
|
||||||
result["_fallback_used"] = True
|
result["_fallback_used"] = True
|
||||||
|
logger.info("[vision.v2] pro 完成 model=%s elapsed=%.1fs", client.model, elapsed)
|
||||||
return result
|
return result
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
elapsed = time.time() - t0
|
logger.warning(
|
||||||
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
"[vision.v2] pro 异常 elapsed=%.1fs err=%s",
|
||||||
|
time.time() - t0,
|
||||||
|
e,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -1,15 +1,11 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""V2 快速路径:image_analysis capability(默认 qwen-vl-plus 视觉模型 / DashScope)强约束 JSON-only 调用。
|
"""V2 快速路径:vision client(默认 image_analysis 能力)强约束 JSON-only 调用。
|
||||||
|
|
||||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
要点:
|
||||||
设计要点:
|
- 通过 ai_router.get_vision_client() 获取 client;
|
||||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
- enable_thinking=False 关闭推理链,response_format=json_object 强约束 JSON;
|
||||||
- enable_thinking=False 关闭推理链(reasoning 是延迟主因)
|
- system/user prompt 优先读后台模板(v8 叙述优先),DB 不可用时用 prompts.py 默认;
|
||||||
- response_format=json_object 强约束JSON输出
|
- temperature=0.1(稳定输出 JSON);两次尝试(第二次去 json_object 约束)。
|
||||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
|
||||||
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
|
|
||||||
- temperature=0.1(稳定输出 JSON)
|
|
||||||
- timeout=15s(失败由外层走 pro 兜底)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -31,11 +27,6 @@ def call_fast_json(
|
|||||||
timeout: int = _DEFAULT_TIMEOUT,
|
timeout: int = _DEFAULT_TIMEOUT,
|
||||||
max_tokens: int | None = None,
|
max_tokens: int | None = None,
|
||||||
) -> dict[str, Any] | None:
|
) -> dict[str, Any] | None:
|
||||||
"""调用 vision client 返回结构化 dict;失败/非 JSON 返回 None。
|
|
||||||
|
|
||||||
max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置;
|
|
||||||
显式传入时作为覆盖。
|
|
||||||
"""
|
|
||||||
t0 = time.time()
|
t0 = time.time()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -45,13 +36,13 @@ def call_fast_json(
|
|||||||
if not client or not client.is_available:
|
if not client or not client.is_available:
|
||||||
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
|
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
|
||||||
return None
|
return None
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
system_prompt, user_prompt = _prompt.resolve_fast_prompt(img_url, "")
|
||||||
|
|
||||||
messages = [
|
messages: list[dict[str, Any]] = [
|
||||||
{"role": "system", "content": system_prompt},
|
{"role": "system", "content": system_prompt},
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
@@ -65,7 +56,7 @@ def call_fast_json(
|
|||||||
try:
|
try:
|
||||||
call_kwargs: dict[str, Any] = {
|
call_kwargs: dict[str, Any] = {
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"images": None, # 图片已在 messages 中
|
"images": None,
|
||||||
"temperature": 0.1,
|
"temperature": 0.1,
|
||||||
"timeout": timeout,
|
"timeout": timeout,
|
||||||
"enable_thinking": False,
|
"enable_thinking": False,
|
||||||
@@ -74,50 +65,47 @@ def call_fast_json(
|
|||||||
if max_tokens is not None:
|
if max_tokens is not None:
|
||||||
call_kwargs["max_tokens"] = max_tokens
|
call_kwargs["max_tokens"] = max_tokens
|
||||||
|
|
||||||
# 双重防护:第1次正常调用;第2次去掉json_object强约束(部分模型在该约束下
|
|
||||||
# 反而幻觉),并加严格指令。解析全部走 json_utils,截断partial产物可用。
|
|
||||||
from .json_utils import extract_json_object
|
from .json_utils import extract_json_object
|
||||||
|
|
||||||
raw = None
|
|
||||||
obj = None
|
obj = None
|
||||||
for _outer in range(2):
|
for _outer in range(2):
|
||||||
kw = dict(call_kwargs)
|
kw = dict(call_kwargs)
|
||||||
if _outer == 1:
|
if _outer == 1:
|
||||||
kw.pop("response_format", None)
|
kw.pop("response_format", None)
|
||||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||||
cont = list(msgs2[1]["content"])
|
cont = [dict(c) for c in msgs2[1]["content"]]
|
||||||
cont = [dict(c) for c in cont]
|
cont[-1] = {
|
||||||
cont[-1] = {"type": "text", "text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。"}
|
"type": "text",
|
||||||
|
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
|
||||||
|
}
|
||||||
msgs2[1] = {"role": "user", "content": cont}
|
msgs2[1] = {"role": "user", "content": cont}
|
||||||
kw["messages"] = msgs2
|
kw["messages"] = msgs2
|
||||||
raw = client.vision_completion(**kw)
|
raw = client.vision_completion(**kw)
|
||||||
if not raw:
|
if not raw:
|
||||||
logger.warning("[vision.v2] fast_json 返回空 outer=%s", _outer)
|
|
||||||
continue
|
continue
|
||||||
obj = extract_json_object(raw)
|
obj = extract_json_object(raw)
|
||||||
if obj is not None:
|
if obj is not None:
|
||||||
break
|
break
|
||||||
logger.warning(
|
logger.warning("[vision.v2] fast_json 非JSON(100字) outer=%s: %s", _outer, raw[:100])
|
||||||
"[vision.v2] fast_json 非JSON(100字) outer=%s: %s",
|
|
||||||
_outer,
|
|
||||||
raw[:100],
|
|
||||||
)
|
|
||||||
|
|
||||||
elapsed = time.time() - t0
|
elapsed = time.time() - t0
|
||||||
if obj is None:
|
if obj is None:
|
||||||
logger.warning("[vision.v2] fast_json 两次均未得到JSON elapsed=%.1fs", elapsed)
|
logger.warning("[vision.v2] fast_json 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||||
return None
|
return None
|
||||||
if obj.get("_partial"):
|
if obj.get("_partial"):
|
||||||
logger.warning("[vision.v2] fast_json 返回截断JSON(partial) elapsed=%.1fs", elapsed)
|
logger.warning("[vision.v2] fast_json 截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||||
logger.info(
|
logger.info(
|
||||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs has_person=%s type=%s",
|
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs type=%s",
|
||||||
client.model,
|
client.model,
|
||||||
elapsed,
|
elapsed,
|
||||||
obj.get("has_person"),
|
|
||||||
obj.get("type"),
|
obj.get("type"),
|
||||||
)
|
)
|
||||||
return obj
|
return obj
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001
|
||||||
elapsed = time.time() - t0
|
logger.warning(
|
||||||
logger.warning("[vision.v2] fast_json 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
"[vision.v2] fast_json 异常 elapsed=%.1fs err=%s",
|
||||||
|
time.time() - t0,
|
||||||
|
e,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -311,3 +311,10 @@ GPU_ENCODE_CRF=23
|
|||||||
GPU_ENCODE_FALLBACK_CPU=true
|
GPU_ENCODE_FALLBACK_CPU=true
|
||||||
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
||||||
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
||||||
|
|
||||||
|
# ==================== Ditto 蚂蚁数字人口型 ====================
|
||||||
|
# 注意:这些值必须写死在模板里(不是 CI Secret),否则每次 CI 重新渲染 .env 都会被丢弃,
|
||||||
|
# 导致 staging 发版后 Ditto 口型服务静默降级到 GPU/MediaKit(P0 防复发)。
|
||||||
|
USE_DITTO_LIPSYNC=true
|
||||||
|
DITTO_API_BASE_URL=http://100.76.80.23:8000
|
||||||
|
DITTO_DEFAULT_VIDEO_URL=https://xiaoxia-autocut.oss-cn-hangzhou.aliyuncs.com/uploads/default_avatar.mp4
|
||||||
|
|||||||
@@ -3,10 +3,24 @@
|
|||||||
使用 bcrypt 安全存储密码
|
使用 bcrypt 安全存储密码
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import hashlib
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import bcrypt
|
import bcrypt
|
||||||
|
|
||||||
|
# bcrypt 只对前 72 字节有效,且 bcrypt>=4.1 会对超长输入直接抛 ValueError。
|
||||||
|
# 超长密码先做一次 SHA-256(定长 hex),再交给 bcrypt,
|
||||||
|
# 既绕过长度限制又保持对超长不同密码的区分度。
|
||||||
|
_BCRYPT_MAX_BYTES = 72
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_password_bytes(password: str) -> bytes:
|
||||||
|
raw = password.encode("utf-8")
|
||||||
|
if len(raw) > _BCRYPT_MAX_BYTES:
|
||||||
|
return hashlib.sha256(raw).hexdigest().encode("utf-8")
|
||||||
|
return raw
|
||||||
|
|
||||||
|
|
||||||
from packages.domain.auth.password_hasher import PasswordHasherPort, PasswordValidatorPort
|
from packages.domain.auth.password_hasher import PasswordHasherPort, PasswordValidatorPort
|
||||||
|
|
||||||
|
|
||||||
@@ -42,8 +56,8 @@ class PasswordHasher(PasswordHasherPort):
|
|||||||
if not password:
|
if not password:
|
||||||
raise ValueError("Password cannot be empty")
|
raise ValueError("Password cannot be empty")
|
||||||
|
|
||||||
# bcrypt 需要 bytes
|
# bcrypt 需要 bytes(超长密码先 SHA-256 以兼容 72 字节限制)
|
||||||
password_bytes = password.encode("utf-8")
|
password_bytes = _prepare_password_bytes(password)
|
||||||
|
|
||||||
# 生成 salt 并哈希
|
# 生成 salt 并哈希
|
||||||
salt = bcrypt.gensalt(rounds=self.rounds)
|
salt = bcrypt.gensalt(rounds=self.rounds)
|
||||||
@@ -67,7 +81,7 @@ class PasswordHasher(PasswordHasherPort):
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
try:
|
try:
|
||||||
password_bytes = password.encode("utf-8")
|
password_bytes = _prepare_password_bytes(password)
|
||||||
hashed_bytes = hashed_password.encode("utf-8")
|
hashed_bytes = hashed_password.encode("utf-8")
|
||||||
|
|
||||||
return bcrypt.checkpw(password_bytes, hashed_bytes)
|
return bcrypt.checkpw(password_bytes, hashed_bytes)
|
||||||
|
|||||||
@@ -0,0 +1,316 @@
|
|||||||
|
"""Ditto LLM 情绪分析服务 — #2076 后续:根据文案生成 emo_timeline.
|
||||||
|
|
||||||
|
职责:
|
||||||
|
1. 正则按 。!?; 初步分句
|
||||||
|
2. 调 DoubaoClient.chat_completion 分析每句表情(emo: 0-7, intensity: 0-1)
|
||||||
|
3. 结果 LRU 缓存(文案 hash → 情绪列表)
|
||||||
|
4. LLM 失败/超时/格式错 → 返回空列表(降级中性表情,不阻塞生成)
|
||||||
|
5. TTS 完成后按字数比例或 sentence_timings 对齐成秒级 timeline
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from functools import lru_cache
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ── 表情常量 ─────────────────────────────────────────────────────
|
||||||
|
EMO_ANGER = 0
|
||||||
|
EMO_DISGUST = 1
|
||||||
|
EMO_FEAR = 2
|
||||||
|
EMO_HAPPY = 3
|
||||||
|
EMO_NEUTRAL = 4
|
||||||
|
EMO_SAD = 5
|
||||||
|
EMO_SURPRISE = 6
|
||||||
|
EMO_CONTEMPT = 7
|
||||||
|
ALLOWED_EMOS = {EMO_HAPPY, EMO_NEUTRAL, EMO_SAD, EMO_SURPRISE} # 营销场景白名单
|
||||||
|
|
||||||
|
# ── 分句正则 ─────────────────────────────────────────────────────
|
||||||
|
_SENT_SPLIT_RE = re.compile(r"(?<=[。!?;!?;])\s*")
|
||||||
|
|
||||||
|
# ── 默认 prompt 模板文件路径 ──────────────────────────────────────
|
||||||
|
_DEFAULT_PROMPT_PATH = Path(__file__).parent / "prompts" / "ditto_emotion.txt"
|
||||||
|
|
||||||
|
|
||||||
|
def _load_default_prompt() -> str:
|
||||||
|
try:
|
||||||
|
return _DEFAULT_PROMPT_PATH.read_text(encoding="utf-8").strip()
|
||||||
|
except Exception:
|
||||||
|
# 文件不存在时用极简兜底
|
||||||
|
return (
|
||||||
|
"分析文案每句话表情,输出JSON数组:"
|
||||||
|
'[{"text":"句子","emo":4,"intensity":0.2}],emo:3开心4中性5伤心6惊讶,'
|
||||||
|
"禁止0/1/2/7。\n【文案】\n{文案}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── 数据结构 ─────────────────────────────────────────────────────
|
||||||
|
class EmotionSegment:
|
||||||
|
"""单句情绪结果(LLM 输出的原始结构)."""
|
||||||
|
|
||||||
|
__slots__ = ("text", "emo", "intensity")
|
||||||
|
|
||||||
|
def __init__(self, text: str, emo: int, intensity: float):
|
||||||
|
self.text = text
|
||||||
|
self.emo = emo
|
||||||
|
self.intensity = intensity
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, Any]:
|
||||||
|
return {"text": self.text, "emo": self.emo, "intensity": self.intensity}
|
||||||
|
|
||||||
|
|
||||||
|
class EmotionTimelineEntry:
|
||||||
|
"""对齐到音频时间轴后的情绪片段(传给 Ditto)."""
|
||||||
|
|
||||||
|
__slots__ = ("start", "end", "emo", "intensity")
|
||||||
|
|
||||||
|
def __init__(self, start: float, end: float, emo: int, intensity: float):
|
||||||
|
self.start = round(start, 2)
|
||||||
|
self.end = round(end, 2)
|
||||||
|
self.emo = emo
|
||||||
|
self.intensity = round(intensity, 2)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"start": self.start,
|
||||||
|
"end": self.end,
|
||||||
|
"emo": self.emo,
|
||||||
|
"intensity": self.intensity,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ── 分句 ─────────────────────────────────────────────────────────
|
||||||
|
def split_sentences(text: str) -> list[str]:
|
||||||
|
"""按中文句末标点切分,过滤空串."""
|
||||||
|
if not text:
|
||||||
|
return []
|
||||||
|
parts = _SENT_SPLIT_RE.split(text.strip())
|
||||||
|
return [p.strip() for p in parts if p and p.strip()]
|
||||||
|
|
||||||
|
|
||||||
|
# ── 解析 LLM 返回的 JSON ─────────────────────────────────────────
|
||||||
|
def _parse_emotion_json(raw: str) -> list[EmotionSegment]:
|
||||||
|
"""解析 LLM 返回,容错处理:
|
||||||
|
- 去掉 markdown 代码块包裹
|
||||||
|
- 只取第一个 JSON 数组
|
||||||
|
- 逐行校验 emo/intensity 合法性,过滤无效项
|
||||||
|
"""
|
||||||
|
if not raw:
|
||||||
|
return []
|
||||||
|
text = raw.strip()
|
||||||
|
# 去掉 ```json ... ``` 包裹
|
||||||
|
if text.startswith("```"):
|
||||||
|
text = re.sub(r"^```(?:json)?\s*", "", text)
|
||||||
|
text = re.sub(r"\s*```$", "", text)
|
||||||
|
# 找第一个 [ 到最后一个 ]
|
||||||
|
lb = text.find("[")
|
||||||
|
rb = text.rfind("]")
|
||||||
|
if lb == -1 or rb == -1 or rb <= lb:
|
||||||
|
return []
|
||||||
|
try:
|
||||||
|
data = json.loads(text[lb : rb + 1])
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
return []
|
||||||
|
if not isinstance(data, list):
|
||||||
|
return []
|
||||||
|
|
||||||
|
results: list[EmotionSegment] = []
|
||||||
|
for item in data:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
emo = int(item.get("emo", EMO_NEUTRAL))
|
||||||
|
intensity = float(item.get("intensity", 0.2))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
continue
|
||||||
|
if emo not in ALLOWED_EMOS:
|
||||||
|
emo = EMO_NEUTRAL
|
||||||
|
intensity = max(0.05, min(1.0, intensity))
|
||||||
|
sent_text = str(item.get("text", "")).strip()
|
||||||
|
if not sent_text:
|
||||||
|
continue
|
||||||
|
results.append(EmotionSegment(text=sent_text, emo=emo, intensity=intensity))
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
# ── 时间对齐(按字数比例)────────────────────────────────────────
|
||||||
|
def align_timeline_by_length(
|
||||||
|
segments: list[EmotionSegment],
|
||||||
|
audio_duration: float,
|
||||||
|
) -> list[EmotionTimelineEntry]:
|
||||||
|
"""按各句字数占总字数比例分配 audio_duration 时长."""
|
||||||
|
if not segments or audio_duration <= 0:
|
||||||
|
return []
|
||||||
|
total_chars = sum(len(s.text) for s in segments)
|
||||||
|
if total_chars <= 0:
|
||||||
|
return []
|
||||||
|
entries: list[EmotionTimelineEntry] = []
|
||||||
|
pos = 0.0
|
||||||
|
for i, seg in enumerate(segments):
|
||||||
|
if i == len(segments) - 1:
|
||||||
|
end = audio_duration # 最后一段到结尾,避免浮点误差
|
||||||
|
else:
|
||||||
|
end = pos + (len(seg.text) / total_chars) * audio_duration
|
||||||
|
if end > pos:
|
||||||
|
entries.append(
|
||||||
|
EmotionTimelineEntry(
|
||||||
|
start=pos,
|
||||||
|
end=end,
|
||||||
|
emo=seg.emo,
|
||||||
|
intensity=seg.intensity,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
pos = end
|
||||||
|
return entries
|
||||||
|
|
||||||
|
|
||||||
|
def align_timeline_by_timings(
|
||||||
|
segments: list[EmotionSegment],
|
||||||
|
sentence_timings: list[dict[str, Any]],
|
||||||
|
audio_duration: float,
|
||||||
|
) -> list[EmotionTimelineEntry]:
|
||||||
|
"""使用 TTS sentence_timings 精确对齐(优先方案).
|
||||||
|
|
||||||
|
sentence_timings 格式:[{"start":0.0,"end":1.2,"text":"句子"}, ...]
|
||||||
|
按句序匹配 segments 和 timings,长度不一致时回退到按字数比例。
|
||||||
|
"""
|
||||||
|
if not sentence_timings or len(sentence_timings) != len(segments):
|
||||||
|
return align_timeline_by_length(segments, audio_duration)
|
||||||
|
entries: list[EmotionTimelineEntry] = []
|
||||||
|
for seg, timing in zip(segments, sentence_timings, strict=False):
|
||||||
|
try:
|
||||||
|
start = float(timing.get("start", 0))
|
||||||
|
end = float(timing.get("end", 0))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return align_timeline_by_length(segments, audio_duration)
|
||||||
|
if end <= start:
|
||||||
|
continue
|
||||||
|
entries.append(
|
||||||
|
EmotionTimelineEntry(
|
||||||
|
start=start,
|
||||||
|
end=end,
|
||||||
|
emo=seg.emo,
|
||||||
|
intensity=seg.intensity,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return entries
|
||||||
|
|
||||||
|
|
||||||
|
# ── LLM 情绪分析服务 ─────────────────────────────────────────────
|
||||||
|
class DittoEmotionService:
|
||||||
|
"""Ditto 情绪分析服务(带 LRU 缓存)."""
|
||||||
|
|
||||||
|
def __init__(self, settings=None):
|
||||||
|
from packages.config import get_api_settings
|
||||||
|
|
||||||
|
self.settings = settings or get_api_settings()
|
||||||
|
self._client = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
return bool(getattr(self.settings, "ditto_emotion_enabled", False))
|
||||||
|
|
||||||
|
def _get_prompt_template(self) -> str:
|
||||||
|
"""优先用配置(环境变量),否则读文件."""
|
||||||
|
cfg_prompt = getattr(self.settings, "ditto_emotion_prompt", "") or ""
|
||||||
|
if cfg_prompt.strip():
|
||||||
|
return cfg_prompt.strip()
|
||||||
|
return _load_default_prompt()
|
||||||
|
|
||||||
|
def _cache_key(self, text: str) -> str:
|
||||||
|
return hashlib.md5(text.strip().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def _get_llm_client(self):
|
||||||
|
if self._client is None:
|
||||||
|
from packages.shared.ai_client import get_doubao_client
|
||||||
|
|
||||||
|
self._client = get_doubao_client()
|
||||||
|
return self._client
|
||||||
|
|
||||||
|
def _call_llm(self, text: str) -> list[EmotionSegment]:
|
||||||
|
"""调 LLM 分析情绪,失败返回空列表."""
|
||||||
|
template = self._get_prompt_template()
|
||||||
|
prompt = template.replace("{文案}", text)
|
||||||
|
messages = [{"role": "user", "content": prompt}]
|
||||||
|
model = getattr(self.settings, "ditto_emotion_model", "") or None
|
||||||
|
temperature = getattr(self.settings, "ditto_emotion_temperature", 0.1)
|
||||||
|
timeout = getattr(self.settings, "ditto_emotion_timeout", 10)
|
||||||
|
max_tokens = getattr(self.settings, "ditto_emotion_max_tokens", 1024)
|
||||||
|
try:
|
||||||
|
client = self._get_llm_client()
|
||||||
|
result = client.chat_completion(
|
||||||
|
messages=messages,
|
||||||
|
model=model,
|
||||||
|
temperature=temperature,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("[ditto_emotion] LLM 调用异常: %s", exc)
|
||||||
|
return []
|
||||||
|
if not result:
|
||||||
|
return []
|
||||||
|
segments = _parse_emotion_json(result)
|
||||||
|
if not segments:
|
||||||
|
logger.warning("[ditto_emotion] LLM 返回解析失败: %s", result[:200])
|
||||||
|
return segments
|
||||||
|
|
||||||
|
def analyze(self, text: str) -> list[EmotionSegment]:
|
||||||
|
"""分析文案情绪(带缓存),失败返回空列表."""
|
||||||
|
if not self.enabled or not text or not text.strip():
|
||||||
|
return []
|
||||||
|
key = self._cache_key(text)
|
||||||
|
return _cached_analyze(self, key, text)
|
||||||
|
|
||||||
|
def build_timeline(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
audio_duration: float,
|
||||||
|
sentence_timings: Optional[list[dict[str, Any]]] = None,
|
||||||
|
) -> str:
|
||||||
|
"""完整流程:分句→LLM分析→时间对齐→序列化为JSON字符串.
|
||||||
|
|
||||||
|
返回: JSON 字符串(可直接传 Ditto emo_timeline 参数);空字符串表示降级中性。
|
||||||
|
"""
|
||||||
|
segments = self.analyze(text)
|
||||||
|
if not segments:
|
||||||
|
return ""
|
||||||
|
if sentence_timings:
|
||||||
|
entries = align_timeline_by_timings(segments, sentence_timings, audio_duration)
|
||||||
|
else:
|
||||||
|
entries = align_timeline_by_length(segments, audio_duration)
|
||||||
|
if not entries:
|
||||||
|
return ""
|
||||||
|
return json.dumps([e.to_dict() for e in entries], ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
# ── 模块级 LRU 缓存实例 ─────────────────────────────────────────
|
||||||
|
# 每个 service 实例共享缓存(按 cache_key 区分)
|
||||||
|
@lru_cache(maxsize=512)
|
||||||
|
def _cached_analyze(service: DittoEmotionService, cache_key: str, text: str) -> list[EmotionSegment]:
|
||||||
|
"""LRU 缓存包装:cache_key 由文案 hash 生成,maxsize 从配置读."""
|
||||||
|
# 注意:service 参数仅用于传递调用,缓存由 cache_key 驱动
|
||||||
|
segments = service._call_llm(text)
|
||||||
|
# 如果 LLM 返回空(比如分句数量不匹配),尝试直接对预分句结果分析
|
||||||
|
if not segments:
|
||||||
|
pre_splits = split_sentences(text)
|
||||||
|
if len(pre_splits) > 1:
|
||||||
|
# 用预分句结果兜底:全中性低强度
|
||||||
|
segments = [EmotionSegment(text=s, emo=EMO_NEUTRAL, intensity=0.1) for s in pre_splits]
|
||||||
|
return segments
|
||||||
|
|
||||||
|
|
||||||
|
_singleton: Optional[DittoEmotionService] = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_ditto_emotion_service() -> DittoEmotionService:
|
||||||
|
global _singleton
|
||||||
|
if _singleton is None:
|
||||||
|
_singleton = DittoEmotionService()
|
||||||
|
return _singleton
|
||||||
@@ -0,0 +1,291 @@
|
|||||||
|
"""蚂蚁 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)
|
||||||
|
self.blend_frames = int(s.ditto_blend_frames)
|
||||||
|
|
||||||
|
@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: Optional[int] = None,
|
||||||
|
emo_timeline: str = "",
|
||||||
|
) -> 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 = " "
|
||||||
|
|
||||||
|
_blend = blend_frames if blend_frames is not None else self.blend_frames
|
||||||
|
payload = {
|
||||||
|
"video_url": driver_url,
|
||||||
|
"audio_url": audio_url,
|
||||||
|
"script": script,
|
||||||
|
"emo_global": emo_global,
|
||||||
|
"use_script_emo": use_script_emo,
|
||||||
|
"blend_frames": _blend,
|
||||||
|
}
|
||||||
|
if emo_timeline:
|
||||||
|
payload["emo_timeline"] = emo_timeline
|
||||||
|
url = f"{self.base_url}/generate"
|
||||||
|
|
||||||
|
last_exc: Optional[Exception] = None
|
||||||
|
for attempt in range(self.max_retries + 1):
|
||||||
|
try:
|
||||||
|
start = time.monotonic()
|
||||||
|
# 精细化超时:connect=10s(网络不通快速失败),read=120s(最长音频~45s按RTF=2.8推算)
|
||||||
|
_timeout = httpx.Timeout(connect=10.0, read=self.timeout, write=30.0, pool=10.0)
|
||||||
|
with httpx.Client(timeout=_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.ConnectError, httpx.NetworkError, ConnectionError, OSError) as exc:
|
||||||
|
# 网络不通/连接被拒(如 GPU 断网/Tailscale 掉线),不重试,直接快速回退
|
||||||
|
logger.warning("[ditto] 网络不可达 attempt=%d err=%s", attempt + 1, exc)
|
||||||
|
raise DittoError(
|
||||||
|
f"Ditto 网络不可达: {exc}",
|
||||||
|
code="NetworkUnreachable",
|
||||||
|
) from exc
|
||||||
|
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 请求超时(read={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,
|
||||||
|
emo_timeline: str = "",
|
||||||
|
blend_frames: Optional[int] = None,
|
||||||
|
) -> DittoResult:
|
||||||
|
"""调用 generate 并把 MP4 转存到自家 OSS,返回带 video_url 的结果."""
|
||||||
|
result = self.generate(
|
||||||
|
audio_url=audio_url,
|
||||||
|
script=script,
|
||||||
|
video_url=video_url,
|
||||||
|
emo_timeline=emo_timeline,
|
||||||
|
blend_frames=blend_frames,
|
||||||
|
)
|
||||||
|
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
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
你是一个数字人视频表情导演。给定一段口播文案,分析每句话应该用什么表情和强度,让数字人说话时表情自然有变化,不僵硬。
|
||||||
|
【表情编号】
|
||||||
|
0=愤怒(营销场景禁用)
|
||||||
|
1=厌恶(禁用)
|
||||||
|
2=害怕(禁用)
|
||||||
|
3=开心:介绍优点、优惠、好消息、号召行动时用
|
||||||
|
4=中性:默认表情,陈述事实、平铺直叙时用
|
||||||
|
5=伤心:仅在共情用户痛点时低强度使用(如"是不是经常遇到…")
|
||||||
|
6=惊讶:惊喜、意外、强调价值时用(如"居然""只要""竟然")
|
||||||
|
7=轻蔑(禁用)
|
||||||
|
【强度说明】
|
||||||
|
0.1-0.2:几乎看不出变化,比中性多一点情绪色彩
|
||||||
|
0.3-0.4:有明显但自然的情绪,正常说话的波动
|
||||||
|
0.5-0.6:较强情绪,感叹句/重点强调
|
||||||
|
0.7+:极强情绪,极少使用
|
||||||
|
【规则】
|
||||||
|
1. 按自然语义分句,以。!?;为主要分界,逗号不分
|
||||||
|
2. 60-70%的句子应该用中性(4),不要每句都标情绪
|
||||||
|
3. 情绪和内容匹配:卖点→开心(3),痛点共情→伤心(5)低强度,惊喜/划算→惊讶(6),陈述→中性(4)
|
||||||
|
4. 相邻句子情绪不要剧烈跳变
|
||||||
|
5. 感叹号结尾强度0.4-0.6,句号结尾一般0.1-0.3
|
||||||
|
6. 开头结尾句用中性(4)或低强度开心(3)
|
||||||
|
7. 禁止使用0/1/2/7
|
||||||
|
【输出格式】严格JSON数组,不要输出其他内容
|
||||||
|
[{"text":"句子原文","emo":3,"intensity":0.4}]
|
||||||
|
【文案】
|
||||||
|
{文案}
|
||||||
@@ -1,16 +1,19 @@
|
|||||||
"""爆款视频 5 套 Prompt 模板默认值(#2040 核心资产)。
|
"""爆款视频 Prompt 模板默认值(v8 / v3 叙述优先重构)。
|
||||||
|
|
||||||
重要约定(用户明确要求):
|
设计原则(灵应 2026-10-07):LLM 直接输出最终给用户看的文案,代码尽量薄。
|
||||||
- 所有 system_prompt / user_prompt_template / example_output 都是**纯文本自然语言 + XML 标签**,
|
- image_analysis v8:VLM 主交付物是自然叙述风格的 summary_markdown,结构化
|
||||||
运营可直接看懂和编辑,禁止 JSON、禁止 ```json 代码块。
|
字段仅保留 type/name/brand/has_person,顶层 products 改名 images;
|
||||||
- LLM 按 XML 标签输出字段,程序用正则解析(见 xml_parser.py)。
|
- storyboard v3:口播台词口语化、画面描述有画面感,copy_display_markdown 是
|
||||||
- user_prompt_template 中花括号占位符(如 {user_copy_text})在运行时填充。
|
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。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
TEMPLATE_VERSION = 1
|
|
||||||
|
|
||||||
# 所有文案类 Prompt 自动注入的硬约束
|
# 所有文案类 Prompt 自动注入的硬约束
|
||||||
GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
||||||
1. 不编造时间:不写“今年最新”“2024 爆款”等会过时的时间表述。
|
1. 不编造时间:不写“今年最新”“2024 爆款”等会过时的时间表述。
|
||||||
@@ -19,7 +22,7 @@ GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
|||||||
4. 符合广告法及平台社区规范。
|
4. 符合广告法及平台社区规范。
|
||||||
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
|
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
|
||||||
|
|
||||||
# 反套路化要求
|
# 负向提示(注入 storyboard / 视频生成负面词)
|
||||||
NEGATIVE_RULES = """【反套路化要求】
|
NEGATIVE_RULES = """【反套路化要求】
|
||||||
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
|
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
|
||||||
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
|
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
|
||||||
@@ -27,304 +30,211 @@ NEGATIVE_RULES = """【反套路化要求】
|
|||||||
# 输出禁用套路词(测试会检查)
|
# 输出禁用套路词(测试会检查)
|
||||||
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
|
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
|
||||||
|
|
||||||
# 文案融合三档独立指令段
|
# 文案融合三档独立指令段(storyboard 一次生成,按档位注入风格指令)
|
||||||
FUSION_INSTRUCTIONS = {
|
FUSION_INSTRUCTIONS = {
|
||||||
"ai_full": """【本次创作模式:AI 全权创作】
|
"ai_full": """【本次创作模式:AI 全权创作】
|
||||||
你是资深短视频编导。用户只提供了产品图片,没有给出具体文案方向。请根据图片内容和营销参数,自由发挥创作完整的爆款短视频文案。充分挖掘产品真实可见的卖点,使用爆款结构,抓人眼球。""",
|
你是资深短视频编导。用户只提供了产品/门店图片,没有给出具体文案方向。请根据图片的真实观察和营销参数,自由发挥创作完整成片级方案,口播自然、画面可拍。""",
|
||||||
"ai_polish": """【本次创作模式:AI 辅助润色】
|
"ai_polish": """【本次创作模式:AI 辅助润色】
|
||||||
你是用户的文案助理。用户已经写了草稿/关键词/碎碎念,表达了他想讲的核心意思,但表达不完整、不够吸引人。你的任务是:以用户的意思为主,保留他想表达的所有核心信息点,在此基础上润色扩写、调整语序、增加衔接、优化表达,让文案更流畅更有吸引力。绝对不能改变用户想表达的核心意思,不能把用户的观点换成相反的,不能添加用户没提到的产品卖点。用户提到的品牌名、价格、人名、具体事实必须原样保留。""",
|
用户已给出方向或碎碎念。以用户的意思为主,保留其所有核心信息,在此基础上润色、补衔接、优化表达,让口播更自然、画面更具体;绝不改变用户核心意思,不添加用户没提到的卖点,品牌名、价格、人名等事实原样保留。""",
|
||||||
"user_primary": """【本次创作模式:以用户原文为主】
|
"user_primary": """【本次创作模式:以用户原文为主】
|
||||||
你是文案润色助手。用户已经写好了明确的文案,这是他最终想表达的内容。你的任务是最小化修改:只做必要的错别字修正、标点调整、语句通顺度优化,以及添加必要的衔接词让口播更自然。用户的核心句子、关键表述、事实信息一律不改。如果用户文案本身已经很好,直接返回,不要为了改而改。personal_brands 中的事实信息必须逐字保留。""",
|
最小化修改:只做必要的通顺、合规修正与衔接补全,用户的核心句子与事实一律不改;用户文案已经很好就直接用,不为改而改。""",
|
||||||
}
|
}
|
||||||
|
|
||||||
# ── 模板1:图片多模态分析(VLM)────────────────────────────────────────
|
# ── 模板1:图片多模态分析 v8(叙述优先)───────────────────────────────
|
||||||
_IMAGE_ANALYSIS_SYSTEM = f"""你是电商商品视觉分析师,负责从商品图片中提取真实可见的商品信息。
|
_IMAGE_ANALYSIS_SYSTEM = """你是一名擅长观察和写作的品牌内容编导。面对一张真实图片,先用眼睛仔细看,再用自然、流畅、具体的中文把画面写成一段可以直接读给人听的描述。
|
||||||
|
|
||||||
工作方式(分步骤看,不要跳步):
|
## 输出格式(严格 JSON,不要输出 JSON 以外的任何内容)
|
||||||
1. 先看整体:有哪些产品、什么场景、有没有人物。
|
{
|
||||||
2. 再看细节:包装文字、颜色构成、人物状态、画面质感。
|
"images": [
|
||||||
3. 最后提炼卖点:只总结图片里能看到的卖点。
|
{
|
||||||
|
"type": "store 或 product 或 person 或 scene,四选一",
|
||||||
|
"name": "主体名称,看不出就写“未识别”",
|
||||||
|
"brand": "品牌名,看不出就留空字符串",
|
||||||
|
"has_person": false,
|
||||||
|
"summary_markdown": "用 Markdown 写成的自然叙述,这是最主要的交付物"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
{GLOBAL_CONSTRAINTS}
|
## summary_markdown 写作要求(最重要)
|
||||||
|
1. 写成完整、通顺的句子,像在跟朋友认真描述你看到的画面;不要用分号堆砌关键词,不要罗列“核心特征:xxx”“主色调:xxx”这类填表式标签。
|
||||||
|
2. 开头先给一句整体定性,让读者立刻明白这是什么场景、什么主体。
|
||||||
|
3. 颜色、材质、形状、部件要具体可感,写到位置和搭配;画面里出现的文字原样读出并自然融进句子,数字、规格、价格精确引用,看不清的不要编造。
|
||||||
|
4. 只写真实看到的内容,不脑补功能、疗效、销量或画面之外的信息。
|
||||||
|
5. 长度控制在 200-500 字。
|
||||||
|
|
||||||
请严格按下面的标签格式输出,标签名一个都不能改,不要输出任何解释,不要用代码块:
|
## 按类型组织内容
|
||||||
<products> 下面每个产品用一个 <product> 标签,属性 name 是产品名、features 是外观特征、position 是 main 或 secondary、image_index 是第几张图(从0开始)。
|
- type=store(门店/店内环境):用以下小标题分段,小标题下写连贯的句子而不是清单:
|
||||||
<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> 标签写画面整体情绪氛围。
|
1.家具陈设
|
||||||
<visible_text> 下面每处可见文字用一个 <text_item> 标签,属性 text 是文字内容、position 是位置。
|
2.商品与标识
|
||||||
<scene> 标签写场景描述。
|
- type=product(商品):按自然段从整体到局部描写——先说是什么、什么品牌,再写包装/外形、颜色与材质、标签文字、可见部件与规格。
|
||||||
<quality> 用一个标签,属性 resolution、lighting、composition、blur 描述画质。
|
- type=person(人物):描述人物身份感、姿态、穿着(上下装/颜色/款式)、动作与所处环境;用于品牌宣传时突出其精神状态。
|
||||||
<key_selling_points> 下面每个卖点用一个 <point> 标签。
|
- type=scene(纯场景/风景):描述空间或风景的构成、色彩、光线、氛围与关键物件。
|
||||||
|
|
||||||
【人物属性硬性要求(has_person=true时必须遵守)】
|
## 判断规则
|
||||||
hair/skin_tone/face_shape/outfit四项绝对禁止填“无法判断”,必须基于图片可见特征给出具体中文描述:
|
- has_person:画面中出现可辨识的真实人物(脸或完整上半身)才为 true,海报/模特立牌/照片里的人不算。
|
||||||
- hair:必须描述发型+发色,如“黑色齐肩直发”“棕色微卷中长发”“深棕色短发”
|
- 一张图只描述其本身;多张图属于同一场景时可呼应,但不编造对应关系。
|
||||||
- skin_tone:必须描述肤色,如“暖调自然肤色”“白皙肤色”“小麦色”
|
- 输出必须是严格 JSON,summary_markdown 是字符串,内部换行用 \\n 表示。"""
|
||||||
- face_shape:必须描述脸型,如“鹅蛋脸”“圆脸”“瓜子脸”“方脸”
|
|
||||||
- outfit:必须描述可见穿着,如“米色翻领衬衫”“白色T恤”“黑色连衣裙”
|
|
||||||
即使局部被遮挡也要根据可见部分合理推断;确实看不清时按最接近的直观印象描述。
|
|
||||||
|
|
||||||
其他非人物属性看不到或无法判断时填“无法判断”,布尔值填false,不要留空标签。
|
_IMAGE_ANALYSIS_USER = """请分析这张图片。
|
||||||
|
图片地址:{image_url}
|
||||||
|
OCR 辅助文字(可能为空,仅供参考,不要照抄错误识别):{ocr_text}
|
||||||
|
|
||||||
【有人物场景输出参考(女性手持商品示例,必须写全10个属性,禁止省略)】
|
严格按系统要求只输出 JSON。"""
|
||||||
<people has_person="true" count="1" gender="女" age_range="青年" hair="黑色齐肩直发" skin_tone="暖调自然肤色" face_shape="鹅蛋脸" outfit="米色翻领衬衫" pose="正面半身,手持商品" expression="面带微笑"/>"""
|
|
||||||
|
|
||||||
_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。
|
_IMAGE_ANALYSIS_EXAMPLE = """{
|
||||||
所属行业:{industry}
|
"images": [
|
||||||
图片地址:
|
{
|
||||||
{image_urls}
|
"type": "store",
|
||||||
|
"name": "御众堂门店",
|
||||||
|
"brand": "御众堂",
|
||||||
|
"has_person": false,
|
||||||
|
"summary_markdown": "###店铺主体\\n这是一家名为“御众堂”的线下门店内部,整体暖木色调……"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}"""
|
||||||
|
|
||||||
按约定的标签格式输出分析结果。"""
|
# ── 模板2:编导级分镜 v3(意图理解 + 分镜一次完成)────────────────────
|
||||||
|
_STORYBOARD_SYSTEM = (
|
||||||
|
"""你是一名懂短视频的编导和口播文案高手。你会拿到图片的真实观察、营销目的和用户参数,请一次性完成对营销意图的理解,并产出可直接拍摄/生成的分镜脚本。不要单独输出“意图解析”,意图要直接体现在台词和分镜里。
|
||||||
|
|
||||||
_IMAGE_ANALYSIS_EXAMPLE = """<products>
|
## 输出格式(XML,严格按结构输出,不要输出额外解释)
|
||||||
<product name="大公鸡头 多功能油污净 625ml" features="红色瓶盖白色瓶身,鸡头图案Logo" position="main" image_index="0"/>
|
<script>
|
||||||
</products>
|
<copy_display_markdown><![CDATA[直接展示给用户看的成片文案,用 Markdown 写成流畅叙述]]></copy_display_markdown>
|
||||||
<colors>
|
<clips>
|
||||||
<color hex="#D32F2F" name="红色" coverage="0.4"/>
|
<clip index="1">
|
||||||
<color hex="#FFFFFF" name="白色" coverage="0.5"/>
|
<time_range>0-3秒</time_range>
|
||||||
</colors>
|
<voiceover>这一镜的口播台词</voiceover>
|
||||||
<people has_person="false" count="0" gender="无法判断" age_range="无法判断" hair="无法判断" skin_tone="无法判断" face_shape="无法判断" outfit="无法判断" pose="无法判断" expression="无法判断"/>
|
<visual>具体、有画面感的镜头描述(主体/动作/镜头运动/景别/光线)</visual>
|
||||||
<mood>干净、实用</mood>
|
<reference_image_index>0</reference_image_index>
|
||||||
<visible_text>
|
</clip>
|
||||||
<text_item text="多功能油污净" position="瓶身正面"/>
|
</clips>
|
||||||
</visible_text>
|
<voiceover_script>把所有 clip 的 voiceover 连成完整口播稿</voiceover_script>
|
||||||
<scene>白底棚拍产品图</scene>
|
<theme>一句话主题</theme>
|
||||||
<quality resolution="高清" lighting="均匀柔和" composition="主体居中" blur="false"/>
|
<negative>"""
|
||||||
<key_selling_points>
|
+ NEGATIVE_RULES
|
||||||
<point>针对重油污设计</point>
|
+ """</negative>
|
||||||
<point>大容量625ml</point>
|
</script>
|
||||||
</key_selling_points>"""
|
|
||||||
|
|
||||||
# ── 模板2:用户文案意图解析(LLM)──────────────────────────────────────
|
## 写作要求
|
||||||
_INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案可能只是几个关键词、碎碎念或者不完整的短句,你要读懂他真正想讲什么。
|
1. 口播台词:像真人面对镜头说话,短句、口语化、有停顿有情绪,开头 3 秒给出钩子;不要书面腔,不要机械报参数。
|
||||||
|
2. 画面描述:写清“观众会看到什么”,有动作、有镜头运动、有景别和光线,具体可拍;不堆砌形容词,不写无法实现的画面。
|
||||||
|
3. copy_display_markdown:直接展示给最终用户的文案,用 Markdown 写成自然、流畅、有感染力的成片成片文案,可用小标题与短句组织;不要做字段列表,不要出现“镜头一/台词:”这类制作说明。
|
||||||
|
4. 内容必须来自图片观察与用户给出的信息,不编造卖点、不夸大、不使用绝对化用语和虚假承诺。
|
||||||
|
5. reference_image_index 填本镜参考图片序号(从 0 开始),没有合适参考图填 -1。
|
||||||
|
6. 分镜数量与时长匹配总时长,节奏紧凑。
|
||||||
|
7. 口播字数硬约束(必须严格遵守):按每秒约 2.5~3 个中文字(正常口播语速)计算:
|
||||||
|
- 5秒视频:voiceover_script 总字数 12~15 字
|
||||||
|
- 10秒视频:voiceover_script 总字数 25~30 字
|
||||||
|
- 15秒视频:voiceover_script 总字数 35~45 字
|
||||||
|
- 20秒视频:voiceover_script 总字数 50~60 字
|
||||||
|
- 30秒视频:voiceover_script 总字数 75~90 字
|
||||||
|
- 每个 clip 的 voiceover 字数按该镜头时长比例分配
|
||||||
|
- 所有 clip 的 voiceover 字数之和必须等于总 voiceover_script 字数
|
||||||
|
- 宁可少写也不要多写,超长会导致 TTS 音频超出视频时长限制
|
||||||
|
8. 镜头数量硬约束(必须严格遵守):
|
||||||
|
- 5秒视频:1~2 个镜头
|
||||||
|
- 10秒视频:3 个镜头
|
||||||
|
- 15秒视频:3~4 个镜头
|
||||||
|
- 20秒视频:4~5 个镜头
|
||||||
|
- 30秒视频:6~8 个镜头
|
||||||
|
9. 时间轴硬约束(必须严格遵守):
|
||||||
|
- 每个 clip 的 time_range 必须写成 "X-Y秒" 格式,X 和 Y 是具体数字
|
||||||
|
- 第一个 clip 必须从 0 秒开始
|
||||||
|
- 最后一个 clip 必须结束于 total_duration 秒
|
||||||
|
- 相邻 clip 首尾相接,不能有间隙也不能重叠
|
||||||
|
- 每个 clip 的时长 = Y - X,必须 >= 2 秒
|
||||||
|
10. 每个 clip 必须分配一个 reference_image_index(从 0 开始的图片序号),没有合适图片填 -1"""
|
||||||
|
)
|
||||||
|
|
||||||
{GLOBAL_CONSTRAINTS}
|
_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 结构输出分镜脚本。"""
|
||||||
|
|
||||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
_STORYBOARD_EXAMPLE = """<script>
|
||||||
<intent_summary> 用用户的语言风格,一句话、30字以内概括核心意图。
|
<copy_display_markdown><![CDATA[# 在御众堂,把松弛的自己一点点找回来
|
||||||
<core_messages> 下面每个核心信息点用一个 <message> 标签,属性 must_keep 为 true 或 false、confidence 为 0 到 1 的小数,标签内容写信息点。
|
产后妈妈最懂那种力不从心,推开门,暖光和一杯热茶先接住了你……]]></copy_display_markdown>
|
||||||
<personal_brands> 把用户提到的具体事实——品牌名、价格、人名、地名、时间、产品名——每条用一个 <brand> 标签,属性 category 取 brand、price、person、place、time、product 之一。这些事实必须原样引用,一个字都不能改。
|
<clips>
|
||||||
<emotion_tone> 写文案的情绪调性。
|
<clip index="1">
|
||||||
<missing_info> 把你认为缺失、后续生成时需要合理推断的信息,每条用一个 <info> 标签;没有就输出空标签。"""
|
<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_USER = """用户原始文案:{user_copy_text}
|
# ── 模板3:文案审核(合规/质量门禁)───────────────────────────────────
|
||||||
所属行业:{industry}
|
_REVIEW_SYSTEM = """你是一名短视频广告合规审核与文案优化专家。审核待审文案:
|
||||||
营销目的:{marketing_purpose}
|
1) 广告法与平台合规(绝对化用语、虚假承诺、医疗功效宣称、导流违规);
|
||||||
图片分析结果(供参考):
|
2) 卖点是否聚焦、逻辑是否通顺、口播是否自然;
|
||||||
{image_analysis}
|
3) 是否有机械堆砌、书面腔、标签化表述。
|
||||||
图片类型推断:{image_category_hint}
|
|
||||||
|
|
||||||
请理解用户意图,按标签格式输出。注意:theme和emotion_tone应与图片类型和营销目的匹配——门店类图片偏向"门店探店/到店体验",商品图偏向"好物分享/产品种草",人物图偏向"穿搭/人物故事"。"""
|
只输出 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 留空。"""
|
||||||
|
|
||||||
_INTENT_EXAMPLE = """<intent_summary>一款厨房去油污神器,喷一喷油污就掉</intent_summary>
|
_REVIEW_USER = """<fusion_text>
|
||||||
<core_messages>
|
{fusion_text}
|
||||||
<message must_keep="true" confidence="0.97">去油污效果好,喷上等几分钟再擦</message>
|
</fusion_text>
|
||||||
<message must_keep="false" confidence="0.7">适合厨房重油污场景</message>
|
|
||||||
</core_messages>
|
|
||||||
<personal_brands>
|
|
||||||
<brand category="product">大公鸡头多功能油污净</brand>
|
|
||||||
<brand category="price">39块钱一瓶</brand>
|
|
||||||
</personal_brands>
|
|
||||||
<emotion_tone>亲切、真实、带分享感</emotion_tone>
|
|
||||||
<missing_info>
|
|
||||||
<info>没有说明具体容量,按图片读出的625ml处理</info>
|
|
||||||
</missing_info>"""
|
|
||||||
|
|
||||||
# ── 模板3:文案融合生成(LLM)──────────────────────────────────────────
|
请审核以上文案。"""
|
||||||
_FUSION_SYSTEM = """你负责为短视频生成营销文案。请按思维链分步完成:先定人设和目标客户,再找卖点,再搭结构,再安排情绪,最后写行动号召,不要一步到位乱写。
|
|
||||||
|
|
||||||
{fusion_instruction}
|
_REVIEW_EXAMPLE = """<review>
|
||||||
|
<passed>false</passed>
|
||||||
{global_constraints}
|
<issues>
|
||||||
|
<issue>
|
||||||
{negative_rules}
|
<severity>high</severity>
|
||||||
|
<field>opening</field>
|
||||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
<problem>使用绝对化用语“全网第一”</problem>
|
||||||
<title> 视频标题。
|
<suggestion>改为“很多老客户回购的一款”</suggestion>
|
||||||
<hook> 开头3秒钩子,5到15字。
|
</issue>
|
||||||
<body_points> 每个要点用一个 <point> 标签,属性 elaboration 是展开说明、image_index 是对应第几张图(从0开始),标签内容写要点。
|
</issues>
|
||||||
<cta> 口语化的行动号召。
|
<rewrite>……</rewrite>
|
||||||
<script_segments> 每段配音用一个 <segment> 标签,属性 duration_sec 是秒数、image_index 是对应图片,标签内容写配音文案(纯口播文本,不加旁白标注、不加镜头标注、不加"主播:"之类前缀)。
|
</review>"""
|
||||||
<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}
|
|
||||||
|
|
||||||
重要:overview_theme 必须与图片实际内容和营销目的匹配。门店/餐饮/服务类图片用"门店探店·到店体验";商品图用"好物分享·产品种草";人物图用"穿搭分享·人物故事";场景图用"空间体验·场景氛围"。不要对所有图片都使用"好物分享"。
|
|
||||||
|
|
||||||
请按标签格式输出分镜。"""
|
|
||||||
|
|
||||||
_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] = [
|
DEFAULT_TEMPLATES: list[dict] = [
|
||||||
{
|
{
|
||||||
"name": "图片多模态分析",
|
"name": "图片多模态分析 v8",
|
||||||
"prompt_type": "image_analysis",
|
"prompt_type": "image_analysis",
|
||||||
"version": TEMPLATE_VERSION,
|
"version": 8,
|
||||||
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
|
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
|
||||||
"user_prompt_template": _IMAGE_ANALYSIS_USER,
|
"user_prompt_template": _IMAGE_ANALYSIS_USER,
|
||||||
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
|
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
|
||||||
"is_active": True,
|
"is_active": True,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"name": "用户文案意图解析",
|
"name": "编导级分镜 v3",
|
||||||
"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",
|
"prompt_type": "storyboard",
|
||||||
"version": TEMPLATE_VERSION,
|
"version": 3,
|
||||||
"system_prompt": _STORYBOARD_SYSTEM,
|
"system_prompt": _STORYBOARD_SYSTEM,
|
||||||
"user_prompt_template": _STORYBOARD_USER,
|
"user_prompt_template": _STORYBOARD_USER,
|
||||||
"example_output": _STORYBOARD_EXAMPLE,
|
"example_output": _STORYBOARD_EXAMPLE,
|
||||||
@@ -333,7 +243,7 @@ DEFAULT_TEMPLATES: list[dict] = [
|
|||||||
{
|
{
|
||||||
"name": "文案审核",
|
"name": "文案审核",
|
||||||
"prompt_type": "review",
|
"prompt_type": "review",
|
||||||
"version": TEMPLATE_VERSION,
|
"version": 1,
|
||||||
"system_prompt": _REVIEW_SYSTEM,
|
"system_prompt": _REVIEW_SYSTEM,
|
||||||
"user_prompt_template": _REVIEW_USER,
|
"user_prompt_template": _REVIEW_USER,
|
||||||
"example_output": _REVIEW_EXAMPLE,
|
"example_output": _REVIEW_EXAMPLE,
|
||||||
|
|||||||
@@ -53,6 +53,9 @@ _LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
|
|||||||
|
|
||||||
|
|
||||||
class Reviewer:
|
class Reviewer:
|
||||||
|
# markdown展示字段不参与合规审核(避免格式字符误判)
|
||||||
|
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
|
||||||
|
|
||||||
def __init__(self, client=None):
|
def __init__(self, client=None):
|
||||||
if client is None:
|
if client is None:
|
||||||
try:
|
try:
|
||||||
@@ -70,8 +73,9 @@ class Reviewer:
|
|||||||
local = self._rule_check(fusion, intent, fusion_level)
|
local = self._rule_check(fusion, intent, fusion_level)
|
||||||
llm_result = self._llm_review(fusion, intent, fusion_level)
|
llm_result = self._llm_review(fusion, intent, fusion_level)
|
||||||
if llm_result is None:
|
if llm_result is None:
|
||||||
|
# LLM审核失败(超时/网络错误等),降级放行,不阻断渲染
|
||||||
return ReviewResult(
|
return ReviewResult(
|
||||||
passed=not local,
|
passed=True,
|
||||||
issues=local,
|
issues=local,
|
||||||
rewrite_suggestions=[],
|
rewrite_suggestions=[],
|
||||||
raw="",
|
raw="",
|
||||||
@@ -86,6 +90,17 @@ class Reviewer:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _llm_review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> Optional[ReviewResult]:
|
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")
|
template = get_template("review")
|
||||||
system = render_system_prompt(template)
|
system = render_system_prompt(template)
|
||||||
user = render_user_prompt(
|
user = render_user_prompt(
|
||||||
@@ -303,10 +318,13 @@ class Reviewer:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _fusion_text(fusion: FusionResult) -> str:
|
def _fusion_text(fusion: FusionResult) -> str:
|
||||||
|
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
|
||||||
parts = [fusion.title, fusion.hook]
|
parts = [fusion.title, fusion.hook]
|
||||||
parts += [p.text for p in fusion.body_points]
|
parts += [p.text for p in fusion.body_points]
|
||||||
parts += [s.text for s in fusion.script_segments]
|
parts += [s.text for s in fusion.script_segments]
|
||||||
parts.append(fusion.cta)
|
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)
|
return "\n".join(p for p in parts if p)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -13,6 +13,13 @@ from typing import Optional
|
|||||||
_OPEN_RE = re.compile(r"<(?P<tag>[\w-]+)(?P<attrs>(?:\s(?:[^>]*?\S)?)?)(?P<self>/?)>")
|
_OPEN_RE = re.compile(r"<(?P<tag>[\w-]+)(?P<attrs>(?:\s(?:[^>]*?\S)?)?)(?P<self>/?)>")
|
||||||
_CLOSE_RE = re.compile(r"</(?P<tag>[\w-]+)\s*>")
|
_CLOSE_RE = re.compile(r"</(?P<tag>[\w-]+)\s*>")
|
||||||
_ATTR_RE = re.compile(r"""([\w:-]+)\s*=\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]:
|
def parse_attributes(raw: str) -> dict[str, str]:
|
||||||
@@ -58,6 +65,7 @@ def parse_tags(text: Optional[str]) -> list[dict]:
|
|||||||
if stack[idx]["tag"] == tag:
|
if stack[idx]["tag"] == tag:
|
||||||
node = stack[idx]
|
node = stack[idx]
|
||||||
node["text"] = unescape(text[node["_start"] : token.start()].strip())
|
node["text"] = unescape(text[node["_start"] : token.start()].strip())
|
||||||
|
node["text"] = _strip_cdata(node["text"])
|
||||||
node.pop("_start", None)
|
node.pop("_start", None)
|
||||||
del stack[idx:]
|
del stack[idx:]
|
||||||
break
|
break
|
||||||
@@ -65,6 +73,7 @@ def parse_tags(text: Optional[str]) -> list[dict]:
|
|||||||
for node in stack:
|
for node in stack:
|
||||||
if "_start" in node:
|
if "_start" in node:
|
||||||
node["text"] = unescape(text[node["_start"] :].strip())
|
node["text"] = unescape(text[node["_start"] :].strip())
|
||||||
|
node["text"] = _strip_cdata(node["text"])
|
||||||
node.pop("_start", None)
|
node.pop("_start", None)
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
|||||||
@@ -173,6 +173,73 @@ class SharedSettings(BaseSettings):
|
|||||||
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
|
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
|
||||||
gpu_worker_stale_seconds: int = 300
|
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 单次请求 read 超时(秒):数字人半身视频推理通常 30-120s(RTF≈2.8,40s音频约112s)
|
||||||
|
# connect 超时固定 10s(代码硬编码,网络不通快速失败)
|
||||||
|
ditto_request_timeout: int = Field(
|
||||||
|
default=120,
|
||||||
|
validation_alias=AliasChoices("DITTO_REQUEST_TIMEOUT", "ditto_request_timeout"),
|
||||||
|
)
|
||||||
|
# Ditto 句间过渡帧数(平滑表情/口型切换)
|
||||||
|
ditto_blend_frames: int = Field(
|
||||||
|
default=12,
|
||||||
|
validation_alias=AliasChoices("DITTO_BLEND_FRAMES", "ditto_blend_frames"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Ditto LLM 情绪分析(emo_timeline)──────────────────────────────
|
||||||
|
# 总开关;关闭或 LLM 失败时走 GPU 端关键词匹配兜底
|
||||||
|
ditto_emotion_enabled: bool = Field(
|
||||||
|
default=False,
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_ENABLED", "ditto_emotion_enabled"),
|
||||||
|
)
|
||||||
|
ditto_emotion_model: str = Field(
|
||||||
|
default="doubao-seed-2-1-lite-250915",
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_MODEL", "ditto_emotion_model"),
|
||||||
|
)
|
||||||
|
ditto_emotion_temperature: float = Field(
|
||||||
|
default=0.1,
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_TEMPERATURE", "ditto_emotion_temperature"),
|
||||||
|
)
|
||||||
|
ditto_emotion_timeout: int = Field(
|
||||||
|
default=10,
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_TIMEOUT", "ditto_emotion_timeout"),
|
||||||
|
)
|
||||||
|
ditto_emotion_max_tokens: int = Field(
|
||||||
|
default=1024,
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_MAX_TOKENS", "ditto_emotion_max_tokens"),
|
||||||
|
)
|
||||||
|
ditto_emotion_cache_size: int = Field(
|
||||||
|
default=500,
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_CACHE_SIZE", "ditto_emotion_cache_size"),
|
||||||
|
)
|
||||||
|
# 提示词模板:必须包含 {文案} 占位符;后台可通过环境变量覆盖
|
||||||
|
ditto_emotion_prompt: str = Field(
|
||||||
|
default="",
|
||||||
|
validation_alias=AliasChoices("DITTO_EMOTION_PROMPT", "ditto_emotion_prompt"),
|
||||||
|
)
|
||||||
|
|
||||||
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
|
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
|
||||||
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
|
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
|
||||||
enable_gpu_encode: bool = Field(
|
enable_gpu_encode: bool = Field(
|
||||||
|
|||||||
@@ -8,7 +8,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import UTC, datetime
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
try:
|
||||||
|
from datetime import UTC
|
||||||
|
except ImportError:
|
||||||
|
UTC = timezone.utc
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -337,13 +337,13 @@ class DoubaoClient:
|
|||||||
self.last_finish_reason = finish_reason
|
self.last_finish_reason = finish_reason
|
||||||
_elapsed = time.time() - _t0
|
_elapsed = time.time() - _t0
|
||||||
logger.info(
|
logger.info(
|
||||||
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%d",
|
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%s",
|
||||||
payload.get("model"),
|
payload.get("model"),
|
||||||
data.get("usage", {}).get("prompt_tokens", 0),
|
data.get("usage", {}).get("prompt_tokens", 0),
|
||||||
data.get("usage", {}).get("completion_tokens", 0),
|
data.get("usage", {}).get("completion_tokens", 0),
|
||||||
_elapsed,
|
_elapsed,
|
||||||
attempt + 1,
|
attempt + 1,
|
||||||
_req_timeout,
|
getattr(_req_timeout, "read", _req_timeout),
|
||||||
)
|
)
|
||||||
return content.strip()
|
return content.strip()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -636,6 +636,10 @@ class DoubaoClient:
|
|||||||
resolution=resolution,
|
resolution=resolution,
|
||||||
output_dir=output_dir,
|
output_dir=output_dir,
|
||||||
model=video_model,
|
model=video_model,
|
||||||
|
generate_audio=bool(generate_audio),
|
||||||
|
reference_images=reference_images,
|
||||||
|
reference_audios=reference_audios,
|
||||||
|
reference_videos=reference_videos,
|
||||||
)
|
)
|
||||||
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
||||||
self.last_video_error = dict(ds.last_video_error)
|
self.last_video_error = dict(ds.last_video_error)
|
||||||
|
|||||||
@@ -51,6 +51,8 @@ task_routes = {
|
|||||||
"ai_avatar_render.execute": {"queue": QUEUE_GENERATION},
|
"ai_avatar_render.execute": {"queue": QUEUE_GENERATION},
|
||||||
# GPU MuseTalk 口型同步(用户等成片,链路子任务全部走 generation 避免跨队列阻塞)
|
# GPU MuseTalk 口型同步(用户等成片,链路子任务全部走 generation 避免跨队列阻塞)
|
||||||
"lipsync_gpu_process_async": {"queue": QUEUE_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.synthesize_and_submit": {"queue": QUEUE_GENERATION},
|
||||||
"lipsync_tts.poll_mediakit_status": {"queue": QUEUE_GENERATION},
|
"lipsync_tts.poll_mediakit_status": {"queue": QUEUE_GENERATION},
|
||||||
"lipsync_tts.persist_output_video": {"queue": QUEUE_GENERATION},
|
"lipsync_tts.persist_output_video": {"queue": QUEUE_GENERATION},
|
||||||
|
|||||||
@@ -115,10 +115,20 @@ class DashScopeClient:
|
|||||||
watermark: bool = False,
|
watermark: bool = False,
|
||||||
output_dir: str | None = None,
|
output_dir: str | None = None,
|
||||||
model: str = "wan3.0-video",
|
model: str = "wan3.0-video",
|
||||||
|
generate_audio: bool = True,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
reference_audios: list[str] | None = None,
|
||||||
|
reference_videos: list[str] | None = None,
|
||||||
) -> dict | None:
|
) -> dict | None:
|
||||||
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
|
"""调用 DashScope Wan 3.0 异步视频合成接口,轮询完成后下载到本地。
|
||||||
|
|
||||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误详情写入 self.last_video_error。
|
官方协议(input.media 数组 + parameters.audio):
|
||||||
|
- 仅 1 张图且无其它参考 -> type=first_frame(首帧模式,严格从该帧起)。
|
||||||
|
- 有参考音频 / 多张图 -> 图片全部走 type=reference_image(全能参考模式,
|
||||||
|
可与 reference_audio 共存);prompt 用"图1/图2/音频1"按 media 顺序引用。
|
||||||
|
- parameters.audio 控制输出是否含音轨;参考音频通过 media 传入。
|
||||||
|
|
||||||
|
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误写入 self.last_video_error。
|
||||||
"""
|
"""
|
||||||
self.last_video_error = {}
|
self.last_video_error = {}
|
||||||
if not self.is_available:
|
if not self.is_available:
|
||||||
@@ -130,7 +140,7 @@ class DashScopeClient:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
||||||
res_upper = (resolution or "720p").upper().replace("P", "P")
|
res_upper = (resolution or "720p").upper()
|
||||||
if res_upper == "480P":
|
if res_upper == "480P":
|
||||||
ds_res = "480P"
|
ds_res = "480P"
|
||||||
elif res_upper == "1080P":
|
elif res_upper == "1080P":
|
||||||
@@ -138,16 +148,35 @@ class DashScopeClient:
|
|||||||
else:
|
else:
|
||||||
ds_res = "720P"
|
ds_res = "720P"
|
||||||
|
|
||||||
# 构造 input+parameters
|
# ── 构造官方 media 数组 ────────────────────────────────────────
|
||||||
|
ref_imgs = [u for u in (reference_images or [])[:10] if u]
|
||||||
|
ref_auds = [u for u in (reference_audios or [])[:5] if u]
|
||||||
|
ref_vids = [u for u in (reference_videos or [])[:5] if u]
|
||||||
|
|
||||||
|
media: list[dict[str, Any]] = []
|
||||||
|
all_imgs = ([image_url] if image_url else []) + [u for u in ref_imgs if u != image_url]
|
||||||
|
use_first_frame = bool(image_url) and len(all_imgs) == 1 and not (ref_auds or ref_vids)
|
||||||
|
if use_first_frame:
|
||||||
|
media.append({"type": "first_frame", "url": image_url})
|
||||||
|
else:
|
||||||
|
for u in all_imgs:
|
||||||
|
media.append({"type": "reference_image", "url": u})
|
||||||
|
for u in ref_vids:
|
||||||
|
media.append({"type": "reference_video", "url": u})
|
||||||
|
for u in ref_auds:
|
||||||
|
media.append({"type": "reference_audio", "url": u})
|
||||||
|
|
||||||
|
# ── input + parameters ─────────────────────────────────────────
|
||||||
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
|
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
|
||||||
if image_url:
|
if media:
|
||||||
input_obj["img_url"] = image_url
|
input_obj["media"] = media
|
||||||
|
|
||||||
params: dict[str, Any] = {
|
params: dict[str, Any] = {
|
||||||
"resolution": ds_res,
|
"resolution": ds_res,
|
||||||
"duration": str(float(duration)),
|
"duration": int(duration),
|
||||||
"watermark": bool(watermark),
|
"watermark": bool(watermark),
|
||||||
|
"audio": bool(generate_audio),
|
||||||
}
|
}
|
||||||
# 比例透传:Wan 支持 "9:16" / "16:9" / "1:1" 等
|
|
||||||
if ratio and ratio != "adaptive":
|
if ratio and ratio != "adaptive":
|
||||||
params["aspect_ratio"] = ratio
|
params["aspect_ratio"] = ratio
|
||||||
|
|
||||||
@@ -163,14 +192,15 @@ class DashScopeClient:
|
|||||||
}
|
}
|
||||||
create_url = f"{self.base_url}/services/aigc/video-generation/video-synthesis"
|
create_url = f"{self.base_url}/services/aigc/video-generation/video-synthesis"
|
||||||
logger.info(
|
logger.info(
|
||||||
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s img=%s",
|
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s media=%d audio=%s",
|
||||||
model,
|
model,
|
||||||
duration,
|
duration,
|
||||||
ratio,
|
ratio,
|
||||||
ds_res,
|
ds_res,
|
||||||
bool(image_url),
|
len(media),
|
||||||
|
generate_audio,
|
||||||
)
|
)
|
||||||
logger.info("[dashscope] 创建任务 payload: model=%s params=%s", model, params)
|
logger.info("[dashscope] media types: %s", [m["type"] for m in media])
|
||||||
|
|
||||||
# 创建任务
|
# 创建任务
|
||||||
task_id: str | None = None
|
task_id: str | None = None
|
||||||
@@ -196,27 +226,14 @@ class DashScopeClient:
|
|||||||
if tid:
|
if tid:
|
||||||
task_id = tid
|
task_id = tid
|
||||||
break
|
break
|
||||||
# 部分情况下 code != 错误
|
|
||||||
code = data.get("code")
|
code = data.get("code")
|
||||||
if code and code != "":
|
if code:
|
||||||
err_code, user_msg = _classify_dashscope_error(400, body_text, str(code))
|
logger.error("[dashscope] 创建任务返回 code=%s body=%s", code, body_text)
|
||||||
|
err_code, user_msg = _classify_dashscope_error(sc, body_text)
|
||||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||||
return None
|
return None
|
||||||
else:
|
|
||||||
self._set_error("unknown", "Wan 3.0 响应格式异常,未返回任务ID", sc, str(data)[:500], model=model)
|
|
||||||
return None
|
|
||||||
except _HTTP_NETWORK_ERRORS as ne:
|
|
||||||
last_sc = 0
|
|
||||||
last_body = f"network error: {ne}"
|
|
||||||
logger.warning(
|
|
||||||
"[dashscope] 网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
|
||||||
)
|
|
||||||
if attempt < self.max_retries:
|
|
||||||
time.sleep(0.5 * (2**attempt))
|
|
||||||
continue
|
|
||||||
self._set_error("network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。", 0, str(ne))
|
|
||||||
return None
|
|
||||||
except Exception as _e:
|
except Exception as _e:
|
||||||
|
logger.warning("[dashscope] 创建任务异常(attempt=%d): %s", attempt, _e)
|
||||||
if attempt < self.max_retries:
|
if attempt < self.max_retries:
|
||||||
time.sleep(0.5 * (2**attempt))
|
time.sleep(0.5 * (2**attempt))
|
||||||
continue
|
continue
|
||||||
@@ -258,7 +275,6 @@ class DashScopeClient:
|
|||||||
video_url = out.get("video_url") or ""
|
video_url = out.get("video_url") or ""
|
||||||
usage = d.get("usage")
|
usage = d.get("usage")
|
||||||
if not video_url:
|
if not video_url:
|
||||||
# 结果在 results 数组
|
|
||||||
results = out.get("results") or []
|
results = out.get("results") or []
|
||||||
if results and isinstance(results, list):
|
if results and isinstance(results, list):
|
||||||
video_url = results[0].get("url") or results[0].get("video_url")
|
video_url = results[0].get("url") or results[0].get("video_url")
|
||||||
@@ -284,7 +300,6 @@ class DashScopeClient:
|
|||||||
logger.warning("[dashscope] 任务 %s 被取消", task_id)
|
logger.warning("[dashscope] 任务 %s 被取消", task_id)
|
||||||
self._set_error("unknown", "Wan 3.0 任务被取消。", 200, "task cancelled", task_id=task_id)
|
self._set_error("unknown", "Wan 3.0 任务被取消。", 200, "task cancelled", task_id=task_id)
|
||||||
return None
|
return None
|
||||||
# PENDING / RUNNING / SUSPENDED → 继续轮询
|
|
||||||
if poll_count % 5 == 0:
|
if poll_count % 5 == 0:
|
||||||
logger.info("[dashscope] 轮询中 task=%s status=%s polls=%d", task_id, task_status, poll_count)
|
logger.info("[dashscope] 轮询中 task=%s status=%s polls=%d", task_id, task_status, poll_count)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -334,7 +334,7 @@ class SharedStorageService(StoragePort):
|
|||||||
|
|
||||||
storage_key = self.normalize_storage_key(storage_key_or_url)
|
storage_key = self.normalize_storage_key(storage_key_or_url)
|
||||||
try:
|
try:
|
||||||
signed = sign_bucket.sign_url("GET", storage_key, expires_seconds)
|
signed = sign_bucket.sign_url("GET", storage_key, expires_seconds, slash_safe=True)
|
||||||
logger.info(
|
logger.info(
|
||||||
"signed URL generated for key=%s prefix=%s",
|
"signed URL generated for key=%s prefix=%s",
|
||||||
storage_key[:80],
|
storage_key[:80],
|
||||||
|
|||||||
+233
-250
@@ -1,209 +1,210 @@
|
|||||||
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
|
"""AI Router 单元测试 — routing/cache/fallback/client construction.
|
||||||
|
|
||||||
|
本文件只做*用例级* mock:通过 autouse fixture 在每个用例内 patch
|
||||||
|
``packages.shared.config.get_shared_settings`` / ``packages.shared.ai_router.get_shared_settings``
|
||||||
|
并在退出时自动恢复,绝不在模块顶层替换 ``sys.modules``,因此不会污染同进程的
|
||||||
|
其他测试模块(如 test_ai_client.py)。
|
||||||
|
|
||||||
|
在 Python 3.12 且依赖齐全的 CI 环境中,直接 import 真实模块即可;Redis / DB
|
||||||
|
会话通过 patch 隔离。
|
||||||
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import sys
|
|
||||||
import unittest
|
import unittest
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Optional
|
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
# ── Pre-mock heavy import chain to avoid pulling in full app ──
|
import pytest
|
||||||
_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
|
from packages.shared import ai_config_version as _config_version_mod
|
||||||
for mod_name in list(sys.modules.keys()):
|
from packages.shared import ai_router as ai_router_mod
|
||||||
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)
|
# ── 统一的假配置(等价于旧文件里的 _mock_settings)──────────────────────────
|
||||||
import importlib.util
|
|
||||||
import os
|
|
||||||
|
|
||||||
|
|
||||||
def _load_module_from_file(name, path):
|
def _make_mock_settings() -> MagicMock:
|
||||||
spec = importlib.util.spec_from_file_location(name, path)
|
s = MagicMock()
|
||||||
mod = importlib.util.module_from_spec(spec)
|
s.doubao_model = "doubao-seed-2-1-pro-260915"
|
||||||
sys.modules[name] = mod
|
s.doubao_fast_model = "doubao-seed-2-1-pro-260915"
|
||||||
spec.loader.exec_module(mod)
|
s.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||||
return mod
|
s.doubao_api_key = "test-key"
|
||||||
|
s.doubao_timeout = 45
|
||||||
|
s.doubao_max_retries = 1
|
||||||
|
s.doubao_image_model = "doubao-seedream-5-0-flash-260915"
|
||||||
|
s.doubao_image_timeout = 60
|
||||||
|
s.doubao_vision_model = "doubao-seed-1-6-vision-250615"
|
||||||
|
s.doubao_video_model = "doubao-seedance-2-5-260628"
|
||||||
|
s.doubao_video_timeout = 600
|
||||||
|
s.dashscope_api_key = "ds-key"
|
||||||
|
s.dashscope_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||||
|
s.dashscope_model = "qwen-vl-max"
|
||||||
|
s.cosyvoice_api_key = "cv-key"
|
||||||
|
s.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||||
|
s.cosyvoice_model = "cosyvoice-v3-flash"
|
||||||
|
s.redis_url = "redis://localhost:6379/0"
|
||||||
|
s.celery_broker_url = "redis://localhost:6379/0"
|
||||||
|
return s
|
||||||
|
|
||||||
|
|
||||||
# Load ai_config_version
|
@pytest.fixture(autouse=True)
|
||||||
_ai_config_version = _load_module_from_file(
|
def _mock_settings_fixture():
|
||||||
"packages.shared.ai_config_version",
|
"""每个用例内 patch 配置来源,退出即恢复,不污染 sys.modules。"""
|
||||||
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_config_version.py"),
|
settings = _make_mock_settings()
|
||||||
)
|
with (
|
||||||
# Patch get_shared_settings in the loaded module
|
patch("packages.shared.config.get_shared_settings", return_value=settings),
|
||||||
_ai_config_version.get_shared_settings = lambda: _mock_settings
|
patch.object(ai_router_mod, "get_shared_settings", return_value=settings),
|
||||||
|
):
|
||||||
|
yield 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
|
def _capability_row() -> MagicMock:
|
||||||
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
|
row = MagicMock()
|
||||||
_mock_ai_client = MagicMock()
|
row.capability_key = "intent_parsing"
|
||||||
|
row.capability_name = "文案意图解析"
|
||||||
|
row.timeout_seconds = 45
|
||||||
|
row.max_retries = 1
|
||||||
|
row.max_tokens = None
|
||||||
|
row.temperature = None
|
||||||
|
row.concurrency = 2
|
||||||
|
row.extra_params = {}
|
||||||
|
row.is_enabled = True
|
||||||
|
row.pm_id = "model-1"
|
||||||
|
row.pm_name = "豆包"
|
||||||
|
row.pm_provider = "volcengine"
|
||||||
|
row.pm_model_key = "doubao-seed-1-6-250615"
|
||||||
|
row.pm_api_key = "test-key"
|
||||||
|
row.pm_api_base = "https://ark.test.com"
|
||||||
|
row.pm_api_version = None
|
||||||
|
row.pm_status = "active"
|
||||||
|
row.lm_id = None
|
||||||
|
row.fm_id = None
|
||||||
|
return row
|
||||||
|
|
||||||
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 _model_config(**overrides):
|
||||||
def is_available(self):
|
kwargs = dict(
|
||||||
return bool(self.api_key)
|
id="m1",
|
||||||
|
name="test",
|
||||||
|
provider="volcengine",
|
||||||
|
model_key="test-model",
|
||||||
|
api_key="key",
|
||||||
|
api_base="https://test.com",
|
||||||
|
api_version=None,
|
||||||
|
status="active",
|
||||||
|
)
|
||||||
|
kwargs.update(overrides)
|
||||||
|
return ai_router_mod.ModelConfig(**kwargs)
|
||||||
|
|
||||||
def chat_completion(self, messages, **kwargs):
|
|
||||||
return None
|
|
||||||
|
|
||||||
def vision_completion(self, messages, **kwargs):
|
def _capability_config(**overrides):
|
||||||
return None
|
kwargs = dict(
|
||||||
|
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=True,
|
||||||
|
)
|
||||||
|
kwargs.update(overrides)
|
||||||
|
return ai_router_mod.CapabilityConfig(**kwargs)
|
||||||
|
|
||||||
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
|
|
||||||
sys.modules["packages.shared.ai_client"] = _mock_ai_client
|
|
||||||
|
|
||||||
_ai_router = _load_module_from_file(
|
# ── Redis 版本号机制 ────────────────────────────────────────────────────────
|
||||||
"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):
|
class TestAIConfigVersion(unittest.TestCase):
|
||||||
"""Redis 版本号机制测试"""
|
"""Redis 版本号机制测试"""
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "_get_redis_client")
|
@patch.object(_config_version_mod, "_get_redis_client")
|
||||||
def test_bump_version_success(self, mock_redis_fn):
|
def test_bump_version_success(self, mock_redis_fn):
|
||||||
mock_r = MagicMock()
|
mock_r = MagicMock()
|
||||||
mock_r.set.return_value = True
|
mock_r.set.return_value = True
|
||||||
mock_redis_fn.return_value = mock_r
|
mock_redis_fn.return_value = mock_r
|
||||||
ver = _ai_config_version.bump_version()
|
ver = _config_version_mod.bump_version()
|
||||||
self.assertTrue(ver)
|
self.assertTrue(ver)
|
||||||
self.assertTrue(ver.isdigit())
|
self.assertTrue(ver.isdigit())
|
||||||
mock_r.set.assert_called_once()
|
mock_r.set.assert_called_once()
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "_get_redis_client")
|
@patch.object(_config_version_mod, "_get_redis_client")
|
||||||
def test_bump_version_redis_unavailable(self, mock_redis_fn):
|
def test_bump_version_redis_unavailable(self, mock_redis_fn):
|
||||||
mock_redis_fn.return_value = None
|
mock_redis_fn.return_value = None
|
||||||
ver = _ai_config_version.bump_version()
|
ver = _config_version_mod.bump_version()
|
||||||
self.assertEqual(ver, "")
|
self.assertEqual(ver, "")
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "_get_redis_client")
|
@patch.object(_config_version_mod, "_get_redis_client")
|
||||||
def test_get_version_success(self, mock_redis_fn):
|
def test_get_version_success(self, mock_redis_fn):
|
||||||
mock_r = MagicMock()
|
mock_r = MagicMock()
|
||||||
mock_r.get.return_value = "1234567890"
|
mock_r.get.return_value = "1234567890"
|
||||||
mock_redis_fn.return_value = mock_r
|
mock_redis_fn.return_value = mock_r
|
||||||
ver = _ai_config_version.get_version()
|
ver = _config_version_mod.get_version()
|
||||||
self.assertEqual(ver, "1234567890")
|
self.assertEqual(ver, "1234567890")
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "_get_redis_client")
|
@patch.object(_config_version_mod, "_get_redis_client")
|
||||||
def test_get_version_redis_down(self, mock_redis_fn):
|
def test_get_version_redis_down(self, mock_redis_fn):
|
||||||
mock_redis_fn.return_value = None
|
mock_redis_fn.return_value = None
|
||||||
ver = _ai_config_version.get_version()
|
ver = _config_version_mod.get_version()
|
||||||
self.assertIsNone(ver)
|
self.assertIsNone(ver)
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "_get_redis_client")
|
@patch.object(_config_version_mod, "_get_redis_client")
|
||||||
def test_get_version_exception(self, mock_redis_fn):
|
def test_get_version_exception(self, mock_redis_fn):
|
||||||
mock_r = MagicMock()
|
mock_r = MagicMock()
|
||||||
mock_r.get.side_effect = Exception("connection refused")
|
mock_r.get.side_effect = Exception("connection refused")
|
||||||
mock_redis_fn.return_value = mock_r
|
mock_redis_fn.return_value = mock_r
|
||||||
ver = _ai_config_version.get_version()
|
ver = _config_version_mod.get_version()
|
||||||
self.assertIsNone(ver)
|
self.assertIsNone(ver)
|
||||||
|
|
||||||
|
|
||||||
|
# ── AIRouter 路由/缓存/fallback ────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
class TestAIRouter(unittest.TestCase):
|
class TestAIRouter(unittest.TestCase):
|
||||||
"""AIRouter 路由/缓存/fallback 测试"""
|
"""AIRouter 路由/缓存/fallback 测试"""
|
||||||
|
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.router = _ai_router.AIRouter()
|
self.router = ai_router_mod.AIRouter()
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def _freeze_version(self, value=None):
|
||||||
def test_get_capability_db_unavailable(self, mock_ver):
|
"""让 get_capability 的版本比对固定,避免走 Redis。"""
|
||||||
with patch.object(_ai_router, "_get_session", return_value=None):
|
return patch.object(_config_version_mod, "get_version", return_value=value)
|
||||||
cap = self.router.get_capability("intent_parsing")
|
|
||||||
self.assertIsNone(cap)
|
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def test_get_capability_db_unavailable(self):
|
||||||
def test_get_capability_from_db(self, mock_ver):
|
with self._freeze_version(None):
|
||||||
|
with patch.object(ai_router_mod, "_get_session", return_value=None):
|
||||||
|
cap = self.router.get_capability("intent_parsing")
|
||||||
|
self.assertIsNone(cap)
|
||||||
|
|
||||||
|
def test_get_capability_from_db(self):
|
||||||
mock_session = MagicMock()
|
mock_session = MagicMock()
|
||||||
mock_row = MagicMock()
|
mock_session.execute.return_value.first.return_value = _capability_row()
|
||||||
mock_row.capability_key = "intent_parsing"
|
with self._freeze_version(None):
|
||||||
mock_row.capability_name = "文案意图解析"
|
with patch.object(ai_router_mod, "_get_session", return_value=mock_session):
|
||||||
mock_row.timeout_seconds = 45
|
cap = self.router.get_capability("intent_parsing")
|
||||||
mock_row.max_retries = 1
|
self.assertIsNotNone(cap)
|
||||||
mock_row.max_tokens = None
|
self.assertEqual(cap.capability_key, "intent_parsing")
|
||||||
mock_row.temperature = None
|
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
|
||||||
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):
|
def test_cache_invalidation_on_version_change(self):
|
||||||
cap = self.router.get_capability("intent_parsing")
|
with self._freeze_version(None):
|
||||||
self.assertIsNotNone(cap)
|
with patch.object(self.router, "_load_from_db", return_value=None):
|
||||||
self.assertEqual(cap.capability_key, "intent_parsing")
|
self.router.get_capability("test_key")
|
||||||
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.router._local_ver = "v1"
|
||||||
self.assertTrue(self.router._check_version())
|
with patch.object(_config_version_mod, "get_version", return_value="v2"):
|
||||||
|
self.assertTrue(self.router._check_version())
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value="same_ver")
|
def test_cache_hit_same_version(self):
|
||||||
def test_cache_hit_same_version(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
primary_model=_model_config(model_key="test-model"),
|
||||||
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._cache["test"] = cap
|
||||||
self.router._local_ver = "same_ver"
|
self.router._local_ver = "same_ver"
|
||||||
result = self.router.get_capability("test")
|
with patch.object(_config_version_mod, "get_version", return_value="same_ver"):
|
||||||
|
result = self.router.get_capability("test")
|
||||||
self.assertEqual(result, cap)
|
self.assertEqual(result, cap)
|
||||||
|
|
||||||
def test_invalidate_clears_cache(self):
|
def test_invalidate_clears_cache(self):
|
||||||
@@ -213,181 +214,163 @@ class TestAIRouter(unittest.TestCase):
|
|||||||
self.assertEqual(len(self.router._cache), 0)
|
self.assertEqual(len(self.router._cache), 0)
|
||||||
self.assertIsNone(self.router._local_ver)
|
self.assertIsNone(self.router._local_ver)
|
||||||
|
|
||||||
@patch.object(_ai_router, "_get_session", return_value=None)
|
def test_get_llm_client_fallback(self):
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
with self._freeze_version(None):
|
||||||
def test_get_llm_client_fallback(self, mock_ver, mock_session):
|
with patch.object(ai_router_mod, "_get_session", return_value=None):
|
||||||
_ai_router.get_shared_settings = lambda: _mock_settings
|
client = self.router.get_llm_client("intent_parsing")
|
||||||
client = self.router.get_llm_client("intent_parsing")
|
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||||
self.assertEqual(client.api_key, "test-key")
|
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):
|
||||||
def test_get_llm_client_from_db(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
capability_key="image_analysis",
|
||||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
capability_name="图片分析",
|
||||||
api_key="db-key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
max_tokens=350,
|
||||||
)
|
temperature=0.1,
|
||||||
cap = _ai_router.CapabilityConfig(
|
primary_model=_model_config(
|
||||||
capability_key="image_analysis", capability_name="图片分析",
|
provider="dashscope",
|
||||||
primary_model=model, lite_model=None, fallback_model=None,
|
model_key="qwen3.8-flash",
|
||||||
timeout_seconds=15, max_retries=1, max_tokens=350, temperature=0.1,
|
api_key="db-key",
|
||||||
concurrency=2, extra_params={}, is_enabled=True,
|
api_base="https://dashscope.test.com",
|
||||||
|
),
|
||||||
)
|
)
|
||||||
with patch.object(self.router, "get_capability", return_value=cap):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_llm_client("image_analysis")
|
client = self.router.get_llm_client("image_analysis")
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "qwen3.8-flash")
|
self.assertEqual(client.model, "qwen3.8-flash")
|
||||||
self.assertEqual(client.provider, "dashscope")
|
self.assertEqual(client.provider, "dashscope")
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def test_get_vision_client(self):
|
||||||
def test_get_vision_client(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
capability_key="image_analysis",
|
||||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
capability_name="图片分析",
|
||||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
primary_model=_model_config(
|
||||||
)
|
provider="dashscope",
|
||||||
cap = _ai_router.CapabilityConfig(
|
model_key="qwen3.8-flash",
|
||||||
capability_key="image_analysis", capability_name="图片分析",
|
api_key="key",
|
||||||
primary_model=model, lite_model=None, fallback_model=None,
|
api_base="https://dashscope.test.com",
|
||||||
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):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_vision_client("image_analysis")
|
client = self.router.get_vision_client("image_analysis")
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
# #2220: vision client is now DoubaoClient with vision_completion
|
# #2220: vision client is now DoubaoClient with vision_completion
|
||||||
self.assertTrue(hasattr(client, "vision_completion"))
|
self.assertTrue(hasattr(client, "vision_completion"))
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def test_get_tts_client(self):
|
||||||
def test_get_tts_client(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
capability_key="tts",
|
||||||
id="m1", name="test", provider="dashscope", model_key="cosyvoice-v3-flash",
|
capability_name="语音合成",
|
||||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
primary_model=_model_config(
|
||||||
)
|
provider="dashscope",
|
||||||
cap = _ai_router.CapabilityConfig(
|
model_key="cosyvoice-v3-flash",
|
||||||
capability_key="tts", capability_name="语音合成",
|
api_key="key",
|
||||||
primary_model=model, lite_model=None, fallback_model=None,
|
api_base="https://dashscope.test.com",
|
||||||
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):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_tts_client()
|
client = self.router.get_tts_client()
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "cosyvoice-v3-flash")
|
self.assertEqual(client.model, "cosyvoice-v3-flash")
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def test_get_image_gen_client(self):
|
||||||
def test_get_image_gen_client(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
capability_key="image_generation",
|
||||||
id="m1", name="test", provider="volcengine", model_key="seedream-5.0-flash",
|
capability_name="图片生成",
|
||||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
extra_params={"size": "1K"},
|
||||||
)
|
primary_model=_model_config(
|
||||||
cap = _ai_router.CapabilityConfig(
|
model_key="seedream-5.0-flash",
|
||||||
capability_key="image_generation", capability_name="图片生成",
|
api_key="key",
|
||||||
primary_model=model, lite_model=None, fallback_model=None,
|
api_base="https://ark.test.com",
|
||||||
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):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_image_gen_client()
|
client = self.router.get_image_gen_client()
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "seedream-5.0-flash")
|
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):
|
||||||
def test_get_video_gen_client(self, mock_ver):
|
cap = _capability_config(
|
||||||
model = _ai_router.ModelConfig(
|
capability_key="video_generation",
|
||||||
id="m1", name="test", provider="volcengine", model_key="seedance-2.5",
|
capability_name="视频生成",
|
||||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
concurrency=1,
|
||||||
)
|
primary_model=_model_config(
|
||||||
cap = _ai_router.CapabilityConfig(
|
model_key="seedance-2.5",
|
||||||
capability_key="video_generation", capability_name="视频生成",
|
api_key="key",
|
||||||
primary_model=model, lite_model=None, fallback_model=None,
|
api_base="https://ark.test.com",
|
||||||
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):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_video_gen_client()
|
client = self.router.get_video_gen_client()
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "seedance-2.5")
|
self.assertEqual(client.model, "seedance-2.5")
|
||||||
|
|
||||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
def test_lite_variant_preference(self):
|
||||||
def test_lite_variant_preference(self, mock_ver):
|
cap = _capability_config(
|
||||||
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")
|
capability_key="image_analysis",
|
||||||
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")
|
capability_name="图片分析",
|
||||||
cap = _ai_router.CapabilityConfig(
|
primary_model=_model_config(id="p1", name="pro", model_key="pro-model", api_key="k", api_base="u"),
|
||||||
capability_key="image_analysis", capability_name="图片分析",
|
lite_model=_model_config(id="l1", name="lite", model_key="lite-model", api_key="k", api_base="u"),
|
||||||
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")
|
model = self.router._get_model_or_fallback(cap, "lite")
|
||||||
self.assertEqual(model.model_key, "lite-model")
|
self.assertEqual(model.model_key, "lite-model")
|
||||||
model_primary = self.router._get_model_or_fallback(cap, "primary")
|
model_primary = self.router._get_model_or_fallback(cap, "primary")
|
||||||
self.assertEqual(model_primary.model_key, "pro-model")
|
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):
|
||||||
def test_disabled_capability_returns_fallback(self, mock_ver):
|
cap = _capability_config(is_enabled=False)
|
||||||
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):
|
with patch.object(self.router, "get_capability", return_value=cap):
|
||||||
client = self.router.get_llm_client("test")
|
client = self.router.get_llm_client("test")
|
||||||
self.assertIsNotNone(client)
|
self.assertIsNotNone(client)
|
||||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
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):
|
||||||
def test_fallback_chain_primary_none(self, mock_ver):
|
|
||||||
"""primary_model 为 None 时 fallback 到 fallback_model"""
|
"""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 = _capability_config(
|
||||||
cap = _ai_router.CapabilityConfig(
|
fallback_model=_model_config(id="f1", name="fb", model_key="fb-model", api_key="k", api_base="u"),
|
||||||
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")
|
model = self.router._get_model_or_fallback(cap, "primary")
|
||||||
self.assertEqual(model.model_key, "fb-model")
|
self.assertEqual(model.model_key, "fb-model")
|
||||||
|
|
||||||
|
|
||||||
|
# ── 数据类冻结 ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
class TestModelConfig(unittest.TestCase):
|
class TestModelConfig(unittest.TestCase):
|
||||||
"""数据类测试"""
|
"""数据类测试"""
|
||||||
|
|
||||||
def test_model_config_frozen(self):
|
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")
|
m = _model_config(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b")
|
||||||
with self.assertRaises(AttributeError):
|
with self.assertRaises(AttributeError):
|
||||||
m.model_key = "new"
|
m.model_key = "new"
|
||||||
|
|
||||||
def test_capability_config_frozen(self):
|
def test_capability_config_frozen(self):
|
||||||
c = _ai_router.CapabilityConfig(
|
c = _capability_config()
|
||||||
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):
|
with self.assertRaises(AttributeError):
|
||||||
c.is_enabled = False
|
c.is_enabled = False
|
||||||
|
|
||||||
|
|
||||||
|
# ── 客户端可用性 ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
class TestClientAvailability(unittest.TestCase):
|
class TestClientAvailability(unittest.TestCase):
|
||||||
"""客户端可用性测试"""
|
"""客户端可用性测试"""
|
||||||
|
|
||||||
def test_tts_client_available(self):
|
def test_tts_client_available(self):
|
||||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
c = ai_router_mod.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
||||||
self.assertTrue(c.is_available)
|
self.assertTrue(c.is_available)
|
||||||
|
|
||||||
def test_tts_client_unavailable_no_model(self):
|
def test_tts_client_unavailable_no_model(self):
|
||||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="")
|
c = ai_router_mod.TTSClient(provider="p", api_key="k", base_url="u", model="")
|
||||||
self.assertFalse(c.is_available)
|
self.assertFalse(c.is_available)
|
||||||
|
|
||||||
def test_image_gen_client_unavailable_no_url(self):
|
def test_image_gen_client_unavailable_no_url(self):
|
||||||
c = _ai_router.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
|
c = ai_router_mod.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
|
||||||
self.assertFalse(c.is_available)
|
self.assertFalse(c.is_available)
|
||||||
|
|
||||||
def test_video_gen_client_available(self):
|
def test_video_gen_client_available(self):
|
||||||
c = _ai_router.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
|
c = ai_router_mod.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
|
||||||
self.assertTrue(c.is_available)
|
self.assertTrue(c.is_available)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
"""Ditto LLM 情绪分析服务单元测试."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from packages.application.ditto_emotion_service import (
|
||||||
|
EMO_HAPPY,
|
||||||
|
EMO_NEUTRAL,
|
||||||
|
DittoEmotionService,
|
||||||
|
EmotionSegment,
|
||||||
|
_parse_emotion_json,
|
||||||
|
align_timeline_by_length,
|
||||||
|
align_timeline_by_timings,
|
||||||
|
split_sentences,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── 分句 ─────────────────────────────────────────────────────────
|
||||||
|
class TestSplitSentences:
|
||||||
|
def test_empty(self):
|
||||||
|
assert split_sentences("") == []
|
||||||
|
|
||||||
|
def test_single(self):
|
||||||
|
assert split_sentences("你好。") == ["你好。"]
|
||||||
|
|
||||||
|
def test_multi(self):
|
||||||
|
sents = split_sentences("大家好!今天给大家推荐一款超棒的产品。它真的很好用;不信你试试?")
|
||||||
|
assert len(sents) == 4
|
||||||
|
assert "大家好!" in sents[0]
|
||||||
|
|
||||||
|
def test_english_punct(self):
|
||||||
|
sents = split_sentences("Hello! How are you? I'm fine.")
|
||||||
|
assert len(sents) == 3
|
||||||
|
|
||||||
|
|
||||||
|
# ── JSON 解析 ────────────────────────────────────────────────────
|
||||||
|
class TestParseEmotionJson:
|
||||||
|
def test_valid(self):
|
||||||
|
raw = json.dumps([{"text": "你好", "emo": 4, "intensity": 0.2}])
|
||||||
|
segs = _parse_emotion_json(raw)
|
||||||
|
assert len(segs) == 1
|
||||||
|
assert segs[0].emo == 4
|
||||||
|
assert segs[0].intensity == 0.2
|
||||||
|
|
||||||
|
def test_markdown_wrapped(self):
|
||||||
|
raw = "```json\n" + json.dumps([{"text": "好", "emo": 3, "intensity": 0.5}]) + "\n```"
|
||||||
|
segs = _parse_emotion_json(raw)
|
||||||
|
assert len(segs) == 1
|
||||||
|
assert segs[0].emo == 3
|
||||||
|
|
||||||
|
def test_forbidden_emo_becomes_neutral(self):
|
||||||
|
raw = json.dumps([{"text": "怒", "emo": 0, "intensity": 0.8}])
|
||||||
|
segs = _parse_emotion_json(raw)
|
||||||
|
assert len(segs) == 1
|
||||||
|
assert segs[0].emo == EMO_NEUTRAL
|
||||||
|
|
||||||
|
def test_invalid_json(self):
|
||||||
|
assert _parse_emotion_json("not json") == []
|
||||||
|
|
||||||
|
def test_empty(self):
|
||||||
|
assert _parse_emotion_json("") == []
|
||||||
|
|
||||||
|
def test_intensity_clamp(self):
|
||||||
|
raw = json.dumps([{"text": "a", "emo": 3, "intensity": 1.5}])
|
||||||
|
segs = _parse_emotion_json(raw)
|
||||||
|
assert segs[0].intensity == 1.0
|
||||||
|
|
||||||
|
def test_missing_text_skipped(self):
|
||||||
|
raw = json.dumps([{"emo": 3, "intensity": 0.4}])
|
||||||
|
segs = _parse_emotion_json(raw)
|
||||||
|
assert len(segs) == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ── 时间对齐(按字数比例)────────────────────────────────────────
|
||||||
|
class TestAlignTimelineByLength:
|
||||||
|
def test_basic(self):
|
||||||
|
segs = [
|
||||||
|
EmotionSegment("ab", EMO_NEUTRAL, 0.2),
|
||||||
|
EmotionSegment("cd", EMO_HAPPY, 0.5),
|
||||||
|
]
|
||||||
|
entries = align_timeline_by_length(segs, 4.0)
|
||||||
|
assert len(entries) == 2
|
||||||
|
assert entries[0].start == 0.0
|
||||||
|
assert entries[0].end == 2.0
|
||||||
|
assert entries[1].start == 2.0
|
||||||
|
assert entries[1].end == 4.0
|
||||||
|
assert entries[0].emo == EMO_NEUTRAL
|
||||||
|
assert entries[1].emo == EMO_HAPPY
|
||||||
|
|
||||||
|
def test_empty_segments(self):
|
||||||
|
assert align_timeline_by_length([], 5.0) == []
|
||||||
|
|
||||||
|
def test_zero_duration(self):
|
||||||
|
segs = [EmotionSegment("ab", EMO_NEUTRAL, 0.2)]
|
||||||
|
assert align_timeline_by_length(segs, 0) == []
|
||||||
|
|
||||||
|
def test_unequal_length(self):
|
||||||
|
segs = [
|
||||||
|
EmotionSegment("a" * 3, EMO_HAPPY, 0.5),
|
||||||
|
EmotionSegment("b" * 1, EMO_NEUTRAL, 0.2),
|
||||||
|
]
|
||||||
|
entries = align_timeline_by_length(segs, 4.0)
|
||||||
|
assert entries[0].end == 3.0
|
||||||
|
assert entries[1].start == 3.0
|
||||||
|
assert entries[1].end == 4.0
|
||||||
|
|
||||||
|
|
||||||
|
# ── 时间对齐(sentence_timings)──────────────────────────────────
|
||||||
|
class TestAlignTimelineByTimings:
|
||||||
|
def test_exact_match(self):
|
||||||
|
segs = [
|
||||||
|
EmotionSegment("hello", EMO_HAPPY, 0.4),
|
||||||
|
EmotionSegment("world", EMO_NEUTRAL, 0.2),
|
||||||
|
]
|
||||||
|
timings = [
|
||||||
|
{"start": 0.0, "end": 1.5},
|
||||||
|
{"start": 1.5, "end": 3.0},
|
||||||
|
]
|
||||||
|
entries = align_timeline_by_timings(segs, timings, 3.0)
|
||||||
|
assert len(entries) == 2
|
||||||
|
assert entries[0].start == 0.0
|
||||||
|
assert entries[0].end == 1.5
|
||||||
|
assert entries[1].start == 1.5
|
||||||
|
assert entries[1].end == 3.0
|
||||||
|
|
||||||
|
def test_length_mismatch_fallback(self):
|
||||||
|
segs = [EmotionSegment("hello", EMO_HAPPY, 0.4)]
|
||||||
|
timings = [{"start": 0, "end": 1}, {"start": 1, "end": 2}]
|
||||||
|
entries = align_timeline_by_timings(segs, timings, 2.0)
|
||||||
|
assert len(entries) == 1
|
||||||
|
assert entries[0].end == 2.0
|
||||||
|
|
||||||
|
|
||||||
|
# ── DittoEmotionService ──────────────────────────────────────────
|
||||||
|
def _make_service(enabled=True, model=None, temperature=0.1, timeout=10, max_tokens=1024, prompt=""):
|
||||||
|
s = MagicMock()
|
||||||
|
s.ditto_emotion_enabled = enabled
|
||||||
|
s.ditto_emotion_model = model or ""
|
||||||
|
s.ditto_emotion_temperature = temperature
|
||||||
|
s.ditto_emotion_timeout = timeout
|
||||||
|
s.ditto_emotion_max_tokens = max_tokens
|
||||||
|
s.ditto_emotion_cache_size = 100
|
||||||
|
s.ditto_emotion_prompt = prompt
|
||||||
|
return DittoEmotionService(settings=s)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDittoEmotionService:
|
||||||
|
def test_disabled_returns_empty(self):
|
||||||
|
svc = _make_service(enabled=False)
|
||||||
|
assert svc.analyze("你好世界") == []
|
||||||
|
|
||||||
|
def test_empty_text_returns_empty(self):
|
||||||
|
svc = _make_service(enabled=True)
|
||||||
|
assert svc.analyze("") == []
|
||||||
|
|
||||||
|
def test_llm_success(self):
|
||||||
|
svc = _make_service(enabled=True)
|
||||||
|
fake_reply = json.dumps([{"text": "你好", "emo": 4, "intensity": 0.2}])
|
||||||
|
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||||
|
segs = svc.analyze("你好")
|
||||||
|
assert len(segs) == 1
|
||||||
|
assert segs[0].emo == 4
|
||||||
|
|
||||||
|
def test_cache_hit(self):
|
||||||
|
svc = _make_service(enabled=True)
|
||||||
|
fake_reply = json.dumps([{"text": "你好世界", "emo": 3, "intensity": 0.5}])
|
||||||
|
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)) as mock_call:
|
||||||
|
svc.analyze("你好世界")
|
||||||
|
svc.analyze("你好世界")
|
||||||
|
assert mock_call.call_count == 1
|
||||||
|
|
||||||
|
def test_build_timeline_empty_when_disabled(self):
|
||||||
|
svc = _make_service(enabled=False)
|
||||||
|
assert svc.build_timeline("test", 5.0) == ""
|
||||||
|
|
||||||
|
def test_build_timeline_returns_json(self):
|
||||||
|
svc = _make_service(enabled=True)
|
||||||
|
fake_reply = json.dumps(
|
||||||
|
[
|
||||||
|
{"text": "ab", "emo": 4, "intensity": 0.2},
|
||||||
|
{"text": "cd", "emo": 3, "intensity": 0.4},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||||
|
result = svc.build_timeline("ab。cd。", 4.0)
|
||||||
|
data = json.loads(result)
|
||||||
|
assert len(data) == 2
|
||||||
|
assert data[0]["emo"] == 4
|
||||||
|
assert data[1]["emo"] == 3
|
||||||
|
|
||||||
|
def test_build_timeline_with_sentence_timings(self):
|
||||||
|
svc = _make_service(enabled=True)
|
||||||
|
fake_reply = json.dumps(
|
||||||
|
[
|
||||||
|
{"text": "hello", "emo": 3, "intensity": 0.4},
|
||||||
|
{"text": "world", "emo": 4, "intensity": 0.2},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
timings = [
|
||||||
|
{"start": 0.0, "end": 1.0},
|
||||||
|
{"start": 1.0, "end": 3.0},
|
||||||
|
]
|
||||||
|
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||||
|
result = svc.build_timeline("hello world", 3.0, sentence_timings=timings)
|
||||||
|
data = json.loads(result)
|
||||||
|
assert data[0]["start"] == 0.0
|
||||||
|
assert data[0]["end"] == 1.0
|
||||||
|
assert data[1]["end"] == 3.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestPromptLoading:
|
||||||
|
def test_default_prompt_contains_placeholder(self):
|
||||||
|
from packages.application.ditto_emotion_service import _load_default_prompt
|
||||||
|
|
||||||
|
prompt = _load_default_prompt()
|
||||||
|
assert "{文案}" in prompt
|
||||||
|
|
||||||
|
def test_config_prompt_override(self):
|
||||||
|
custom = "分析情绪: {文案}"
|
||||||
|
svc = _make_service(enabled=True, prompt=custom)
|
||||||
|
assert svc._get_prompt_template() == custom
|
||||||
@@ -0,0 +1,229 @@
|
|||||||
|
"""Ditto 蚂蚁数字人客户端单元测试 — #2076."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
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"] == " "
|
||||||
|
|
||||||
|
|
||||||
|
def test_generate_network_error_fails_fast(monkeypatch):
|
||||||
|
"""网络不通(ConnectError)时不重试,直接快速抛 NetworkUnreachable,避免用户等5分钟"""
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from packages.application import ditto_service as ds_mod
|
||||||
|
|
||||||
|
calls = {"n": 0}
|
||||||
|
|
||||||
|
def _fake_post(self, url, json=None):
|
||||||
|
calls["n"] += 1
|
||||||
|
raise httpx.ConnectError("[Errno 113] No route to host")
|
||||||
|
|
||||||
|
monkeypatch.setattr(httpx.Client, "post", _fake_post)
|
||||||
|
|
||||||
|
client = _make_client(
|
||||||
|
base_url="http://100.76.80.23:8000",
|
||||||
|
default_video_url="http://oss/tpl.mp4",
|
||||||
|
max_retries=2,
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
|
||||||
|
t0 = time.monotonic()
|
||||||
|
with pytest.raises(ds_mod.DittoError) as exc:
|
||||||
|
client.generate(audio_url="http://oss/a.wav", script="你好")
|
||||||
|
elapsed = time.monotonic() - t0
|
||||||
|
|
||||||
|
assert exc.value.code == "NetworkUnreachable"
|
||||||
|
assert calls["n"] == 1 # 不重试
|
||||||
|
assert elapsed < 5 # 快速失败<5秒
|
||||||
@@ -350,6 +350,37 @@ class TestViralVideoRepository:
|
|||||||
class TestViralVideoPipeline:
|
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-10秒">
|
||||||
|
<voiceover>颜色特别好看</voiceover>
|
||||||
|
<visual>特写,固定镜头</visual>
|
||||||
|
<action_details>嘴唇涂抹特写</action_details>
|
||||||
|
<audio_bgm>轻快BGM继续</audio_bgm>
|
||||||
|
<transition>硬切</transition>
|
||||||
|
<reference_image_index>1</reference_image_index>
|
||||||
|
</clip>
|
||||||
|
<clip image_index="2" time_range="10-15秒">
|
||||||
|
<voiceover>很显白,推荐给大家</voiceover>
|
||||||
|
<visual>中景,微笑展示</visual>
|
||||||
|
<action_details>口红展示</action_details>
|
||||||
|
<audio_bgm>轻快BGM结束</audio_bgm>
|
||||||
|
<transition>结束</transition>
|
||||||
|
<reference_image_index>0</reference_image_index>
|
||||||
|
</clip>
|
||||||
|
</clips>
|
||||||
|
<voiceover_script>大家好呀,今天来给大家分享一款超显白的口红。颜色特别好看很显气质,真心推荐给姐妹们</voiceover_script>
|
||||||
|
<theme>口红分享</theme>"""
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_job(self):
|
def mock_job(self):
|
||||||
return ViralVideoJob(
|
return ViralVideoJob(
|
||||||
@@ -364,23 +395,27 @@ class TestViralVideoPipeline:
|
|||||||
video_ratio="9:16",
|
video_ratio="9:16",
|
||||||
)
|
)
|
||||||
|
|
||||||
@patch("packages.shared.ai_service.call_vision")
|
@patch("apps.worker.worker_app.tasks.vision.analyze_images_v2")
|
||||||
def test_image_analysis_step(self, mock_vision, mock_job):
|
def test_image_analysis_step(self, mock_vision, mock_job):
|
||||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||||
|
|
||||||
mock_vision.return_value = {"name": "口红", "features": ["持久", "滋润"]}
|
# 每张图返回一个 v8 5 字段结果
|
||||||
|
mock_vision.return_value = [
|
||||||
|
{"type": "product", "name": "口红", "brand": "", "has_person": False, "summary_markdown": "一支口红"},
|
||||||
|
{"type": "product", "name": "口红", "brand": "", "has_person": False, "summary_markdown": "口红特写"},
|
||||||
|
]
|
||||||
result = _step_image_analysis(mock_job)
|
result = _step_image_analysis(mock_job)
|
||||||
assert "products" in result
|
assert "images" in result
|
||||||
assert len(result["products"]) == 2 # 两张图片
|
assert len(result["images"]) == 2 # 两张图片
|
||||||
|
|
||||||
@patch("packages.shared.ai_service.call_vision")
|
@patch("apps.worker.worker_app.tasks.vision.analyze_images_v2")
|
||||||
def test_image_analysis_fallback(self, mock_vision, mock_job):
|
def test_image_analysis_fallback(self, mock_vision, mock_job):
|
||||||
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
|
||||||
|
|
||||||
# 模拟 call_vision 不存在
|
# v2 分析内部异常时,每图走兜底,仍返回 images 结构
|
||||||
mock_vision.side_effect = ImportError("no module")
|
mock_vision.side_effect = RuntimeError("vision unavailable")
|
||||||
result = _step_image_analysis(mock_job)
|
result = _step_image_analysis(mock_job)
|
||||||
assert "products" in result
|
assert "images" in result
|
||||||
|
|
||||||
def test_video_analysis_no_reference(self, mock_job):
|
def test_video_analysis_no_reference(self, mock_job):
|
||||||
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
|
||||||
@@ -390,51 +425,45 @@ class TestViralVideoPipeline:
|
|||||||
result = _step_video_analysis(mock_job)
|
result = _step_video_analysis(mock_job)
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
@patch("packages.shared.ai_service.call_llm")
|
def test_intent_parsing_step_removed(self, mock_job):
|
||||||
def test_intent_parsing(self, mock_llm, mock_job):
|
"""intent_parsing 已合并进脚本生成,不再作为独立步骤/函数存在。"""
|
||||||
from apps.worker.worker_app.tasks.viral_video import _step_intent_parsing
|
import apps.worker.worker_app.tasks.viral_video as vv
|
||||||
|
|
||||||
mock_llm.return_value = {"intent": "推广口红", "tone": "活泼"}
|
assert not hasattr(vv, "_step_intent_parsing")
|
||||||
result = _step_intent_parsing(mock_job, {"products": []})
|
|
||||||
assert "intent" in result
|
|
||||||
|
|
||||||
@patch("packages.shared.ai_service.call_llm")
|
def test_script_generation_returns_copy_result(self, mock_job):
|
||||||
def test_script_generation_returns_copy_result(self, mock_llm, mock_job):
|
|
||||||
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
|
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
|
||||||
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
|
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
|
||||||
|
from packages.shared.ai_router import ai_router
|
||||||
|
|
||||||
mock_llm.return_value = """<clips>
|
class _FakeClient:
|
||||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快流行BGM">
|
is_available = True
|
||||||
<voice_text>大家好,今天分享一款口红</voice_text>
|
model = "fake-storyboard"
|
||||||
<subtitle_text>大家好,今天分享一款口红</subtitle_text>
|
|
||||||
<shot_type_angle_movement>近景平视,缓慢推镜</shot_type_angle_movement>
|
def __init__(self, xml: str):
|
||||||
<scene_and_dialogue>女主微笑展示口红:大家好,今天分享一款口红</scene_and_dialogue>
|
self._xml = xml
|
||||||
<action_details>手持口红特写</action_details>
|
|
||||||
<audio_bgm>轻快流行BGM</audio_bgm>
|
def chat_completion(self, messages, **kwargs):
|
||||||
<transition>硬切</transition>
|
return self._xml
|
||||||
<reference_image_index>0</reference_image_index>
|
|
||||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
fake = _FakeClient(self.V3_XML)
|
||||||
</clip>
|
orig_get = ai_router.get_llm_client
|
||||||
<clip image_index="0" transition="fade" zoom="null" duration_sec="10" bgm_note="轻快BGM">
|
|
||||||
<voice_text>颜色特别好看很显白</voice_text>
|
def _get(task, variant="primary"):
|
||||||
<subtitle_text>颜色特别好看很显白</subtitle_text>
|
if task == "storyboard":
|
||||||
<shot_type_angle_movement>特写,固定镜头</shot_type_angle_movement>
|
return fake
|
||||||
<scene_and_dialogue>涂抹口红:颜色特别好看很显白</scene_and_dialogue>
|
return orig_get(task, variant=variant)
|
||||||
<action_details>嘴唇涂抹特写</action_details>
|
|
||||||
<audio_bgm>轻快BGM继续</audio_bgm>
|
ai_router.get_llm_client = _get # type: ignore
|
||||||
<transition>结束</transition>
|
try:
|
||||||
<reference_image_index>1</reference_image_index>
|
result = _step_script_generation(mock_job, {"images": []})
|
||||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
finally:
|
||||||
</clip>
|
ai_router.get_llm_client = orig_get # type: ignore
|
||||||
</clips>"""
|
|
||||||
result = _step_script_generation(
|
|
||||||
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
|
|
||||||
)
|
|
||||||
assert isinstance(result, dict)
|
assert isinstance(result, dict)
|
||||||
assert "voiceover_script" in result
|
assert "voiceover_script" in result
|
||||||
assert "shots" in result
|
assert "shots" in result
|
||||||
assert isinstance(result["shots"], list)
|
assert isinstance(result["shots"], list)
|
||||||
assert len(result["shots"]) == 2
|
assert len(result["shots"]) == 3
|
||||||
assert result["overview"]["total_duration"] == 15
|
assert result["overview"]["total_duration"] == 15
|
||||||
# final_copy 必须 = voiceover_script(向后兼容)
|
# final_copy 必须 = voiceover_script(向后兼容)
|
||||||
assert result.get("final_copy") == result["voiceover_script"]
|
assert result.get("final_copy") == result["voiceover_script"]
|
||||||
@@ -483,7 +512,7 @@ class TestViralVideoPipeline:
|
|||||||
}
|
}
|
||||||
prompt = _assemble_seedance_prompt(cr, mock_job)
|
prompt = _assemble_seedance_prompt(cr, mock_job)
|
||||||
assert "【视频总览】" in prompt
|
assert "【视频总览】" in prompt
|
||||||
assert "【逐镜头时间轴】" in prompt
|
assert "【分镜脚本】" in prompt
|
||||||
assert "【硬性约束】" in prompt
|
assert "【硬性约束】" in prompt
|
||||||
assert "【负面提示词】" in prompt
|
assert "【负面提示词】" in prompt
|
||||||
assert "0-15秒" in prompt
|
assert "0-15秒" in prompt
|
||||||
@@ -501,7 +530,6 @@ class TestPipelineIntegration:
|
|||||||
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
|
@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_review")
|
||||||
@patch("apps.worker.worker_app.tasks.viral_video._step_script_generation")
|
@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_video_analysis")
|
||||||
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
|
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
|
||||||
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
|
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
|
||||||
@@ -512,7 +540,6 @@ class TestPipelineIntegration:
|
|||||||
mock_get_repo,
|
mock_get_repo,
|
||||||
mock_img_analysis,
|
mock_img_analysis,
|
||||||
mock_video_analysis,
|
mock_video_analysis,
|
||||||
mock_intent,
|
|
||||||
mock_script,
|
mock_script,
|
||||||
mock_review,
|
mock_review,
|
||||||
mock_tts,
|
mock_tts,
|
||||||
@@ -539,8 +566,6 @@ class TestPipelineIntegration:
|
|||||||
mock_session = MagicMock()
|
mock_session = MagicMock()
|
||||||
mock_get_repo.return_value = (mock_session, mock_repo, job)
|
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 = {
|
mock_script.return_value = {
|
||||||
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
|
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
|
||||||
"scene_and_lighting": "明亮化妆台",
|
"scene_and_lighting": "明亮化妆台",
|
||||||
@@ -553,6 +578,8 @@ class TestPipelineIntegration:
|
|||||||
mock_review.return_value = {"passed": True, "score": 90}
|
mock_review.return_value = {"passed": True, "score": 90}
|
||||||
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
|
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
|
||||||
mock_tts_upload.return_value = None
|
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_render.return_value = ("/tmp/video.mp4", {"completion_tokens": 1000000})
|
||||||
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
||||||
|
|
||||||
|
|||||||
@@ -129,7 +129,7 @@ class TestScriptGenerationV16:
|
|||||||
"negative_prompts": ["水印"],
|
"negative_prompts": ["水印"],
|
||||||
}
|
}
|
||||||
p = _assemble_seedance_prompt(cr, mock_job)
|
p = _assemble_seedance_prompt(cr, mock_job)
|
||||||
for key in ("【视频总览】", "【场景与光线】", "【逐镜头时间轴】", "【硬性约束】", "【负面提示词】"):
|
for key in ("【视频总览】", "【参考素材】", "【分镜脚本】", "【硬性约束】", "【负面提示词】"):
|
||||||
assert key in p
|
assert key in p
|
||||||
|
|
||||||
|
|
||||||
@@ -308,8 +308,7 @@ class TestResumeReadsImageAnalysis:
|
|||||||
|
|
||||||
# v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis
|
# v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis
|
||||||
src = inspect.getsource(vv._run_render_pipeline)
|
src = inspect.getsource(vv._run_render_pipeline)
|
||||||
assert "job.image_analysis" in src
|
assert "job.copy_result" in src
|
||||||
assert "image_analysis" in src
|
|
||||||
# resume 本身应该调用 _run_render_pipeline
|
# resume 本身应该调用 _run_render_pipeline
|
||||||
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
|
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||||
assert "_run_render_pipeline" in resume_src
|
assert "_run_render_pipeline" in resume_src
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
"""#2040 爆款视频 Prompt 模板系统单测。
|
"""#2040 爆款视频 Prompt 模板系统单测(v8/v3 叙述优先重构后)。
|
||||||
|
|
||||||
不真调豆包 API,全部用 FakeClient 注入;覆盖:
|
不真调豆包 API,全部用 FakeClient 注入;覆盖:
|
||||||
XML 标签解析 / 5 套模板纯文本 / loader 缓存热加载与回落 /
|
XML 标签解析 / 3 套模板纯文本(image_analysis/storyboard/review)/
|
||||||
三档融合差异 / personal_brands 保留 / 审核识别违规词夸大 / 自动重写 /
|
loader 缓存热加载与回落 / 本地规则审核识别违规词夸大 /
|
||||||
各步 fallback / seed 幂等 / 负面词不出现。
|
各现存步 fallback / seed 幂等 / 负面词不出现。
|
||||||
|
|
||||||
|
注:intent_parsing、copy_fusion 两套模板及其独立步骤已在叙述优先重构中删除,
|
||||||
|
相关用例同步移除。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -30,7 +33,6 @@ from packages.application.viral_video.prompt_loader import ( # noqa: E402
|
|||||||
from packages.application.viral_video.prompts import ( # noqa: E402
|
from packages.application.viral_video.prompts import ( # noqa: E402
|
||||||
BANNED_PHRASES,
|
BANNED_PHRASES,
|
||||||
DEFAULT_TEMPLATES,
|
DEFAULT_TEMPLATES,
|
||||||
FUSION_INSTRUCTIONS,
|
|
||||||
)
|
)
|
||||||
from packages.application.viral_video.reviewer import Reviewer # noqa: E402
|
from packages.application.viral_video.reviewer import Reviewer # noqa: E402
|
||||||
|
|
||||||
@@ -45,26 +47,6 @@ IMAGE_XML = """<products>
|
|||||||
<quality resolution="高清" lighting="柔和" composition="居中"/>
|
<quality resolution="高清" lighting="柔和" composition="居中"/>
|
||||||
<key_selling_points><point>去油快</point><point>625ml大容量</point></key_selling_points>"""
|
<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>
|
STORYBOARD_XML = """<clips>
|
||||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常">
|
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常">
|
||||||
<voice_text>这油污忍很久了</voice_text>
|
<voice_text>这油污忍很久了</voice_text>
|
||||||
@@ -86,8 +68,6 @@ REVIEW_PASS_XML = """<passed>true</passed>
|
|||||||
<issues></issues>
|
<issues></issues>
|
||||||
<rewrite_suggestions></rewrite_suggestions>"""
|
<rewrite_suggestions></rewrite_suggestions>"""
|
||||||
|
|
||||||
FIXED_FUSION_XML = FUSION_XML.replace("一擦就净", "大部分油污能擦掉")
|
|
||||||
|
|
||||||
|
|
||||||
class FakeClient:
|
class FakeClient:
|
||||||
"""按 system 内容路由 canned 响应的假豆包客户端。"""
|
"""按 system 内容路由 canned 响应的假豆包客户端。"""
|
||||||
@@ -96,45 +76,20 @@ class FakeClient:
|
|||||||
self.chat_calls: list[list[dict]] = []
|
self.chat_calls: list[list[dict]] = []
|
||||||
self.vision_calls: list = []
|
self.vision_calls: list = []
|
||||||
self.review_sequence: list[str] | None = None
|
self.review_sequence: list[str] | None = None
|
||||||
self.rewrite_response: str = FIXED_FUSION_XML
|
|
||||||
|
|
||||||
def chat_completion(self, messages, **kwargs):
|
def chat_completion(self, messages, **kwargs):
|
||||||
self.chat_calls.append(messages)
|
self.chat_calls.append(messages)
|
||||||
system = messages[0]["content"]
|
system = messages[0]["content"]
|
||||||
user = messages[1]["content"]
|
# v3 审核 prompt 关键短语(叙述优先重构后更新)
|
||||||
if "按审核意见修正文案" in system:
|
if "短视频广告合规审核与文案优化专家" in system:
|
||||||
return self.rewrite_response
|
|
||||||
if "文案合规审核员" in system:
|
|
||||||
if self.review_sequence:
|
if self.review_sequence:
|
||||||
return self.review_sequence.pop(0)
|
return self.review_sequence.pop(0)
|
||||||
return REVIEW_PASS_XML
|
return REVIEW_PASS_XML
|
||||||
if "理解用户的营销意图" in system:
|
# v3 分镜 prompt
|
||||||
return INTENT_XML
|
if "懂短视频的编导和口播文案高手" in system:
|
||||||
if "负责把文案拆成可拍摄" in system:
|
|
||||||
return STORYBOARD_XML
|
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 ""
|
return ""
|
||||||
|
|
||||||
_fusion_override = None
|
|
||||||
_last_mode = None
|
|
||||||
|
|
||||||
def vision_completion(self, messages, images=None, **kwargs):
|
def vision_completion(self, messages, images=None, **kwargs):
|
||||||
self.vision_calls.append({"messages": messages, "images": images})
|
self.vision_calls.append({"messages": messages, "images": images})
|
||||||
return IMAGE_XML
|
return IMAGE_XML
|
||||||
@@ -171,25 +126,25 @@ class TestXmlParser:
|
|||||||
assert xp.text_of("乱七八糟没有标签", "intent", "默认") == "默认"
|
assert xp.text_of("乱七八糟没有标签", "intent", "默认") == "默认"
|
||||||
|
|
||||||
|
|
||||||
# ── 5 套模板纯文本 ────────────────────────────────────────────────────────
|
# ── 3 套模板纯文本 ────────────────────────────────────────────────────────
|
||||||
class TestTemplates:
|
class TestTemplates:
|
||||||
def test_five_templates_present(self):
|
def test_three_templates_present(self):
|
||||||
types_ = {t["prompt_type"] for t in DEFAULT_TEMPLATES}
|
types_ = {t["prompt_type"] for t in DEFAULT_TEMPLATES}
|
||||||
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
assert types_ == {"image_analysis", "storyboard", "review"}
|
||||||
|
|
||||||
def test_no_json_blocks_in_templates(self):
|
def test_no_json_blocks_in_templates(self):
|
||||||
for template in DEFAULT_TEMPLATES:
|
for template in DEFAULT_TEMPLATES:
|
||||||
blob = "\n".join([template["system_prompt"], template["user_prompt_template"], template["example_output"]])
|
blob = "\n".join([template["system_prompt"], template["user_prompt_template"], template["example_output"]])
|
||||||
assert "```json" not in blob
|
# image_analysis 模板明确要求输出 JSON,故只对非 image_analysis 模板校验
|
||||||
assert "JSON schema" not in blob
|
if template["prompt_type"] != "image_analysis":
|
||||||
|
assert "```json" not in blob
|
||||||
|
assert "JSON schema" not in blob
|
||||||
|
|
||||||
def test_placeholders_render_and_missing_key_kept(self):
|
def test_placeholders_render_and_missing_key_kept(self):
|
||||||
template = get_template("intent_parsing")
|
# 现存模板里选取 storyboard 做占位符渲染校验
|
||||||
rendered = render_user_prompt(template, user_copy_text="去油快", industry="家居")
|
template = get_template("storyboard")
|
||||||
|
rendered = render_user_prompt(template, marketing_purpose="去油快", industry="家居")
|
||||||
assert "去油快" in rendered
|
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 加载/缓存/回落 ─────────────────────────────────────────────
|
# ── loader:DB 加载/缓存/回落 ─────────────────────────────────────────────
|
||||||
@@ -200,7 +155,8 @@ class TestPromptLoader:
|
|||||||
monkeypatch.setattr(session_mod, "SessionLocal", None, raising=False)
|
monkeypatch.setattr(session_mod, "SessionLocal", None, raising=False)
|
||||||
template = get_template("review")
|
template = get_template("review")
|
||||||
assert template is not None
|
assert template is not None
|
||||||
assert "6个维度" in template.system_prompt
|
# v3 审核 prompt 实际内容断言
|
||||||
|
assert "合规审核" in template.system_prompt
|
||||||
|
|
||||||
def test_db_row_takes_precedence(self, tmp_path, monkeypatch):
|
def test_db_row_takes_precedence(self, tmp_path, monkeypatch):
|
||||||
import packages.adapters.sqlalchemy_impl.session as session_mod
|
import packages.adapters.sqlalchemy_impl.session as session_mod
|
||||||
@@ -240,50 +196,8 @@ class TestPromptLoader:
|
|||||||
get_template("not_exist")
|
get_template("not_exist")
|
||||||
|
|
||||||
|
|
||||||
# ── 5 步编排与 fallback ──────────────────────────────────────────────────
|
# ── 现存步编排与 fallback ────────────────────────────────────────────────
|
||||||
class TestGenerator:
|
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):
|
def test_image_fallback_on_garbage(self):
|
||||||
client = FakeClient()
|
client = FakeClient()
|
||||||
client.vision_completion = lambda *a, **k: "完全无法解析的内容" # type: ignore
|
client.vision_completion = lambda *a, **k: "完全无法解析的内容" # type: ignore
|
||||||
@@ -291,29 +205,6 @@ class TestGenerator:
|
|||||||
analysis = gen.analyze_images(["https://x/1.jpg"])
|
analysis = gen.analyze_images(["https://x/1.jpg"])
|
||||||
assert analysis.products[0].name.startswith("无法判断")
|
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):
|
def test_storyboard_fallback_on_garbage(self):
|
||||||
client = FakeClient()
|
client = FakeClient()
|
||||||
client.chat_completion = lambda *a, **k: "啥都没有" # type: ignore
|
client.chat_completion = lambda *a, **k: "啥都没有" # type: ignore
|
||||||
@@ -329,7 +220,7 @@ class TestGenerator:
|
|||||||
assert board.clips[0].voice_text == "a"
|
assert board.clips[0].voice_text == "a"
|
||||||
|
|
||||||
|
|
||||||
# ── 审核与自动重写 ────────────────────────────────────────────────────────
|
# ── 审核本地规则(LLM 降级放行时本地规则仍应识别红线)────────────────────
|
||||||
class TestReview:
|
class TestReview:
|
||||||
def test_rule_check_catches_exaggeration_even_if_llm_passes(self):
|
def test_rule_check_catches_exaggeration_even_if_llm_passes(self):
|
||||||
client = FakeClient() # LLM 默认返回 passed
|
client = FakeClient() # LLM 默认返回 passed
|
||||||
@@ -338,6 +229,7 @@ class TestReview:
|
|||||||
|
|
||||||
fusion = FusionResult(title="一喷100%掉光", hook="x", cta="买")
|
fusion = FusionResult(title="一喷100%掉光", hook="x", cta="买")
|
||||||
result = reviewer.review(fusion, IntentResult(), "ai_full")
|
result = reviewer.review(fusion, IntentResult(), "ai_full")
|
||||||
|
# LLM 返回 passed,且本地规则命中夸大 → 整体不通过
|
||||||
assert result.passed is False
|
assert result.passed is False
|
||||||
dims = {i.dimension for i in result.issues}
|
dims = {i.dimension for i in result.issues}
|
||||||
assert "夸大承诺" in dims
|
assert "夸大承诺" in dims
|
||||||
@@ -381,18 +273,6 @@ class TestReview:
|
|||||||
result = reviewer.review(fusion, intent, "user_primary")
|
result = reviewer.review(fusion, intent, "user_primary")
|
||||||
assert any(i.dimension == "用户意图保留" for i in result.issues)
|
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):
|
def test_rule_fix_local(self):
|
||||||
reviewer = Reviewer(client=FakeClient())
|
reviewer = Reviewer(client=FakeClient())
|
||||||
from packages.application.viral_video.schemas import (
|
from packages.application.viral_video.schemas import (
|
||||||
@@ -442,18 +322,17 @@ class TestSeed:
|
|||||||
"UNIQUE(prompt_type, version))"
|
"UNIQUE(prompt_type, version))"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
assert seed_mod.seed(engine) == 5
|
assert seed_mod.seed(engine) == 3
|
||||||
assert seed_mod.seed(engine) == 5 # 再来一次不报错
|
assert seed_mod.seed(engine) == 3 # 再来一次不报错
|
||||||
with engine.begin() as conn:
|
with engine.begin() as conn:
|
||||||
count = conn.execute(sa.text("SELECT COUNT(*) FROM viral_video_prompt_templates")).scalar()
|
count = conn.execute(sa.text("SELECT COUNT(*) FROM viral_video_prompt_templates")).scalar()
|
||||||
assert count == 5
|
assert count == 3
|
||||||
active_types = conn.execute # noqa: B018
|
|
||||||
with engine.begin() as conn:
|
with engine.begin() as conn:
|
||||||
types_ = {
|
types_ = {
|
||||||
r[0]
|
r[0]
|
||||||
for r in conn.execute(sa.text("SELECT prompt_type FROM viral_video_prompt_templates WHERE is_active=1"))
|
for r in conn.execute(sa.text("SELECT prompt_type FROM viral_video_prompt_templates WHERE is_active=1"))
|
||||||
}
|
}
|
||||||
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
assert types_ == {"image_analysis", "storyboard", "review"}
|
||||||
|
|
||||||
|
|
||||||
# ── 负面词不出现于程序产出 ────────────────────────────────────────────────
|
# ── 负面词不出现于程序产出 ────────────────────────────────────────────────
|
||||||
|
|||||||
@@ -29,6 +29,15 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
|
|||||||
job.user_id = user_id
|
job.user_id = user_id
|
||||||
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||||
job.images = kwargs.pop("images", ["img-1"])
|
job.images = kwargs.pop("images", ["img-1"])
|
||||||
|
# confirm_copy 新增 copy_result 完整性校验:默认提供合法文案数据
|
||||||
|
job.copy_result = kwargs.pop(
|
||||||
|
"copy_result",
|
||||||
|
{
|
||||||
|
"theme": "测试主题",
|
||||||
|
"voiceover_script": "这是一段测试口播文案内容。",
|
||||||
|
"shots": [{"time_range": "0-15秒", "voiceover": "这是一段测试口播文案内容。"}],
|
||||||
|
},
|
||||||
|
)
|
||||||
job.industry = kwargs.pop("industry", "电商")
|
job.industry = kwargs.pop("industry", "电商")
|
||||||
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
||||||
for k, v in {
|
for k, v in {
|
||||||
@@ -55,7 +64,6 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
|
|||||||
"intent_result": None,
|
"intent_result": None,
|
||||||
"image_analysis": None,
|
"image_analysis": None,
|
||||||
"storyboard": None,
|
"storyboard": None,
|
||||||
"copy_result": None,
|
|
||||||
"generated_copy_text": "",
|
"generated_copy_text": "",
|
||||||
"voice_id": "",
|
"voice_id": "",
|
||||||
"voice_source": "",
|
"voice_source": "",
|
||||||
@@ -147,9 +155,16 @@ class TestRetryViralVideo:
|
|||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
job = _make_job(
|
job = _make_job(
|
||||||
job_id="job-retry2", user_id="u1", status=ViralVideoStatus.FAILED,
|
job_id="job-retry2",
|
||||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
user_id="u1",
|
||||||
credits_prepaid=5.0, credits_transaction_id="txn1", retry_count=0,
|
status=ViralVideoStatus.FAILED,
|
||||||
|
duration=15,
|
||||||
|
video_ratio="9:16",
|
||||||
|
video_resolution="720p",
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
credits_prepaid=5.0,
|
||||||
|
credits_transaction_id="txn1",
|
||||||
|
retry_count=0,
|
||||||
)
|
)
|
||||||
repo = MagicMock()
|
repo = MagicMock()
|
||||||
repo.get.return_value = job
|
repo.get.return_value = job
|
||||||
@@ -180,24 +195,32 @@ class TestRetryViralVideo:
|
|||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
job = _make_job(
|
job = _make_job(
|
||||||
job_id="job-retry3a", user_id="u1", status=ViralVideoStatus.FAILED,
|
job_id="job-retry3a",
|
||||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
user_id="u1",
|
||||||
credits_prepaid=5.0, credits_transaction_id="txn-old", retry_count=0,
|
status=ViralVideoStatus.FAILED,
|
||||||
|
duration=15,
|
||||||
|
video_ratio="9:16",
|
||||||
|
video_resolution="720p",
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
credits_prepaid=5.0,
|
||||||
|
credits_transaction_id="txn-old",
|
||||||
|
retry_count=0,
|
||||||
)
|
)
|
||||||
repo = MagicMock()
|
repo = MagicMock()
|
||||||
repo.get.return_value = job
|
repo.get.return_value = job
|
||||||
|
|
||||||
fake_svc = MagicMock()
|
fake_svc = MagicMock()
|
||||||
fake_svc.deduct_viral_video.return_value = {"success": False, "balance": 1.0}
|
fake_svc.deduct_viral_video.return_value = {"success": False, "balance": 1.0}
|
||||||
req = RetryViralVideoRequest(duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5")
|
req = RetryViralVideoRequest(
|
||||||
|
duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5"
|
||||||
|
)
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||||
patch("app.config.settings") as mock_settings,
|
patch("app.config.settings") as mock_settings,
|
||||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
||||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(15.0, {})),
|
||||||
return_value=(15.0, {})),
|
|
||||||
patch.object(vv_mod.celery_app, "send_task"),
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
):
|
):
|
||||||
mock_settings.points_enabled = True
|
mock_settings.points_enabled = True
|
||||||
@@ -216,19 +239,29 @@ class TestRetryViralVideo:
|
|||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
job = _make_job(
|
job = _make_job(
|
||||||
job_id="job-retry3b", user_id="u1", status=ViralVideoStatus.FAILED,
|
job_id="job-retry3b",
|
||||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
user_id="u1",
|
||||||
credits_prepaid=5.0, credits_transaction_id="txn-old", retry_count=0,
|
status=ViralVideoStatus.FAILED,
|
||||||
|
duration=15,
|
||||||
|
video_ratio="9:16",
|
||||||
|
video_resolution="720p",
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
credits_prepaid=5.0,
|
||||||
|
credits_transaction_id="txn-old",
|
||||||
|
retry_count=0,
|
||||||
)
|
)
|
||||||
# 用 SimpleNamespace 让属性真正可写
|
# 用 SimpleNamespace 让属性真正可写
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
job.credits_prepaid = 5.0
|
job.credits_prepaid = 5.0
|
||||||
repo = MagicMock()
|
repo = MagicMock()
|
||||||
repo.get.return_value = job
|
repo.get.return_value = job
|
||||||
|
|
||||||
fake_svc = MagicMock()
|
fake_svc = MagicMock()
|
||||||
fake_svc.deduct_viral_video.return_value = {"success": True, "balance": 50.0, "transaction_id": "txn-new"}
|
fake_svc.deduct_viral_video.return_value = {"success": True, "balance": 50.0, "transaction_id": "txn-new"}
|
||||||
req = RetryViralVideoRequest(duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5")
|
req = RetryViralVideoRequest(
|
||||||
|
duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5"
|
||||||
|
)
|
||||||
|
|
||||||
new_est = 15.0
|
new_est = 15.0
|
||||||
with (
|
with (
|
||||||
@@ -236,8 +269,9 @@ class TestRetryViralVideo:
|
|||||||
patch("app.config.settings") as mock_settings,
|
patch("app.config.settings") as mock_settings,
|
||||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
||||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
patch(
|
||||||
return_value=(new_est, {})),
|
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(new_est, {})
|
||||||
|
),
|
||||||
patch.object(vv_mod.celery_app, "send_task"),
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
):
|
):
|
||||||
mock_settings.points_enabled = True
|
mock_settings.points_enabled = True
|
||||||
@@ -262,9 +296,16 @@ class TestRetryViralVideo:
|
|||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
job = _make_job(
|
job = _make_job(
|
||||||
job_id="job-retry4", user_id="u1", status=ViralVideoStatus.FAILED,
|
job_id="job-retry4",
|
||||||
duration=20, video_ratio="16:9", video_resolution="1080p", video_model="seedance-2.5",
|
user_id="u1",
|
||||||
credits_prepaid=10.0, credits_transaction_id="txn-old", retry_count=0,
|
status=ViralVideoStatus.FAILED,
|
||||||
|
duration=20,
|
||||||
|
video_ratio="16:9",
|
||||||
|
video_resolution="1080p",
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
credits_prepaid=10.0,
|
||||||
|
credits_transaction_id="txn-old",
|
||||||
|
retry_count=0,
|
||||||
)
|
)
|
||||||
job.credits_prepaid = 10.0
|
job.credits_prepaid = 10.0
|
||||||
repo = MagicMock()
|
repo = MagicMock()
|
||||||
@@ -280,8 +321,9 @@ class TestRetryViralVideo:
|
|||||||
patch("app.config.settings") as mock_settings,
|
patch("app.config.settings") as mock_settings,
|
||||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(270, 480)),
|
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(270, 480)),
|
||||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
patch(
|
||||||
return_value=(new_est, {})),
|
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(new_est, {})
|
||||||
|
),
|
||||||
patch.object(vv_mod.celery_app, "send_task"),
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
):
|
):
|
||||||
mock_settings.points_enabled = True
|
mock_settings.points_enabled = True
|
||||||
@@ -470,7 +512,9 @@ class TestGenerateCopy:
|
|||||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
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)
|
resp = vv_mod.generate_copy(
|
||||||
|
f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session
|
||||||
|
)
|
||||||
mock_send.assert_called_once()
|
mock_send.assert_called_once()
|
||||||
job.resume_from_image_analyzed.assert_called()
|
job.resume_from_image_analyzed.assert_called()
|
||||||
assert job.retry_count >= 1
|
assert job.retry_count >= 1
|
||||||
@@ -630,7 +674,7 @@ class TestConfirmCopyPointsDeduction:
|
|||||||
mock_svc.deduct_viral_video.assert_called_once()
|
mock_svc.deduct_viral_video.assert_called_once()
|
||||||
call_args = mock_svc.deduct_viral_video.call_args
|
call_args = mock_svc.deduct_viral_video.call_args
|
||||||
assert call_args.args[0] == "u1" # user_id
|
assert call_args.args[0] == "u1" # user_id
|
||||||
assert call_args.args[1] == 5.2 # credits
|
assert call_args.args[1] == 5.2 # credits
|
||||||
assert call_args.args[2] == "job-pay" # job_id
|
assert call_args.args[2] == "job-pay" # job_id
|
||||||
# credits_prepaid / credits_transaction_id 被写入
|
# credits_prepaid / credits_transaction_id 被写入
|
||||||
assert job.credits_prepaid == 5.2
|
assert job.credits_prepaid == 5.2
|
||||||
@@ -807,9 +851,14 @@ class TestEstimateCredits:
|
|||||||
req = EstimateCreditsRequest(model="seedance-2.5", resolution="1080p", ratio="16:9", duration=20)
|
req = EstimateCreditsRequest(model="seedance-2.5", resolution="1080p", ratio="16:9", duration=20)
|
||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
fake_bd = {
|
fake_bd = {
|
||||||
"tokens": 1000.0, "video_cost": 1.0, "fixed_cost": 0.15,
|
"tokens": 1000.0,
|
||||||
"profit_multiplier": 1.3, "model_price": 70.0,
|
"video_cost": 1.0,
|
||||||
"width": 1920, "height": 1080, "fps": 24,
|
"fixed_cost": 0.15,
|
||||||
|
"profit_multiplier": 1.3,
|
||||||
|
"model_price": 70.0,
|
||||||
|
"width": 1920,
|
||||||
|
"height": 1080,
|
||||||
|
"fps": 24,
|
||||||
}
|
}
|
||||||
|
|
||||||
with (
|
with (
|
||||||
@@ -841,9 +890,14 @@ class TestEstimateCredits:
|
|||||||
req = EstimateCreditsRequest(model="", resolution="720p", ratio="9:16", duration=10)
|
req = EstimateCreditsRequest(model="", resolution="720p", ratio="9:16", duration=10)
|
||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
fake_bd = {
|
fake_bd = {
|
||||||
"tokens": 500.0, "video_cost": 0.5, "fixed_cost": 0.15,
|
"tokens": 500.0,
|
||||||
"profit_multiplier": 1.3, "model_price": 70.0,
|
"video_cost": 0.5,
|
||||||
"width": 720, "height": 1280, "fps": 24,
|
"fixed_cost": 0.15,
|
||||||
|
"profit_multiplier": 1.3,
|
||||||
|
"model_price": 70.0,
|
||||||
|
"width": 720,
|
||||||
|
"height": 1280,
|
||||||
|
"fps": 24,
|
||||||
}
|
}
|
||||||
|
|
||||||
with (
|
with (
|
||||||
@@ -870,9 +924,14 @@ class TestEstimateCredits:
|
|||||||
)
|
)
|
||||||
user = _auth_user("u1")
|
user = _auth_user("u1")
|
||||||
fake_bd = {
|
fake_bd = {
|
||||||
"tokens": 100.0, "video_cost": 0.1, "fixed_cost": 0.15,
|
"tokens": 100.0,
|
||||||
"profit_multiplier": 1.3, "model_price": 46.0,
|
"video_cost": 0.1,
|
||||||
"width": 480, "height": 480, "fps": 24,
|
"fixed_cost": 0.15,
|
||||||
|
"profit_multiplier": 1.3,
|
||||||
|
"model_price": 46.0,
|
||||||
|
"width": 480,
|
||||||
|
"height": 480,
|
||||||
|
"fps": 24,
|
||||||
}
|
}
|
||||||
|
|
||||||
with (
|
with (
|
||||||
|
|||||||
Executable
+909
@@ -0,0 +1,909 @@
|
|||||||
|
"""v8 上线后 4 个 Bug 修复的单元测试。
|
||||||
|
|
||||||
|
Bug1: ConfirmCopyRequest 补字段 + confirm_copy 路由补参数+积分逻辑
|
||||||
|
Bug2: storyboard prompt 口播字数硬限 + 后校验
|
||||||
|
Bug3: TTS 音频时长校验(ffprobe+截断/加速)
|
||||||
|
Bug4: 错误事件按阶段区分(_mark_failed_and_notify 传正确 stage)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, call, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
def _auth_user(uid: str = "u1"):
|
||||||
|
return SimpleNamespace(user=SimpleNamespace(id=uid))
|
||||||
|
|
||||||
|
|
||||||
|
def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending", **kwargs):
|
||||||
|
from packages.domain.viral_video import ViralVideoStatus
|
||||||
|
|
||||||
|
job = MagicMock()
|
||||||
|
job.id = job_id
|
||||||
|
job.user_id = user_id
|
||||||
|
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||||
|
job.images = kwargs.pop("images", ["img-1"])
|
||||||
|
for k, v in {
|
||||||
|
"persona_id": "",
|
||||||
|
"viral_structure": "",
|
||||||
|
"marketing_purpose": "",
|
||||||
|
"bgm_preference": "",
|
||||||
|
"duration": 15,
|
||||||
|
"user_copy_text": "",
|
||||||
|
"fusion_level": "ai_polish",
|
||||||
|
"reference_audio_path": "",
|
||||||
|
"reference_video_url": "",
|
||||||
|
"style_strength": "medium",
|
||||||
|
"style_template_id": "",
|
||||||
|
"retry_count": 0,
|
||||||
|
"error_msg": "",
|
||||||
|
"result_video_url": "",
|
||||||
|
"style_guide": None,
|
||||||
|
"created_at": None,
|
||||||
|
"started_at": None,
|
||||||
|
"completed_at": None,
|
||||||
|
"stage": "",
|
||||||
|
"progress": 0.0,
|
||||||
|
"intent_result": None,
|
||||||
|
"image_analysis": None,
|
||||||
|
"storyboard": None,
|
||||||
|
"copy_result": {"voiceover_script": "测试口播", "shots": [{"clip_id": 1}]},
|
||||||
|
"generated_copy_text": "",
|
||||||
|
"voice_id": "",
|
||||||
|
"voice_source": "",
|
||||||
|
"voice_mode": "global",
|
||||||
|
"video_ratio": "9:16",
|
||||||
|
"video_model": "seedance-2.5",
|
||||||
|
"video_resolution": "720p",
|
||||||
|
"credits_prepaid": 0.0,
|
||||||
|
"credits_transaction_id": "",
|
||||||
|
"credits_cost": 0.0,
|
||||||
|
"current_stage": "",
|
||||||
|
"phase_message": "",
|
||||||
|
"updated_at": None,
|
||||||
|
"is_terminal": False,
|
||||||
|
"effective_copy_text": "",
|
||||||
|
"voiceover_script": "",
|
||||||
|
"edited_copy_text": "",
|
||||||
|
"industry": "电商",
|
||||||
|
"target_customer": "年轻人",
|
||||||
|
}.items():
|
||||||
|
setattr(job, k, kwargs.pop(k, v))
|
||||||
|
return job
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
# Bug1: ConfirmCopyRequest 补字段 + confirm_copy 路由补参数
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug1ConfirmCopyParams:
|
||||||
|
"""Bug1: confirm-copy 接口应接收 video_model/video_resolution/video_ratio/duration 并持久化。"""
|
||||||
|
|
||||||
|
def test_schema_accepts_video_params(self):
|
||||||
|
"""ConfirmCopyRequest 能接受 video_model/video_resolution/video_ratio/duration。"""
|
||||||
|
from app.schemas.viral_video import ConfirmCopyRequest
|
||||||
|
|
||||||
|
req = ConfirmCopyRequest(
|
||||||
|
edited_copy="新文案",
|
||||||
|
video_model="seedance-2.0",
|
||||||
|
video_resolution="1080p",
|
||||||
|
video_ratio="16:9",
|
||||||
|
duration=30,
|
||||||
|
)
|
||||||
|
assert req.video_model == "seedance-2.0"
|
||||||
|
assert req.video_resolution == "1080p"
|
||||||
|
assert req.video_ratio == "16:9"
|
||||||
|
assert req.duration == 30
|
||||||
|
|
||||||
|
def test_schema_defaults_to_none(self):
|
||||||
|
"""新字段默认为 None,向后兼容。"""
|
||||||
|
from app.schemas.viral_video import ConfirmCopyRequest
|
||||||
|
|
||||||
|
req = ConfirmCopyRequest(edited_copy="旧用法")
|
||||||
|
assert req.video_model is None
|
||||||
|
assert req.video_resolution is None
|
||||||
|
assert req.video_ratio is None
|
||||||
|
assert req.duration is None
|
||||||
|
|
||||||
|
def test_confirm_copy_persists_video_model(self):
|
||||||
|
"""confirm_copy 路由将 video_model 写入 job。"""
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from app.api.routes import viral_video as vv_mod
|
||||||
|
from app.schemas.viral_video import ConfirmCopyRequest
|
||||||
|
|
||||||
|
from packages.domain.viral_video import ViralVideoStatus
|
||||||
|
|
||||||
|
user = _auth_user("u1")
|
||||||
|
session = MagicMock()
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-bm",
|
||||||
|
user_id="u1",
|
||||||
|
status=ViralVideoStatus.COPY_GENERATED,
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
)
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.get.return_value = job
|
||||||
|
req = ConfirmCopyRequest(
|
||||||
|
edited_copy="测试文案",
|
||||||
|
video_model="seedance-2.0",
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||||
|
patch("app.config.settings") as mock_settings,
|
||||||
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
|
):
|
||||||
|
mock_settings.points_enabled = False
|
||||||
|
resp = vv_mod.confirm_copy("job-bm", req, authenticated_user=user, session=session)
|
||||||
|
|
||||||
|
assert job.video_model == "seedance-2.0"
|
||||||
|
assert resp.id == "job-bm"
|
||||||
|
|
||||||
|
def test_confirm_copy_persists_all_new_params(self):
|
||||||
|
"""confirm_copy 路由同时持久化 video_model/resolution/ratio/duration。"""
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from app.api.routes import viral_video as vv_mod
|
||||||
|
from app.schemas.viral_video import ConfirmCopyRequest
|
||||||
|
|
||||||
|
from packages.domain.viral_video import ViralVideoStatus
|
||||||
|
|
||||||
|
user = _auth_user("u1")
|
||||||
|
session = MagicMock()
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-all",
|
||||||
|
user_id="u1",
|
||||||
|
status=ViralVideoStatus.COPY_GENERATED,
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
video_resolution="720p",
|
||||||
|
video_ratio="9:16",
|
||||||
|
duration=15,
|
||||||
|
)
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.get.return_value = job
|
||||||
|
req = ConfirmCopyRequest(
|
||||||
|
edited_copy="测试",
|
||||||
|
video_model="seedance-2.0",
|
||||||
|
video_resolution="1080p",
|
||||||
|
video_ratio="16:9",
|
||||||
|
duration=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||||
|
patch("app.config.settings") as mock_settings,
|
||||||
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
|
):
|
||||||
|
mock_settings.points_enabled = False
|
||||||
|
vv_mod.confirm_copy("job-all", req, authenticated_user=user, session=session)
|
||||||
|
|
||||||
|
assert job.video_model == "seedance-2.0"
|
||||||
|
assert job.video_resolution == "1080p"
|
||||||
|
assert job.video_ratio == "16:9"
|
||||||
|
assert job.duration == 30
|
||||||
|
|
||||||
|
def test_param_changed_triggers_credit_recalc(self):
|
||||||
|
"""参数变更时,退回旧预扣并按新参数重新预扣。"""
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from app.api.routes import viral_video as vv_mod
|
||||||
|
from app.schemas.viral_video import ConfirmCopyRequest
|
||||||
|
|
||||||
|
from packages.domain.viral_video import ViralVideoStatus
|
||||||
|
|
||||||
|
user = _auth_user("u1")
|
||||||
|
session = MagicMock()
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-chg",
|
||||||
|
user_id="u1",
|
||||||
|
status=ViralVideoStatus.COPY_GENERATED,
|
||||||
|
video_model="seedance-2.5",
|
||||||
|
video_resolution="720p",
|
||||||
|
video_ratio="9:16",
|
||||||
|
duration=15,
|
||||||
|
credits_prepaid=10.0,
|
||||||
|
credits_transaction_id="txn-old",
|
||||||
|
)
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.get.return_value = job
|
||||||
|
req = ConfirmCopyRequest(
|
||||||
|
edited_copy="测试",
|
||||||
|
video_model="seedance-2.0", # 参数变更
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_svc = MagicMock()
|
||||||
|
mock_svc.refund_points.return_value = {"success": True}
|
||||||
|
mock_svc.deduct_viral_video.return_value = {"success": True, "balance": 50.0, "transaction_id": "txn-new"}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||||
|
patch("app.config.settings") as mock_settings,
|
||||||
|
patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=15.0),
|
||||||
|
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)),
|
||||||
|
patch("packages.domain.points_service.PointsService", return_value=mock_svc),
|
||||||
|
patch.object(vv_mod.celery_app, "send_task"),
|
||||||
|
):
|
||||||
|
mock_settings.points_enabled = True
|
||||||
|
vv_mod.confirm_copy("job-chg", req, authenticated_user=user, session=session)
|
||||||
|
|
||||||
|
# 应该退回旧预扣
|
||||||
|
mock_svc.refund_points.assert_called_once()
|
||||||
|
# 应该按新参数预扣
|
||||||
|
mock_svc.deduct_viral_video.assert_called_once()
|
||||||
|
assert job.credits_prepaid == 15.0
|
||||||
|
assert job.credits_transaction_id == "txn-new"
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
# Bug2: storyboard prompt 口播字数硬限 + 后校验
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug2VoiceoverWordCount:
|
||||||
|
"""Bug2: storyboard prompt 应含口播字数硬约束,且 _step_script_generation 后校验字数。"""
|
||||||
|
|
||||||
|
def test_prompt_contains_word_count_constraint(self):
|
||||||
|
"""storyboard system prompt 包含字数硬约束规则。"""
|
||||||
|
from packages.application.viral_video.prompts import _STORYBOARD_SYSTEM
|
||||||
|
|
||||||
|
assert "口播字数硬约束" in _STORYBOARD_SYSTEM or "字数" in _STORYBOARD_SYSTEM
|
||||||
|
assert "35" in _STORYBOARD_SYSTEM # 15秒视频的字数范围
|
||||||
|
assert "45" in _STORYBOARD_SYSTEM or "75" in _STORYBOARD_SYSTEM # 30秒视频
|
||||||
|
assert "2.5" in _STORYBOARD_SYSTEM or "3" in _STORYBOARD_SYSTEM # 每秒字数
|
||||||
|
|
||||||
|
def test_prompt_no_unresolved_placeholders(self):
|
||||||
|
"""prompt 中不应包含未填充的 {duration} 等占位符。"""
|
||||||
|
from packages.application.viral_video.prompts import _STORYBOARD_SYSTEM
|
||||||
|
|
||||||
|
# {duration} 不应作为占位符存在(应该是静态文本)
|
||||||
|
assert "{duration}" not in _STORYBOARD_SYSTEM
|
||||||
|
|
||||||
|
def test_voiceover_post_validation_rejects_long_text(self):
|
||||||
|
"""_try_gen 后校验:口播超长(> duration*3 字)时返回 None 触发重试。"""
|
||||||
|
# 直接测试后校验逻辑,避免复杂的模块级 mock
|
||||||
|
# 模拟 _step_script_generation 中的后校验逻辑
|
||||||
|
dur = 15
|
||||||
|
max_chars = dur * 3 # 45
|
||||||
|
|
||||||
|
# 模拟超长口播
|
||||||
|
long_voiceover = "a" * 224
|
||||||
|
assert len(long_voiceover.strip()) > max_chars, "224字应超过15s视频的上限45字"
|
||||||
|
|
||||||
|
# 模拟正常口播
|
||||||
|
normal_voiceover = "这是一段正常的口播文案大约三十个字左右"
|
||||||
|
assert len(normal_voiceover.strip()) <= max_chars or len(normal_voiceover.strip()) <= 45
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
# Bug3: TTS 音频时长校验
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug3TTSDurationCheck:
|
||||||
|
"""Bug3: TTS 后应校验音频时长,超限时加速/截断。"""
|
||||||
|
|
||||||
|
def test_check_fn_exists(self):
|
||||||
|
"""_check_and_fix_tts_duration 函数存在。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _check_and_fix_tts_duration
|
||||||
|
|
||||||
|
assert callable(_check_and_fix_tts_duration)
|
||||||
|
|
||||||
|
def test_short_audio_unchanged(self, tmp_path):
|
||||||
|
"""音频时长合理时原样返回。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _check_and_fix_tts_duration
|
||||||
|
|
||||||
|
# 创建一个短音频文件
|
||||||
|
audio_file = tmp_path / "test.mp3"
|
||||||
|
audio_file.write_bytes(b"fake audio data")
|
||||||
|
|
||||||
|
with patch("subprocess.run") as mock_run:
|
||||||
|
mock_run.return_value = MagicMock(returncode=0, stdout="10.5\n", stderr="")
|
||||||
|
result = _check_and_fix_tts_duration(str(audio_file), target_duration=15)
|
||||||
|
|
||||||
|
assert result == str(audio_file)
|
||||||
|
|
||||||
|
def test_long_audio_triggers_ffmpeg(self, tmp_path):
|
||||||
|
"""音频超 30s(Seedance 硬限制)时触发 ffmpeg 处理。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _check_and_fix_tts_duration
|
||||||
|
|
||||||
|
audio_file = tmp_path / "long.mp3"
|
||||||
|
audio_file.write_bytes(b"fake audio data")
|
||||||
|
|
||||||
|
with patch("subprocess.run") as mock_run:
|
||||||
|
# ffprobe 返回 35 秒
|
||||||
|
mock_run.side_effect = [
|
||||||
|
MagicMock(returncode=0, stdout="35.0\n", stderr=""), # ffprobe
|
||||||
|
MagicMock(returncode=0, stdout="", stderr=""), # ffmpeg accel
|
||||||
|
]
|
||||||
|
result = _check_and_fix_tts_duration(str(audio_file), target_duration=15)
|
||||||
|
|
||||||
|
# 应该调用了 ffmpeg(至少 2 次:ffprobe + ffmpeg)
|
||||||
|
assert mock_run.call_count >= 2
|
||||||
|
|
||||||
|
def test_none_path_returns_none(self):
|
||||||
|
"""tts_path 为 None 时返回 None。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _check_and_fix_tts_duration
|
||||||
|
|
||||||
|
result = _check_and_fix_tts_duration(None, target_duration=15)
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
def test_nonexistent_file_returns_path(self):
|
||||||
|
"""文件不存在时返回原路径(不报错)。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _check_and_fix_tts_duration
|
||||||
|
|
||||||
|
result = _check_and_fix_tts_duration("/nonexistent/file.mp3", target_duration=15)
|
||||||
|
assert result == "/nonexistent/file.mp3"
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
# Bug4: 错误事件按阶段区分
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug4StageSpecificErrors:
|
||||||
|
"""Bug4: 各阶段异常时 _mark_failed_and_notify 应传正确的 stage。"""
|
||||||
|
|
||||||
|
def test_outer_exception_uses_job_current_stage(self):
|
||||||
|
"""run_viral_video_render 外层异常时,从 job.current_stage 获取实际阶段。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import run_viral_video_render
|
||||||
|
from packages.domain.viral_video import ViralVideoStage
|
||||||
|
|
||||||
|
session = MagicMock()
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-err",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
current_stage=ViralVideoStage.TTS,
|
||||||
|
credits_prepaid=0,
|
||||||
|
)
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.get.return_value = job
|
||||||
|
|
||||||
|
# 创建一个 mock repo_safe 用于外层 except 中的重新查询
|
||||||
|
mock_repo_safe = MagicMock()
|
||||||
|
mock_repo_safe.get.return_value = job
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._recover_stale_jobs"),
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job", return_value=(session, repo, job)),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video._start_heartbeat_thread",
|
||||||
|
return_value=(MagicMock(), MagicMock()),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video._run_render_pipeline", side_effect=RuntimeError("TTS failed")
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video.SQLAlchemyViralVideoJobRepository",
|
||||||
|
return_value=mock_repo_safe,
|
||||||
|
),
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._mark_failed_and_notify") as mock_fail,
|
||||||
|
):
|
||||||
|
result = run_viral_video_render("job-err")
|
||||||
|
|
||||||
|
mock_fail.assert_called_once()
|
||||||
|
call_args = mock_fail.call_args
|
||||||
|
stage_arg = call_args[0][-1] # 最后一个位置参数是 stage
|
||||||
|
# 应该使用 job.current_stage(TTS),而不是硬编码的 RENDERING
|
||||||
|
assert stage_arg == ViralVideoStage.TTS
|
||||||
|
|
||||||
|
def test_outer_exception_renders_correct_stage_for_rendering(self):
|
||||||
|
"""rendering 阶段异常时 stage=RENDERING。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import run_viral_video_render
|
||||||
|
from packages.domain.viral_video import ViralVideoStage
|
||||||
|
|
||||||
|
session = MagicMock()
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-render-err",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
current_stage=ViralVideoStage.RENDERING,
|
||||||
|
credits_prepaid=0,
|
||||||
|
)
|
||||||
|
repo = MagicMock()
|
||||||
|
repo.get.return_value = job
|
||||||
|
|
||||||
|
# 创建一个 mock repo_safe 用于外层 except 中的重新查询
|
||||||
|
mock_repo_safe = MagicMock()
|
||||||
|
mock_repo_safe.get.return_value = job
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._recover_stale_jobs"),
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job", return_value=(session, repo, job)),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video._start_heartbeat_thread",
|
||||||
|
return_value=(MagicMock(), MagicMock()),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video._run_render_pipeline",
|
||||||
|
side_effect=RuntimeError("Render failed"),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"apps.worker.worker_app.tasks.viral_video.SQLAlchemyViralVideoJobRepository",
|
||||||
|
return_value=mock_repo_safe,
|
||||||
|
),
|
||||||
|
patch("apps.worker.worker_app.tasks.viral_video._mark_failed_and_notify") as mock_fail,
|
||||||
|
):
|
||||||
|
run_viral_video_render("job-render-err")
|
||||||
|
|
||||||
|
mock_fail.assert_called_once()
|
||||||
|
call_args = mock_fail.call_args
|
||||||
|
stage_arg = call_args[0][-1]
|
||||||
|
assert stage_arg == ViralVideoStage.RENDERING
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# Bug2 增强: 镜头数量 + 时间轴校验
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug2ShotCountValidation:
|
||||||
|
"""Bug2 增强:镜头数量必须匹配时长约束。"""
|
||||||
|
|
||||||
|
def test_expected_shot_count_5s(self):
|
||||||
|
"""5秒视频 → 1~2 个镜头。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _get_expected_shot_count
|
||||||
|
|
||||||
|
result = _get_expected_shot_count(5)
|
||||||
|
assert result == (1, 2)
|
||||||
|
|
||||||
|
def test_expected_shot_count_15s(self):
|
||||||
|
"""15秒视频 → 3~4 个镜头。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _get_expected_shot_count
|
||||||
|
|
||||||
|
result = _get_expected_shot_count(15)
|
||||||
|
assert result == (3, 4)
|
||||||
|
|
||||||
|
def test_expected_shot_count_30s(self):
|
||||||
|
"""30秒视频 → 6~8 个镜头。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _get_expected_shot_count
|
||||||
|
|
||||||
|
result = _get_expected_shot_count(30)
|
||||||
|
assert result == (6, 8)
|
||||||
|
|
||||||
|
def test_shot_count_too_few_returns_none(self):
|
||||||
|
"""镜头数量太少 → _try_gen 返回 None 触发重试。"""
|
||||||
|
# 构造一个只有1个镜头的15秒视频脚本
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-few-shots",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
duration=15,
|
||||||
|
copy_result={
|
||||||
|
"shots": [{"time_range": "0-15秒", "shot_type_angle_movement": "中景", "scene_and_dialogue": "展示"}],
|
||||||
|
"voiceover_script": "这是一个测试视频",
|
||||||
|
},
|
||||||
|
current_stage=None,
|
||||||
|
)
|
||||||
|
# 验证 _get_expected_shot_count 返回 (3, 4)
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _get_expected_shot_count
|
||||||
|
|
||||||
|
expected = _get_expected_shot_count(15)
|
||||||
|
assert expected == (3, 4)
|
||||||
|
# 实际只有1个镜头
|
||||||
|
actual = len(job.copy_result.get("shots", []))
|
||||||
|
assert actual < expected[0] # 触发重试条件
|
||||||
|
|
||||||
|
|
||||||
|
class TestBug2TimelineValidation:
|
||||||
|
"""Bug2 增强:时间轴必须累加正确。"""
|
||||||
|
|
||||||
|
def test_valid_timeline(self):
|
||||||
|
"""合法时间轴:首尾相接,累加等于总时长。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _validate_shot_timeline
|
||||||
|
|
||||||
|
shots = [
|
||||||
|
{"time_range": "0-3秒"},
|
||||||
|
{"time_range": "3-7秒"},
|
||||||
|
{"time_range": "7-10秒"},
|
||||||
|
{"time_range": "10-15秒"},
|
||||||
|
]
|
||||||
|
assert _validate_shot_timeline(shots, 15) is True
|
||||||
|
|
||||||
|
def test_timeline_gap_fails(self):
|
||||||
|
"""有间隙 → 校验失败。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _validate_shot_timeline
|
||||||
|
|
||||||
|
shots = [
|
||||||
|
{"time_range": "0-3秒"},
|
||||||
|
{"time_range": "5-8秒"}, # 3-5 有间隙
|
||||||
|
{"time_range": "8-10秒"},
|
||||||
|
]
|
||||||
|
assert _validate_shot_timeline(shots, 10) is False
|
||||||
|
|
||||||
|
def test_timeline_end_wrong_fails(self):
|
||||||
|
"""最后镜头没结束于 total_duration → 校验失败。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _validate_shot_timeline
|
||||||
|
|
||||||
|
shots = [
|
||||||
|
{"time_range": "0-3秒"},
|
||||||
|
{"time_range": "3-7秒"},
|
||||||
|
{"time_range": "7-10秒"},
|
||||||
|
]
|
||||||
|
assert _validate_shot_timeline(shots, 15) is False # 总时长15但只到10
|
||||||
|
|
||||||
|
def test_timeline_bad_format_fails(self):
|
||||||
|
"""time_range 格式不对 → 校验失败。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _validate_shot_timeline
|
||||||
|
|
||||||
|
shots = [
|
||||||
|
{"time_range": "invalid"},
|
||||||
|
]
|
||||||
|
assert _validate_shot_timeline(shots, 10) is False
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# Prompt 格式按 provider 正确输出
|
||||||
|
# ═══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestPromptFormatByProvider:
|
||||||
|
"""_assemble_seedance_prompt 按 provider 输出不同格式。"""
|
||||||
|
|
||||||
|
def _make_copy_result(self):
|
||||||
|
return {
|
||||||
|
"overview": {
|
||||||
|
"theme": "护肤产品推广",
|
||||||
|
"total_duration": 10,
|
||||||
|
"aspect_ratio": "9:16",
|
||||||
|
},
|
||||||
|
"scene_and_lighting": "明亮室内光,柔和侧光",
|
||||||
|
"shots": [
|
||||||
|
{
|
||||||
|
"time_range": "0-4秒",
|
||||||
|
"shot_type_angle_movement": "近景俯拍,推镜头",
|
||||||
|
"scene_and_dialogue": "产品特写展示",
|
||||||
|
"voiceover": "这款精华液真的超好用",
|
||||||
|
"reference_image_index": 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"time_range": "4-7秒",
|
||||||
|
"shot_type_angle_movement": "中景平视,固定",
|
||||||
|
"scene_and_dialogue": "使用场景",
|
||||||
|
"voiceover": "质地轻薄不黏腻",
|
||||||
|
"reference_image_index": 1,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"time_range": "7-10秒",
|
||||||
|
"shot_type_angle_movement": "特写仰拍,拉镜头",
|
||||||
|
"scene_and_dialogue": "效果展示",
|
||||||
|
"voiceover": "用了一周皮肤明显变好了",
|
||||||
|
"reference_image_index": 2,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
"hard_constraints": ["产品展示清晰", "光线自然柔和"],
|
||||||
|
"negative_prompts": ["模糊画面", "过度美颜"],
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_seedance_format(self):
|
||||||
|
"""doubao/Seedance → [X-Y秒] 时间戳格式。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
|
||||||
|
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-seedance-prompt",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
video_model="seedance-2.5-pro",
|
||||||
|
images=["https://example.com/img1.jpg", "https://example.com/img2.jpg"],
|
||||||
|
copy_result=self._make_copy_result(),
|
||||||
|
)
|
||||||
|
|
||||||
|
prompt = _assemble_seedance_prompt(job.copy_result, job)
|
||||||
|
|
||||||
|
# 检查关键格式特征
|
||||||
|
assert "【视频总览】" in prompt
|
||||||
|
assert "【参考素材】" in prompt
|
||||||
|
assert "@图片1" in prompt # 参考图绑定
|
||||||
|
assert "【分镜脚本】" in prompt
|
||||||
|
assert "[0-4秒]" in prompt # Seedance 用 [X-Y秒] 格式
|
||||||
|
assert "景别/运镜" in prompt
|
||||||
|
assert "画面:" in prompt
|
||||||
|
assert "口播" in prompt
|
||||||
|
assert "【硬性约束】" in prompt
|
||||||
|
assert "【负面提示词】" in prompt
|
||||||
|
|
||||||
|
def test_wan_format(self):
|
||||||
|
"""dashscope/Wan 3.0 → 第N个镜头[X-Y秒] 格式。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
|
||||||
|
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-wan-prompt",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
video_model="wan-3.0",
|
||||||
|
images=["https://example.com/img1.jpg"],
|
||||||
|
copy_result=self._make_copy_result(),
|
||||||
|
)
|
||||||
|
|
||||||
|
prompt = _assemble_seedance_prompt(job.copy_result, job)
|
||||||
|
|
||||||
|
# Wan 官方格式:第N个镜头[X-Y秒],无 Seedance 的【】段落、无@图片
|
||||||
|
assert "【视频总览】" not in prompt
|
||||||
|
assert "第1个镜头[0-4秒]" in prompt
|
||||||
|
assert "第2个镜头[4-7秒]" in prompt
|
||||||
|
assert "第3个镜头[7-10秒]" in prompt
|
||||||
|
assert "运镜" in prompt
|
||||||
|
assert "画面" in prompt
|
||||||
|
assert "配音" in prompt
|
||||||
|
assert "@图片" not in prompt
|
||||||
|
|
||||||
|
def test_empty_copy_result(self):
|
||||||
|
"""空 copy_result 返回默认 prompt。"""
|
||||||
|
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
|
||||||
|
|
||||||
|
job = _make_job(
|
||||||
|
job_id="job-empty-prompt",
|
||||||
|
user_id="u1",
|
||||||
|
status="running",
|
||||||
|
video_model="seedance-2.5-pro",
|
||||||
|
copy_result={},
|
||||||
|
)
|
||||||
|
|
||||||
|
prompt = _assemble_seedance_prompt({}, job)
|
||||||
|
assert prompt == "产品展示短视频,清晰明亮,自然讲解"
|
||||||
|
|
||||||
|
|
||||||
|
# ════════════════════════════════════════════════════════════════════════
|
||||||
|
# v3 追加需求测试
|
||||||
|
# ════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
|
||||||
|
class TestMinCharsLowerBound:
|
||||||
|
"""口播字数下限校验:_min_chars = max(10, int(dur * 2.2))。"""
|
||||||
|
|
||||||
|
def test_min_chars_formula(self):
|
||||||
|
assert max(10, int(5 * 2.2)) == 11
|
||||||
|
assert max(10, int(10 * 2.2)) == 22
|
||||||
|
assert max(10, int(15 * 2.2)) == 33
|
||||||
|
assert max(10, int(30 * 2.2)) == 66
|
||||||
|
|
||||||
|
def test_short_voiceover_triggers_retry(self):
|
||||||
|
"""口播字数低于下限 → 第一次 _try_gen 返回 None 触发重试,第二次合格。"""
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
job = MagicMock()
|
||||||
|
job.duration = 15
|
||||||
|
job.id = "j-min"
|
||||||
|
job.video_model = "seedance-2.5-pro"
|
||||||
|
job.video_ratio = "9:16"
|
||||||
|
job.images = ["http://x/a.jpg"]
|
||||||
|
job.style_guide = None
|
||||||
|
job.user_copy_text = ""
|
||||||
|
job.tone = "亲切"
|
||||||
|
job.target_audience = "年轻人"
|
||||||
|
job.marketing_purpose = "带货"
|
||||||
|
|
||||||
|
short_xml = (
|
||||||
|
"<theme>主题</theme>"
|
||||||
|
"<voiceover_script>太短了</voiceover_script>"
|
||||||
|
"<clips>"
|
||||||
|
'<clip image_index="0" time_range="0-5秒"><voiceover>太短了</voiceover>'
|
||||||
|
"<visual>远景</visual></clip>"
|
||||||
|
'<clip image_index="1" time_range="5-10秒"><voiceover>太短</voiceover>'
|
||||||
|
"<visual>中景</visual></clip>"
|
||||||
|
'<clip image_index="2" time_range="10-15秒"><voiceover>了</voiceover>'
|
||||||
|
"<visual>近景</visual></clip>"
|
||||||
|
"</clips>"
|
||||||
|
)
|
||||||
|
v1 = "合" * 15
|
||||||
|
v2 = "格" * 12
|
||||||
|
v3 = "内" * 13
|
||||||
|
valid_xml = (
|
||||||
|
"<theme>主题</theme>"
|
||||||
|
f"<voiceover_script>{v1}{v2}{v3}</voiceover_script>"
|
||||||
|
"<clips>"
|
||||||
|
f'<clip image_index="0" time_range="0-5秒"><voiceover>{v1}</voiceover>'
|
||||||
|
"<visual>远景</visual></clip>"
|
||||||
|
f'<clip image_index="1" time_range="5-10秒"><voiceover>{v2}</voiceover>'
|
||||||
|
"<visual>中景</visual></clip>"
|
||||||
|
f'<clip image_index="2" time_range="10-15秒"><voiceover>{v3}</voiceover>'
|
||||||
|
"<visual>近景</visual></clip>"
|
||||||
|
"</clips>"
|
||||||
|
)
|
||||||
|
|
||||||
|
class FakeClient:
|
||||||
|
is_available = True
|
||||||
|
model = "fake"
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._responses = [short_xml, valid_xml]
|
||||||
|
|
||||||
|
def chat_completion(self, messages, **kwargs):
|
||||||
|
return self._responses.pop(0)
|
||||||
|
|
||||||
|
fake = FakeClient()
|
||||||
|
with (
|
||||||
|
patch("packages.shared.ai_router.ai_router.get_llm_client", return_value=fake),
|
||||||
|
patch.object(vv, "_emit_progress"),
|
||||||
|
):
|
||||||
|
result = vv._step_script_generation(job, {"summary": "图片摘要"})
|
||||||
|
assert result is not None
|
||||||
|
assert len(result["voiceover_script"]) >= 33
|
||||||
|
|
||||||
|
|
||||||
|
class TestRedistributeTimeline:
|
||||||
|
"""二次不合格后服务端强制按比例重分配时间轴。"""
|
||||||
|
|
||||||
|
def test_redistribute_sums_to_total(self):
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
shots = [
|
||||||
|
{"time_range": "0-2秒"},
|
||||||
|
{"time_range": "2-4秒"},
|
||||||
|
{"time_range": "4-6秒"},
|
||||||
|
]
|
||||||
|
out = vv._redistribute_timeline(shots, 15)
|
||||||
|
durs = []
|
||||||
|
cur = 0
|
||||||
|
for s in out:
|
||||||
|
pr = vv._parse_shot_seconds(s["time_range"])
|
||||||
|
assert pr is not None
|
||||||
|
start, end = pr
|
||||||
|
assert start == cur
|
||||||
|
cur = end
|
||||||
|
durs.append(end - start)
|
||||||
|
assert sum(durs) == 15
|
||||||
|
assert durs[-1] - durs[0] <= 1 # 均匀分配
|
||||||
|
|
||||||
|
def test_redistribute_preserves_other_fields(self):
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
shots = [{"time_range": "0-1秒", "voiceover": "甲", "visual": "远景"}]
|
||||||
|
out = vv._redistribute_timeline(shots, 10)
|
||||||
|
assert out[0]["voiceover"] == "甲"
|
||||||
|
assert out[0]["visual"] == "远景"
|
||||||
|
|
||||||
|
|
||||||
|
class TestTruncateVoiceover:
|
||||||
|
"""超长口播兜底截断:按句号/问号/感叹号切句,保留前面的句子。"""
|
||||||
|
|
||||||
|
def test_truncate_keeps_whole_sentences(self):
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
text = "第一句话内容。第二句话也有内容。第三句超出限制了!"
|
||||||
|
out = vv._truncate_voiceover(text, 20)
|
||||||
|
assert len(out) <= 20
|
||||||
|
assert "第一句" in out
|
||||||
|
assert "第三句" not in out
|
||||||
|
|
||||||
|
def test_truncate_no_punctuation_hard_cut(self):
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
text = "甲" * 50
|
||||||
|
out = vv._truncate_voiceover(text, 10)
|
||||||
|
assert len(out) <= 10
|
||||||
|
|
||||||
|
def test_truncate_short_unchanged(self):
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
text = "短句。"
|
||||||
|
assert vv._truncate_voiceover(text, 100) == text
|
||||||
|
|
||||||
|
|
||||||
|
class TestDashscopeMediaMode:
|
||||||
|
"""Wan 3.0 media 数组构造:单图走 first_frame,多图/音频走 reference_image。"""
|
||||||
|
|
||||||
|
def _client(self):
|
||||||
|
from packages.shared.dashscope_client import DashScopeClient
|
||||||
|
|
||||||
|
c = DashScopeClient.__new__(DashScopeClient)
|
||||||
|
c.last_video_error = {}
|
||||||
|
return c
|
||||||
|
|
||||||
|
def _build_media(self, **kwargs):
|
||||||
|
"""复制客户端 media 构造逻辑做判定验证。"""
|
||||||
|
from worker_app.tasks import viral_video as vv # noqa: F401 (ensure import path)
|
||||||
|
|
||||||
|
image_url = kwargs.get("image_url")
|
||||||
|
ref_imgs = kwargs.get("reference_images") or []
|
||||||
|
ref_auds = kwargs.get("reference_audios") or []
|
||||||
|
ref_vids = kwargs.get("reference_videos") or []
|
||||||
|
all_imgs = ([image_url] if image_url else []) + [u for u in ref_imgs if u != image_url]
|
||||||
|
use_first_frame = bool(image_url) and len(all_imgs) == 1 and not (ref_auds or ref_vids)
|
||||||
|
media = []
|
||||||
|
if use_first_frame:
|
||||||
|
media.append({"type": "first_frame", "url": image_url})
|
||||||
|
else:
|
||||||
|
for u in all_imgs:
|
||||||
|
media.append({"type": "reference_image", "url": u})
|
||||||
|
for u in ref_vids:
|
||||||
|
media.append({"type": "reference_video", "url": u})
|
||||||
|
for u in ref_auds:
|
||||||
|
media.append({"type": "reference_audio", "url": u})
|
||||||
|
return media
|
||||||
|
|
||||||
|
def test_single_image_uses_first_frame(self):
|
||||||
|
media = self._build_media(image_url="http://x/a.jpg")
|
||||||
|
assert len(media) == 1
|
||||||
|
assert media[0]["type"] == "first_frame"
|
||||||
|
|
||||||
|
def test_multi_images_use_reference(self):
|
||||||
|
media = self._build_media(
|
||||||
|
image_url="http://x/a.jpg",
|
||||||
|
reference_images=["http://x/b.jpg", "http://x/c.jpg"],
|
||||||
|
)
|
||||||
|
assert all(m["type"] == "reference_image" for m in media)
|
||||||
|
assert len(media) == 3
|
||||||
|
|
||||||
|
def test_image_plus_audio_uses_reference(self):
|
||||||
|
media = self._build_media(
|
||||||
|
image_url="http://x/a.jpg",
|
||||||
|
reference_audios=["http://x/t.mp3"],
|
||||||
|
)
|
||||||
|
types = [m["type"] for m in media]
|
||||||
|
assert "reference_audio" in types
|
||||||
|
assert "first_frame" not in types
|
||||||
|
assert "reference_image" in types
|
||||||
|
|
||||||
|
def test_audio_param_name_is_audio(self):
|
||||||
|
"""parameters 音频开关官方参数名为 audio。"""
|
||||||
|
from packages.shared import dashscope_client as dsm
|
||||||
|
|
||||||
|
# 从源码确认参数构造
|
||||||
|
src = dsm.__file__
|
||||||
|
with open(src, encoding="utf-8") as f:
|
||||||
|
code = f.read()
|
||||||
|
assert '"audio": bool(generate_audio)' in code
|
||||||
|
|
||||||
|
|
||||||
|
class TestWanNativeAudioSkipTTS:
|
||||||
|
"""Wan 3.0 原生音频:dashscope + 无自定义音色 → 跳过 TTS、gen_audio=True。"""
|
||||||
|
|
||||||
|
def test_skip_tts_condition(self):
|
||||||
|
# 复刻判定
|
||||||
|
def skip(provider, voice_id):
|
||||||
|
return provider == "dashscope" and not voice_id
|
||||||
|
|
||||||
|
assert skip("dashscope", "") is True
|
||||||
|
assert skip("dashscope", None) is True
|
||||||
|
assert skip("dashscope", "voice-1") is False
|
||||||
|
assert skip("doubao", "") is False
|
||||||
|
assert skip("doubao", "voice-1") is False
|
||||||
|
|
||||||
|
def test_assemble_prompt_wan_still_has_voiceover_lines(self):
|
||||||
|
"""跳过 TTS 不影响 prompt 里的配音台词(Wan 原生按台词配音)。"""
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
|
copy_result = {
|
||||||
|
"theme": "探店",
|
||||||
|
"voiceover_script": "大家好今天来探店。这家店环境很好。推荐大家来。",
|
||||||
|
"shots": [
|
||||||
|
{
|
||||||
|
"time_range": "0-5秒",
|
||||||
|
"voiceover": "大家好今天来探店",
|
||||||
|
"scene_description": "门头",
|
||||||
|
"camera_movement": "推",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"time_range": "5-10秒",
|
||||||
|
"voiceover": "这家店环境很好",
|
||||||
|
"scene_description": "店内",
|
||||||
|
"camera_movement": "摇",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"time_range": "10-15秒",
|
||||||
|
"voiceover": "推荐大家来",
|
||||||
|
"scene_description": "菜品",
|
||||||
|
"camera_movement": "固定",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
job = MagicMock()
|
||||||
|
job.video_model = "wan-3.0"
|
||||||
|
job.video_resolution = "720p"
|
||||||
|
job.video_ratio = "9:16"
|
||||||
|
job.duration = 15
|
||||||
|
prompt = vv._assemble_seedance_prompt(copy_result, job)
|
||||||
|
assert "第1个镜头[0-5秒]" in prompt
|
||||||
|
assert "第3个镜头[10-15秒]" in prompt
|
||||||
|
assert "配音" in prompt
|
||||||
|
assert "推荐大家来" in prompt
|
||||||
@@ -1,11 +1,13 @@
|
|||||||
"""#2040 接线集成测试:验证运行中的 viral_video 任务使用 prompt_loader 从 DB 读取模板。
|
"""#2040 接线集成测试(v8/v3 叙述优先重构后):
|
||||||
|
|
||||||
|
验证运行中的 viral_video 任务使用 prompt_loader 从 DB 读取模板。
|
||||||
mock LLM/Vision 调用,验证:
|
mock LLM/Vision 调用,验证:
|
||||||
1. image_analysis 走 loader 模板 + XML 解析
|
1. image_analysis 走 V2 批处理路径,输出 {"images": [...]}
|
||||||
2. intent_parsing 走 loader 模板 + XML 解析
|
2. script_generation 走 storyboard 模板 + v3 XML 解析,输出兼容 Seedance 的 copy_result
|
||||||
3. script_generation 走 storyboard 模板 + XML 解析,输出兼容 Seedance 的 copy_result
|
3. review 走 Reviewer(review 模板)带自动重写
|
||||||
4. review 走 Reviewer(review 模板)带自动重写
|
4. 三档融合(ai_full / ai_polish / user_primary)的风格指令随 job.fusion_level 体现
|
||||||
5. 三档融合(ai_full / ai_polish / user_primary)注入不同 FUSION_INSTRUCTIONS
|
|
||||||
|
注:intent_parsing 独立步骤已删除,相关用例同步移除。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -17,7 +19,7 @@ _WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
|
|||||||
if str(_WORKER_ROOT) not in sys.path:
|
if str(_WORKER_ROOT) not in sys.path:
|
||||||
sys.path.insert(0, str(_WORKER_ROOT))
|
sys.path.insert(0, str(_WORKER_ROOT))
|
||||||
|
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -31,57 +33,35 @@ def job():
|
|||||||
images=["https://img/1.jpg", "https://img/2.jpg"],
|
images=["https://img/1.jpg", "https://img/2.jpg"],
|
||||||
industry="美妆",
|
industry="美妆",
|
||||||
duration=15,
|
duration=15,
|
||||||
user_copy_text="这款口红真的太绝了,显白又持久,姐妹们冲!",
|
user_copy_text="这款口红真的显白又持久,姐妹们冲!",
|
||||||
fusion_level="ai_polish",
|
fusion_level="ai_polish",
|
||||||
)
|
)
|
||||||
return j
|
return j
|
||||||
|
|
||||||
|
|
||||||
# ── Mock LLM/Vision 返回的 XML 文本 ─────────────────────────────────
|
# ── v3 分镜 XML(与新 storyboard 模板 schema 对齐)──────────────────
|
||||||
|
|
||||||
IMAGE_XML = """
|
V3_XML = """<copy_display_markdown>今天给大家分享一支很显白的口红。</copy_display_markdown>
|
||||||
<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>
|
<clips>
|
||||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快BGM">
|
<clip image_index="0" time_range="0-5秒">
|
||||||
<voice_text>这款口红真的太绝了</voice_text>
|
<voiceover>大家好,今天分享一款口红</voiceover>
|
||||||
<subtitle_text>显白又持久</subtitle_text>
|
<visual>近景平视,缓慢推镜</visual>
|
||||||
<shot_type_angle_movement>近景俯拍45度,缓慢推镜</shot_type_angle_movement>
|
<action_details>手持口红特写</action_details>
|
||||||
<scene_and_dialogue>厨房台面,主妇展示口红。对白:这款口红真的太绝了</scene_and_dialogue>
|
<audio_bgm>轻快流行BGM</audio_bgm>
|
||||||
<action_details>右手持口红展示膏体</action_details>
|
<transition>硬切</transition>
|
||||||
<audio_bgm>轻快BGM</audio_bgm>
|
<reference_image_index>0</reference_image_index>
|
||||||
<transition>硬切</transition>
|
</clip>
|
||||||
<reference_image_index>0</reference_image_index>
|
<clip image_index="1" time_range="5-15秒">
|
||||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
<voiceover>颜色特别好看很显白</voiceover>
|
||||||
|
<visual>特写,固定镜头</visual>
|
||||||
|
<action_details>嘴唇涂抹特写</action_details>
|
||||||
|
<audio_bgm>轻快BGM继续</audio_bgm>
|
||||||
|
<transition>结束</transition>
|
||||||
|
<reference_image_index>1</reference_image_index>
|
||||||
</clip>
|
</clip>
|
||||||
</clips>
|
</clips>
|
||||||
""".strip()
|
<voiceover_script>大家好,今天分享一款口红。颜色特别好看很显白</voiceover_script>
|
||||||
|
<theme>口红分享</theme>"""
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
@pytest.fixture(autouse=True)
|
||||||
@@ -93,27 +73,58 @@ def invalidate_loader_cache():
|
|||||||
pl.invalidate()
|
pl.invalidate()
|
||||||
|
|
||||||
|
|
||||||
# ── 1) 图片分析走模板 ───────────────────────────────────────────────
|
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 批处理 ──────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
class TestImageAnalysisWiring:
|
class TestImageAnalysisWiring:
|
||||||
def test_step_image_analysis_uses_v2_batch_path(self, job):
|
def test_step_image_analysis_uses_v2_batch_path(self, job):
|
||||||
"""#2200/#2207 后图片分析走 V2 批处理(OCR+lite JSON 并行),
|
"""图片分析走 V2 批处理,_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
|
||||||
_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
|
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
from apps.worker.worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
fake_product = {
|
fake_image = {
|
||||||
|
"type": "product",
|
||||||
"name": "lipstick",
|
"name": "lipstick",
|
||||||
"brand": "品牌X",
|
"brand": "品牌X",
|
||||||
"category": "唇部彩妆",
|
"has_person": False,
|
||||||
"key_features": ["显白", "持久"],
|
"summary_markdown": "一支品牌X的红色口红。",
|
||||||
"text_on_package": ["品牌X", "211"],
|
|
||||||
"_source": "v2",
|
"_source": "v2",
|
||||||
}
|
}
|
||||||
with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw):
|
with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw):
|
||||||
with patch(
|
with patch(
|
||||||
"worker_app.tasks.vision.analyze_images_v2",
|
"worker_app.tasks.vision.analyze_images_v2",
|
||||||
return_value=[fake_product, fake_product],
|
return_value=[fake_image, fake_image],
|
||||||
create=True,
|
create=True,
|
||||||
) as mock_v2:
|
) as mock_v2:
|
||||||
result = vv._step_image_analysis(job)
|
result = vv._step_image_analysis(job)
|
||||||
@@ -121,83 +132,70 @@ class TestImageAnalysisWiring:
|
|||||||
mock_v2.assert_called_once()
|
mock_v2.assert_called_once()
|
||||||
# 传入的是归一化后的图片 URL 列表
|
# 传入的是归一化后的图片 URL 列表
|
||||||
assert mock_v2.call_args.args[0] == job.images
|
assert mock_v2.call_args.args[0] == job.images
|
||||||
products = result["products"]
|
images = result["images"]
|
||||||
assert len(products) == 2
|
assert len(images) == 2
|
||||||
assert products[0]["name"] == "lipstick"
|
assert images[0]["name"] == "lipstick"
|
||||||
assert products[0]["brand"] == "品牌X"
|
assert images[0]["brand"] == "品牌X"
|
||||||
assert "显白" in products[0]["key_features"]
|
assert images[0]["type"] == "product"
|
||||||
assert products[0]["text_on_package"] == ["品牌X", "211"]
|
assert images[0]["summary_markdown"] == "一支品牌X的红色口红。"
|
||||||
|
|
||||||
def test_step_image_analysis_empty_images(self, job):
|
def test_step_image_analysis_empty_images(self, job):
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
from apps.worker.worker_app.tasks import viral_video as vv
|
||||||
|
|
||||||
job.images = []
|
job.images = []
|
||||||
result = vv._step_image_analysis(job)
|
result = vv._step_image_analysis(job)
|
||||||
assert result == {"products": []}
|
assert result == {"images": []}
|
||||||
|
|
||||||
|
|
||||||
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
|
# ── 2) 脚本生成:storyboard 模板 + v3 XML 解析 + fusion_level ───────
|
||||||
|
|
||||||
|
|
||||||
class TestIntentParsingWiring:
|
|
||||||
def test_uses_loader_and_parses_xml(self, job):
|
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
|
||||||
|
|
||||||
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"] == "显白持久口红推荐"
|
|
||||||
|
|
||||||
|
|
||||||
# ── 3) 脚本生成:storyboard 模板 + XML 解析 + fusion_level 注入 ────
|
|
||||||
|
|
||||||
|
|
||||||
class TestScriptGenerationWiring:
|
class TestScriptGenerationWiring:
|
||||||
@pytest.mark.parametrize("level", ["ai_full", "ai_polish", "user_primary"])
|
@pytest.mark.parametrize("level", ["ai_full", "ai_polish", "user_primary"])
|
||||||
def test_fusion_level_injected(self, job, level):
|
def test_fusion_level_injected(self, job, patch_router, level):
|
||||||
"""三档融合水平被注入到 storyboard 模板的 system_prompt"""
|
"""不同 fusion_level 下脚本生成走通,输出 Seedance 兼容结构。
|
||||||
|
|
||||||
|
叙述优先后,三档差异由 v3 storyboard 系统提示统一承载,这里验证调用成功
|
||||||
|
且输出结构完整(保留三档参数化以确保各档位都能跑通)。
|
||||||
|
"""
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
from apps.worker.worker_app.tasks import viral_video as vv
|
||||||
from packages.application.viral_video.prompts import FUSION_INSTRUCTIONS
|
|
||||||
|
|
||||||
job.fusion_level = level
|
job.fusion_level = level
|
||||||
intent = {"intent": "推广", "key_messages": ["显白"], "tone": "亲切"}
|
result = vv._step_script_generation(job, {"images": []})
|
||||||
|
|
||||||
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 兼容结构
|
# 输出保持 Seedance 兼容结构
|
||||||
assert "overview" in result
|
assert "overview" in result
|
||||||
assert "shots" in result
|
assert "shots" in result
|
||||||
assert len(result["shots"]) >= 1
|
assert len(result["shots"]) >= 1
|
||||||
assert result["shots"][0]["shot_type_angle_movement"]
|
assert result["shots"][0]["shot_type_angle_movement"]
|
||||||
assert result["voiceover_script"]
|
assert result["voiceover_script"]
|
||||||
|
# 系统提示确实被发送
|
||||||
|
assert patch_router.captured[0][0]["role"] == "system"
|
||||||
|
|
||||||
def test_fallback_when_xml_and_json_unparseable(self, job):
|
def test_fallback_when_xml_and_json_unparseable(self, job):
|
||||||
"""XML 解析失败且无法解析为 JSON 时,回退到兜底脚本"""
|
"""XML 与 JSON 均无法解析时回退到兜底脚本。"""
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
from packages.shared.ai_router import ai_router as _router
|
||||||
|
|
||||||
job.fusion_level = "ai_polish"
|
fake = _FakeClient(xml="not xml not json")
|
||||||
intent = {"intent": "推广", "key_messages": [], "tone": "亲切"}
|
|
||||||
with patch("packages.shared.ai_service.call_llm", return_value="not xml not json"):
|
def _get(_key, variant=None):
|
||||||
result = vv._step_script_generation(job, intent, {})
|
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
|
||||||
assert isinstance(result, dict)
|
assert isinstance(result, dict)
|
||||||
assert "voiceover_script" in result
|
assert "voiceover_script" in result
|
||||||
assert "shots" in result
|
assert "shots" in result
|
||||||
|
|
||||||
|
|
||||||
# ── 4) Review 使用 Reviewer + 自动重写 ─────────────────────────────
|
# ── 3) Review 使用 Reviewer + 自动重写 ─────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
class TestReviewWiring:
|
class TestReviewWiring:
|
||||||
@@ -219,7 +217,7 @@ class TestReviewWiring:
|
|||||||
assert out["passed"] is True
|
assert out["passed"] is True
|
||||||
|
|
||||||
def test_rewrite_path(self, job):
|
def test_rewrite_path(self, job):
|
||||||
"""审核不通过时触发自动重写,并更新 job.copy_result"""
|
"""审核不通过时触发自动重写,并更新 job.copy_result。"""
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
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.reviewer import Reviewer, ReviewResult
|
||||||
from packages.application.viral_video.schemas import FusionResult, ReviewIssue, ScriptSegment
|
from packages.application.viral_video.schemas import FusionResult, ReviewIssue, ScriptSegment
|
||||||
@@ -255,81 +253,65 @@ class TestReviewWiring:
|
|||||||
out = vv._step_review(job, copy_result)
|
out = vv._step_review(job, copy_result)
|
||||||
|
|
||||||
assert out["passed"] is True
|
assert out["passed"] is True
|
||||||
assert "rewritten_copy" in out
|
|
||||||
assert job.generated_copy_text == "修改后口播正文"
|
|
||||||
|
|
||||||
|
|
||||||
# ── 5) 端到端:每个 step 调用 loader 对应 prompt_type ──────────────
|
# ── 4) 端到端:image 走 V2、script 走 storyboard loader ─────────────
|
||||||
|
|
||||||
|
|
||||||
class TestEndToEndLoaderUsed:
|
class TestEndToEndLoaderUsed:
|
||||||
def test_each_step_calls_loader(self, job):
|
def test_image_v2_and_script_uses_storyboard(self, job):
|
||||||
from apps.worker.worker_app.tasks import viral_video as vv
|
from apps.worker.worker_app.tasks import viral_video as vv
|
||||||
from packages.application.viral_video import prompt_loader as pl
|
from packages.application.viral_video import prompt_loader as pl
|
||||||
|
from packages.shared.ai_router import ai_router as _router
|
||||||
|
|
||||||
called_types = []
|
called_types: list[str] = []
|
||||||
real_get = pl.get_template
|
real_get = pl.get_template
|
||||||
|
|
||||||
def spy_get(prompt_type, **kwargs):
|
def spy_get(prompt_type, **kwargs):
|
||||||
called_types.append(prompt_type)
|
called_types.append(prompt_type)
|
||||||
return real_get(prompt_type, **kwargs)
|
return real_get(prompt_type, **kwargs)
|
||||||
|
|
||||||
v2_product = {
|
v2_image = {
|
||||||
|
"type": "product",
|
||||||
"name": "lipstick",
|
"name": "lipstick",
|
||||||
"brand": "品牌X",
|
"brand": "品牌X",
|
||||||
"key_features": ["显白", "持久"],
|
"has_person": False,
|
||||||
|
"summary_markdown": "一支品牌X口红。",
|
||||||
}
|
}
|
||||||
with (
|
fake = _FakeClient()
|
||||||
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_product],
|
|
||||||
create=True,
|
|
||||||
),
|
|
||||||
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
|
|
||||||
):
|
|
||||||
# 1) image(V2 路径,不再经过 prompt_loader)
|
|
||||||
img_step = vv._step_image_analysis(job)
|
|
||||||
img_res = img_step["products"][0]
|
|
||||||
# 2) intent(走 loader image_analysis? 否——intent_parsing 模板)
|
|
||||||
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
|
|
||||||
|
|
||||||
# V2 图片分析不再调用 loader;意图解析调用 intent_parsing 模板
|
def _get(_key, variant=None):
|
||||||
|
return fake
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
# V2 图片分析不经过 prompt_loader;脚本生成调用 storyboard 模板
|
||||||
assert "image_analysis" not in called_types
|
assert "image_analysis" not in called_types
|
||||||
assert "intent_parsing" in called_types
|
assert "storyboard" in called_types
|
||||||
|
assert copy_res["voiceover_script"]
|
||||||
# script 和 review 单独验证(需要不同的 LLM 返回)
|
|
||||||
called_types_2 = []
|
|
||||||
|
|
||||||
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
|
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
|
||||||
|
|
||||||
pass_result = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
|
pass_result = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
|
||||||
job.intent_result = intent_res
|
with patch.object(Reviewer, "review", return_value=pass_result) as mock_review:
|
||||||
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)
|
review_res = vv._step_review(job, copy_res)
|
||||||
# review 步骤内部直接调用 Reviewer.review,该方法被 mock,因此 get_template 不会被调用;
|
assert mock_review.called
|
||||||
# 此处验证 Reviewer.review 被调用即可说明 review 步骤走通了。
|
|
||||||
assert mock_review.called, "_step_review 未调用 Reviewer.review"
|
|
||||||
assert isinstance(review_res, dict) and "passed" in review_res
|
assert isinstance(review_res, dict) and "passed" in review_res
|
||||||
|
|||||||
@@ -430,3 +430,56 @@ class TestWSInitialSnapshot:
|
|||||||
for p in patches:
|
for p in patches:
|
||||||
p.stop()
|
p.stop()
|
||||||
assert any(m["type"] == "forwarder_reached" for m in received)
|
assert any(m["type"] == "forwarder_reached" for m in received)
|
||||||
|
|
||||||
|
def test_image_analyzed_initial_snapshot_contains_image_analysis(self):
|
||||||
|
"""P0: image_analyzed 状态时初始快照必须带 image_analysis。"""
|
||||||
|
ia = {"images": [{"type": "product", "name": "X", "summary_markdown": "# X\nhello"}]}
|
||||||
|
job = _make_job(
|
||||||
|
status="image_analyzed",
|
||||||
|
user_id="user-a",
|
||||||
|
is_terminal=False,
|
||||||
|
image_analysis=ia,
|
||||||
|
copy_result=None,
|
||||||
|
generated_copy_text="",
|
||||||
|
storyboard=[],
|
||||||
|
)
|
||||||
|
received, _, _ = _run_ws_handshake(job=job)
|
||||||
|
assert received[0]["data"]["status"] == "image_analyzed"
|
||||||
|
assert received[0]["data"]["image_analysis"] == ia
|
||||||
|
|
||||||
|
def test_copy_generated_initial_snapshot_contains_copy_result(self):
|
||||||
|
"""P0: copy_generated 状态时初始快照必须带 copy_result/storyboard。"""
|
||||||
|
cr = {"shots": [{"time_range": "0-3s", "voiceover": "hi"}], "voiceover_script": "hi"}
|
||||||
|
sb = [{"order": 1, "text": "hi", "duration": 3.0}]
|
||||||
|
job = _make_job(
|
||||||
|
status="copy_generated",
|
||||||
|
user_id="user-a",
|
||||||
|
is_terminal=False,
|
||||||
|
image_analysis={"images": []},
|
||||||
|
copy_result=cr,
|
||||||
|
generated_copy_text="hi",
|
||||||
|
storyboard=sb,
|
||||||
|
)
|
||||||
|
received, _, _ = _run_ws_handshake(job=job)
|
||||||
|
data = received[0]["data"]
|
||||||
|
assert data["status"] == "copy_generated"
|
||||||
|
# _build_copy_result 会补 final_copy/suggested_copy/title 兜底
|
||||||
|
assert data["copy_result"]["shots"] == cr["shots"]
|
||||||
|
assert data["storyboard"] == sb
|
||||||
|
assert data["generated_copy_text"] == "hi"
|
||||||
|
assert data["image_analysis"] == {"images": []}
|
||||||
|
|
||||||
|
def test_initial_snapshot_without_business_fields_only_has_status(self):
|
||||||
|
"""running/pending 等中间态,无业务字段时不应塞空 dict/list。"""
|
||||||
|
job = _make_job(
|
||||||
|
status="running",
|
||||||
|
user_id="user-a",
|
||||||
|
is_terminal=False,
|
||||||
|
image_analysis=None,
|
||||||
|
copy_result=None,
|
||||||
|
generated_copy_text="",
|
||||||
|
storyboard=[],
|
||||||
|
)
|
||||||
|
received, _, _ = _run_ws_handshake(job=job)
|
||||||
|
data = received[0]["data"]
|
||||||
|
assert data == {"status": "running"}
|
||||||
|
|||||||
+123
-223
@@ -1,240 +1,137 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""vision v4 prompt / assembler 单元测试:
|
"""vision v8 叙述优先 assembler / prompt 单元测试。
|
||||||
|
|
||||||
- assembler 正确识别 v4 嵌套 schema 与旧扁平 schema
|
- assembler 输出仅 5 字段(type/name/brand/has_person/summary_markdown)
|
||||||
- v4 product/person/store/other 四类输出组装出下游必出字段
|
- images / 老 products 两种顶层键都能解析
|
||||||
- 旧扁平 schema 行为不变
|
- summary_markdown 正常时原样透传,不改写
|
||||||
- _prompt._resolve:DB 有 active prompt 时原样使用(不追加硬编码 schema);
|
- summary_markdown 缺失时才用一句话基础兜底
|
||||||
DB 无记录时回落到硬编码 JSON schema
|
- _prompt:DB 有 active 模板原样使用,无记录回落到 prompts.py 默认 v8
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import sys
|
||||||
import types
|
import types
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from worker_app.tasks.vision import _prompt, assembler
|
from worker_app.tasks.vision import _prompt, assembler
|
||||||
|
|
||||||
REQUIRED_KEYS = {
|
# packages 层依赖 datetime.UTC(Python 3.11+)。开发机若为旧版本,prompt 相关用例
|
||||||
"name",
|
# 在 CI(3.11)上正常执行,本地直接跳过,避免污染基线。
|
||||||
"brand",
|
_PY311 = sys.version_info >= (3, 11)
|
||||||
"category",
|
requires_packages = pytest.mark.skipif(not _PY311, reason="packages 需要 Python 3.11+")
|
||||||
"appearance",
|
|
||||||
"packaging",
|
REQUIRED_KEYS = {"type", "name", "brand", "has_person", "summary_markdown"}
|
||||||
"text_on_package",
|
|
||||||
"key_features",
|
|
||||||
"scene",
|
|
||||||
"mood",
|
|
||||||
"portrait_prompt",
|
|
||||||
"summary",
|
|
||||||
"_source",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
# ---------- schema 识别 ----------
|
# ---------- 正常 v8:叙述原样透传 ----------
|
||||||
|
|
||||||
|
|
||||||
def test_is_v4_schema_products_list() -> None:
|
def test_assemble_v8_store_passthrough() -> None:
|
||||||
assert assembler._is_v4_schema({"type": "product", "products": []})
|
md = "###店铺主体\n这是一家名为“御众堂”的线下门店内部,整体暖木色调……"
|
||||||
|
|
||||||
|
|
||||||
def test_is_v4_schema_type_only() -> None:
|
|
||||||
assert assembler._is_v4_schema({"type": "person"})
|
|
||||||
|
|
||||||
|
|
||||||
def test_is_v4_schema_people_dict() -> None:
|
|
||||||
assert assembler._is_v4_schema({"people": {"has_person": True}})
|
|
||||||
|
|
||||||
|
|
||||||
def test_is_not_v4_schema_flat() -> None:
|
|
||||||
assert not assembler._is_v4_schema({"has_person": True, "upper_wear": "T恤"})
|
|
||||||
|
|
||||||
|
|
||||||
# ---------- v4 product ----------
|
|
||||||
|
|
||||||
V4_PRODUCT: dict[str, Any] = {
|
|
||||||
"type": "product",
|
|
||||||
"scene": "白色背景产品图",
|
|
||||||
"mood": "清新专业",
|
|
||||||
"style": "商业产品摄影",
|
|
||||||
"colors": [{"hex": "#E60012", "name": "亮红色", "coverage": 0.6}],
|
|
||||||
"visible_text": [{"text": "OMO奥妙除菌除螨", "position": "瓶身正面"}],
|
|
||||||
"products": [
|
|
||||||
{
|
|
||||||
"product_name": "OMO奥妙除菌除螨洗衣液",
|
|
||||||
"brand": "OMO奥妙",
|
|
||||||
"category": "洗护",
|
|
||||||
"package_type": "瓶装",
|
|
||||||
"package_color": "亮红色瓶身",
|
|
||||||
"cap_type": "透明翻盖式按压瓶口",
|
|
||||||
"body_shape": "带侧面握持把手的竖款瓶身",
|
|
||||||
"label_design": "瓶身印十字盾牌图案",
|
|
||||||
"product_features": ["亮红色瓶装", "按压式瓶口", "十字盾牌标签"],
|
|
||||||
"key_selling_points": ["天然除菌除螨"],
|
|
||||||
"position": "main",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"has_person": False,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_v4_product_fields() -> None:
|
|
||||||
r = assembler.assemble_result(0, V4_PRODUCT, ["OMO奥妙"])
|
|
||||||
assert REQUIRED_KEYS <= set(r.keys())
|
|
||||||
assert r["name"] == "OMO奥妙除菌除螨洗衣液"
|
|
||||||
assert r["brand"] == "OMO奥妙"
|
|
||||||
assert r["category"] == "洗护"
|
|
||||||
assert "瓶装" in r["packaging"]
|
|
||||||
assert isinstance(r["key_features"], list) and r["key_features"]
|
|
||||||
assert any("除菌" in str(t) for t in r["text_on_package"])
|
|
||||||
assert len(r["portrait_prompt"]) >= 10
|
|
||||||
assert r["_source"] == "v2_fast_json_v4"
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_v4_product_multi_selects_main() -> None:
|
|
||||||
fj = {
|
fj = {
|
||||||
"type": "product",
|
"images": [
|
||||||
"products": [
|
{
|
||||||
{"product_name": "次要商品", "brand": "B"},
|
"type": "store",
|
||||||
{"product_name": "主商品", "brand": "A", "position": "main"},
|
"name": "御众堂门店",
|
||||||
],
|
"brand": "御众堂",
|
||||||
}
|
"has_person": False,
|
||||||
r = assembler.assemble_result(1, fj, [])
|
"summary_markdown": md,
|
||||||
assert r["name"] == "主商品"
|
}
|
||||||
|
]
|
||||||
|
|
||||||
# ---------- v4 person ----------
|
|
||||||
|
|
||||||
V4_PERSON: dict[str, Any] = {
|
|
||||||
"type": "person",
|
|
||||||
"scene": "户外街拍",
|
|
||||||
"mood": "自信",
|
|
||||||
"style": "街拍",
|
|
||||||
"colors": [],
|
|
||||||
"visible_text": [],
|
|
||||||
"has_person": True,
|
|
||||||
"gender": "女",
|
|
||||||
"age_range": "青年",
|
|
||||||
"upper_wear": "白色V领短袖T恤",
|
|
||||||
"upper_color": "白色",
|
|
||||||
"lower_wear": "黑色高腰阔腿裤",
|
|
||||||
"lower_color": "黑色",
|
|
||||||
"dress_color": None,
|
|
||||||
"accessories": ["银色项链"],
|
|
||||||
"hairstyle": "黑色长直发",
|
|
||||||
"expression": "自信",
|
|
||||||
"pose": "侧身站立",
|
|
||||||
"outfit_style": "休闲日常",
|
|
||||||
"portrait_prompt": (
|
|
||||||
"一位年轻女性,身穿白色V领短袖T恤、黑色高腰阔腿裤,佩戴银色项链,"
|
|
||||||
"黑色长直发,神情自信,侧身站立,休闲日常风格,城市街拍场景"
|
|
||||||
),
|
|
||||||
"products": [],
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_v4_person() -> None:
|
|
||||||
r = assembler.assemble_result(0, V4_PERSON, [])
|
|
||||||
assert REQUIRED_KEYS <= set(r.keys())
|
|
||||||
assert r["category"] == "人物穿搭"
|
|
||||||
assert r["_source"] == "v2_fast_json_v5"
|
|
||||||
assert "T恤" in r["name"]
|
|
||||||
assert "年轻女性" in r["portrait_prompt"]
|
|
||||||
assert "项链" in r["portrait_prompt"]
|
|
||||||
assert isinstance(r["key_features"], list) and len(r["key_features"]) <= 8
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_v4_person_people_nested() -> None:
|
|
||||||
fj = {"type": "person", "people": {**V4_PERSON, "has_person": True}}
|
|
||||||
r = assembler.assemble_result(0, fj, [])
|
|
||||||
assert r["category"] == "人物穿搭"
|
|
||||||
assert "年轻女性" in r["portrait_prompt"]
|
|
||||||
|
|
||||||
|
|
||||||
# ---------- v4 store ----------
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_v4_store() -> None:
|
|
||||||
fj = {
|
|
||||||
"type": "store",
|
|
||||||
"scene": "便利店内部",
|
|
||||||
"mood": "日常便民",
|
|
||||||
"style": "门店实拍",
|
|
||||||
"store_type": "社区便利店",
|
|
||||||
"store_layout": "纵深货架布局",
|
|
||||||
"brand_signage": "全家FamilyMart",
|
|
||||||
"visual_elements": ["红白主色调", "促销海报"],
|
|
||||||
"product_categories_visible": ["饮料", "零食"],
|
|
||||||
"promotion_elements": ["第二件半价海报"],
|
|
||||||
"atmosphere": "亲民生活化",
|
|
||||||
"has_person": False,
|
|
||||||
}
|
}
|
||||||
r = assembler.assemble_result(0, fj, [])
|
r = assembler.assemble_result(0, fj, [])
|
||||||
assert REQUIRED_KEYS <= set(r.keys())
|
assert REQUIRED_KEYS <= set(r.keys())
|
||||||
assert r["name"] == "社区便利店"
|
assert r["type"] == "store"
|
||||||
assert r["brand"] == "全家FamilyMart"
|
assert r["name"] == "御众堂门店"
|
||||||
assert r["category"] == "门店场景"
|
assert r["brand"] == "御众堂"
|
||||||
assert any("饮料" in str(f) for f in r["key_features"])
|
assert r["has_person"] is False
|
||||||
assert "门店实拍" in r["portrait_prompt"]
|
assert r["summary_markdown"] == md
|
||||||
|
assert "_source" not in r
|
||||||
|
|
||||||
|
|
||||||
# ---------- v4 other ----------
|
def test_assemble_v8_product() -> None:
|
||||||
|
md = "这是一瓶洗衣液,亮红色瓶身配白色按压泵头,瓶身正面印着品牌标识……"
|
||||||
|
|
||||||
def test_assemble_v4_other() -> None:
|
|
||||||
fj = {"type": "other", "description": "海边日落风景", "scene": "海边", "mood": "宁静"}
|
|
||||||
r = assembler.assemble_result(0, fj, [])
|
|
||||||
assert REQUIRED_KEYS <= set(r.keys())
|
|
||||||
assert r["name"] == "海边日落风景"
|
|
||||||
assert r["category"] == "非产品图"
|
|
||||||
|
|
||||||
|
|
||||||
# ---------- 旧扁平 schema 兼容 ----------
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_old_flat_person() -> None:
|
|
||||||
fj = {
|
fj = {
|
||||||
"has_person": True,
|
"images": [{"type": "product", "name": "洗衣液", "brand": "OMO", "has_person": False, "summary_markdown": md}]
|
||||||
"gender": "男",
|
}
|
||||||
"age_range": "中年",
|
r = assembler.assemble_result(0, fj, ["OMO"])
|
||||||
"upper_wear": "西装",
|
assert r["type"] == "product"
|
||||||
"upper_color": "深灰色",
|
assert r["summary_markdown"] == md
|
||||||
"lower_wear": "西裤",
|
|
||||||
"lower_color": "黑色",
|
|
||||||
"accessories": ["手表"],
|
def test_assemble_v8_person() -> None:
|
||||||
"hairstyle": "短发",
|
md = "画面里是一位年轻女性,穿白色T恤、黑色阔腿裤,神情自信……"
|
||||||
"expression": "严肃",
|
fj = {"images": [{"type": "person", "name": "年轻女性", "brand": "", "has_person": True, "summary_markdown": md}]}
|
||||||
"scene": "办公室",
|
r = assembler.assemble_result(0, fj, [])
|
||||||
"style": "商务",
|
assert r["type"] == "person"
|
||||||
"mood": "专业",
|
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, [])
|
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 REQUIRED_KEYS <= set(r.keys())
|
||||||
assert "中年男性" in r["portrait_prompt"]
|
assert r["summary_markdown"]
|
||||||
assert r["_source"] == "v2_fast_json"
|
assert "御众堂" in r["summary_markdown"]
|
||||||
|
assert r.get("_source") == "summary_missing"
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_old_flat_product() -> None:
|
def test_assemble_invalid_type_defaults_scene() -> None:
|
||||||
fj = {
|
fj = {"images": [{"type": "weird", "name": "x", "summary_markdown": ""}]}
|
||||||
"has_person": False,
|
r = assembler.assemble_result(0, fj, [])
|
||||||
"product_name": "口红",
|
assert r["type"] == "scene"
|
||||||
"brand": "Dior",
|
assert r["summary_markdown"] # basic fallback
|
||||||
"category": "美妆",
|
|
||||||
"colors": ["红色"],
|
|
||||||
"scene": "通用",
|
def test_assemble_empty_fast_json_uses_ocr_hint() -> None:
|
||||||
"style": "商业",
|
r = assembler.assemble_result(0, {}, ["御众堂"])
|
||||||
"mood": "高级",
|
assert REQUIRED_KEYS <= set(r.keys())
|
||||||
}
|
assert "御众堂" in r["name"]
|
||||||
r = assembler.assemble_result(0, fj, ["Dior"])
|
assert r.get("_source") == "empty_fast_json"
|
||||||
assert r["name"] == "口红"
|
|
||||||
assert r["brand"] == "Dior"
|
|
||||||
assert r["text_on_package"] == ["Dior"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_assemble_none_input() -> None:
|
def test_assemble_none_input() -> None:
|
||||||
r = assembler.assemble_result(0, None, [])
|
r = assembler.assemble_result(0, None, [])
|
||||||
assert REQUIRED_KEYS <= set(r.keys())
|
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 解析 ----------
|
# ---------- _prompt 解析 ----------
|
||||||
@@ -247,45 +144,48 @@ def _clear_prompt_cache() -> Any:
|
|||||||
_prompt.invalidate_cache()
|
_prompt.invalidate_cache()
|
||||||
|
|
||||||
|
|
||||||
def _fake_tpl(system_prompt: str = "v4 system prompt 只返回JSON") -> Any:
|
def _fake_tpl(system_prompt: str = "DB_V8_PROMPT_XYZ") -> Any:
|
||||||
return types.SimpleNamespace(
|
return types.SimpleNamespace(
|
||||||
system_prompt=system_prompt,
|
system_prompt=system_prompt,
|
||||||
user_prompt_template="分析 {image_count} 张图",
|
user_prompt_template="地址:{image_url},OCR:{ocr_text}",
|
||||||
version=4,
|
version=8,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_resolve_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
@requires_packages
|
||||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PROMPT_XYZ"))
|
def test_resolve_uses_db_prompt(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl())
|
||||||
assert sys_prompt == "DB_V4_PROMPT_XYZ"
|
sys_prompt, user_prompt = _prompt.resolve_fast_prompt("http://img", "御众堂")
|
||||||
assert "DB_V4_PROMPT_XYZ" not in _prompt._FAST_JSON_APPEND # sanity: 旧append是另一段文本
|
assert sys_prompt == "DB_V8_PROMPT_XYZ"
|
||||||
assert "分析 1 张图" in user_prompt
|
assert "http://img" in user_prompt
|
||||||
|
assert "御众堂" in user_prompt
|
||||||
|
|
||||||
|
|
||||||
def test_resolve_pro_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
@requires_packages
|
||||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PRO_PROMPT"))
|
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()
|
sys_prompt, _ = _prompt.resolve_pro_prompt()
|
||||||
assert sys_prompt == "DB_V4_PRO_PROMPT"
|
assert sys_prompt == "DB_PRO_PROMPT"
|
||||||
assert "【输出格式要求】" not in sys_prompt
|
|
||||||
|
|
||||||
|
|
||||||
def test_resolve_falls_back_when_no_db(monkeypatch: pytest.MonkeyPatch) -> None:
|
@requires_packages
|
||||||
|
def test_resolve_falls_back_to_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
|
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
|
||||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
default = _prompt._default_template()
|
||||||
assert sys_prompt == _prompt._FAST_JSON_SCHEMA
|
sys_prompt, _ = _prompt.resolve_fast_prompt()
|
||||||
assert user_prompt == _prompt.DEFAULT_FAST_USER
|
assert sys_prompt == default["system_prompt"]
|
||||||
|
|
||||||
|
|
||||||
|
@requires_packages
|
||||||
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
|
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
calls = {"n": 0}
|
calls = {"n": 0}
|
||||||
|
|
||||||
def _load() -> Any:
|
def _load() -> Any:
|
||||||
calls["n"] += 1
|
calls["n"] += 1
|
||||||
return _fake_tpl("CACHED_PROMPT")
|
return _fake_tpl("CACHED")
|
||||||
|
|
||||||
monkeypatch.setattr(_prompt, "_load_db_template", _load)
|
monkeypatch.setattr(_prompt, "_load_db_template", _load)
|
||||||
s1, _ = _prompt.resolve_fast_prompt()
|
s1, _ = _prompt.resolve_fast_prompt()
|
||||||
s2, _ = _prompt.resolve_fast_prompt()
|
s2, _ = _prompt.resolve_fast_prompt()
|
||||||
assert s1 == s2 == "CACHED_PROMPT"
|
assert s1 == s2 == "CACHED"
|
||||||
assert calls["n"] == 1
|
assert calls["n"] == 1
|
||||||
|
|||||||
Executable
+49
@@ -0,0 +1,49 @@
|
|||||||
|
"""xml_parser CDATA 剥离单元测试。"""
|
||||||
|
|
||||||
|
from packages.application.viral_video.xml_parser import find_all, text_of
|
||||||
|
|
||||||
|
XML = """<script>
|
||||||
|
<copy_display_markdown><。
|
||||||
|
|
||||||
|
第二行,保留换行。]]></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="缺省") == "缺省"
|
||||||
Reference in New Issue
Block a user