Compare commits
56 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3e6f87a8b5 | |||
| 5da46945fa | |||
| bff20b03d7 | |||
| 63c8496fa8 | |||
| 99ba9b7c58 | |||
| ef82192679 | |||
| 0bd4123ae5 | |||
| 9139c697b0 | |||
| b60de7202a | |||
| f1bd816449 | |||
| 249b70e53e | |||
| dbc6db02e0 | |||
| ec28699806 | |||
| afc7a37d17 | |||
| 731ad37217 | |||
| c551ecbcc5 | |||
| 873fa89305 | |||
| 225406cc2e | |||
| 4ccb395dd1 | |||
| f5162455c2 | |||
| 5b4b844f1a | |||
| 587eefd8fa | |||
| cbc9fc885b | |||
| 1fb156745d | |||
| 607e989b95 | |||
| c277e87ad3 | |||
| c25269051e | |||
| 54842c1387 | |||
| d1c137af8a | |||
| f0eaf4c31f | |||
| a8a20d6f87 | |||
| 0fa2397f33 | |||
| 52e757f227 | |||
| 8d4ad83212 | |||
| 92bacf3853 | |||
| 8f949ae5d8 | |||
| ed118d4444 | |||
| 5a5833166a | |||
| b3e97e6bae | |||
| 6b20a568ec | |||
| b429117e97 | |||
| 4000a81f91 | |||
| b32a14ada6 | |||
| d0ff504a5f | |||
| 498fe7b38c | |||
| 4266f4a6a2 | |||
| a50cbc995f | |||
| cfca2443c4 | |||
| 6cb9cf0b27 | |||
| b8c6091a11 | |||
| 3b9a3dd426 | |||
| 617c40e1d4 | |||
| 2597962528 | |||
| 8c56694599 | |||
| cd618c3f29 | |||
| 015fd2c381 |
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,222 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""099: AI 模型路由层 seed — 补齐缺失模型和能力配置.
|
||||
|
||||
幂等:所有 INSERT 先检查存在性。
|
||||
- ai_models: 补齐 qwen3.7-plus, seedream, seedance, embedding, wan3.0 等
|
||||
- ai_capability_configs: 补齐 image_generation, video_generation, embedding
|
||||
- 更新已有 capability 的 lite_model_id
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "099_ai_model_router_seed"
|
||||
down_revision = "098_viral_video_image_analysis_v5"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# CI 环境下 ai_models 表可能尚未创建(由 ORM 自动建表,非 migration)
|
||||
# 如果表不存在则跳过 seed,由应用启动时 ORM 建表后首次访问时生效
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
# ai_models 表不存在,跳过所有 seed(CI 环境)
|
||||
return
|
||||
|
||||
# ── 1. 补齐 ai_models 缺失记录 ────────────────────────────────────────────
|
||||
existing_models = {
|
||||
row[0]
|
||||
for row in conn.execute(
|
||||
sa.text("SELECT model_key FROM ai_models WHERE deleted_at IS NULL")
|
||||
).fetchall()
|
||||
}
|
||||
|
||||
# 从已有 active 记录获取 API key(复用,不硬编码)
|
||||
dashscope_key_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT api_key FROM ai_models WHERE provider='dashscope' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
dashscope_key = dashscope_key_row[0] if dashscope_key_row else ""
|
||||
|
||||
volcengine_key_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT api_key FROM ai_models WHERE provider='volcengine' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
volcengine_key = volcengine_key_row[0] if volcengine_key_row else ""
|
||||
|
||||
new_models = [
|
||||
{
|
||||
"model_key": "qwen3.7-plus",
|
||||
"name": "通义千问3.7 Plus(VLM 兜底)",
|
||||
"provider": "dashscope",
|
||||
"api_key": dashscope_key,
|
||||
"api_base": "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
"description": "阿里云百炼 Qwen3.7 Plus 多模态模型,用于 VLM 兜底分析",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seedream-5-0-flash-260915",
|
||||
"name": "Seedream 5.0 Flash(图片生成)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎 Seedream 5.0 Flash 文生图模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seedance-2-5-260628",
|
||||
"name": "Seedance 2.5(视频生成)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎 Seedance 2.5 图/文生视频模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-embedding-vision-251215",
|
||||
"name": "豆包多模态向量嵌入",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎豆包多模态向量嵌入模型",
|
||||
},
|
||||
{
|
||||
"model_key": "wan3.0-video",
|
||||
"name": "Wan 3.0 视频生成",
|
||||
"provider": "dashscope",
|
||||
"api_key": dashscope_key,
|
||||
"api_base": "https://dashscope.aliyuncs.com/api/v1",
|
||||
"description": "阿里云百炼 Wan 3.0 视频生成模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seed-2-1-pro-260915",
|
||||
"name": "豆包 Seed 2.1 Pro(高精度推理)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎豆包 Seed 2.1 Pro 深度思考+多模态",
|
||||
},
|
||||
]
|
||||
|
||||
for m in new_models:
|
||||
if m["model_key"] not in existing_models:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base, description, status, is_default, usage_today, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :name, :provider, :model_key, :api_key, :api_base, :description, 'active', false, 0, now(), now())
|
||||
"""
|
||||
),
|
||||
m,
|
||||
)
|
||||
|
||||
# ── 2. 补齐 ai_capability_configs 缺失项 ──────────────────────────────────
|
||||
cap_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not cap_table_check:
|
||||
return
|
||||
|
||||
existing_caps = {
|
||||
row[0]
|
||||
for row in conn.execute(
|
||||
sa.text("SELECT capability_key FROM ai_capability_configs")
|
||||
).fetchall()
|
||||
}
|
||||
|
||||
def _get_model_id(model_key: str) -> str | None:
|
||||
row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = :key AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
),
|
||||
{"key": model_key},
|
||||
).first()
|
||||
return row[0] if row else None
|
||||
|
||||
# image_generation
|
||||
if "image_generation" not in existing_caps:
|
||||
mid = _get_model_id("doubao-seedream-5-0-flash-260915")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 60, 1, 2, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "image_generation",
|
||||
"cn": "图片生成(Seedream)",
|
||||
"pm": mid,
|
||||
"ep": json.dumps({"size": "1K"}),
|
||||
},
|
||||
)
|
||||
|
||||
# video_generation
|
||||
if "video_generation" not in existing_caps:
|
||||
mid = _get_model_id("doubao-seedance-2-5-260628")
|
||||
fb_mid = _get_model_id("wan3.0-video")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, fallback_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, :fm, 600, 1, 1, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "video_generation",
|
||||
"cn": "视频生成(Seedance/Wan)",
|
||||
"pm": mid,
|
||||
"fm": fb_mid,
|
||||
"ep": json.dumps({}),
|
||||
},
|
||||
)
|
||||
|
||||
# embedding
|
||||
if "embedding" not in existing_caps:
|
||||
mid = _get_model_id("doubao-embedding-vision-251215")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 30, 2, 5, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "embedding",
|
||||
"cn": "向量嵌入",
|
||||
"pm": mid,
|
||||
"ep": json.dumps({}),
|
||||
},
|
||||
)
|
||||
|
||||
# ── 3. 更新 image_analysis 的 lite_model_id ─────────────────────────────
|
||||
lite_model_id = _get_model_id("qwen3.8-flash")
|
||||
if lite_model_id:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET lite_model_id = :lite WHERE capability_key = 'image_analysis' AND lite_model_id IS NULL"
|
||||
),
|
||||
{"lite": lite_model_id},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
# 安全检查表是否存在
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
return
|
||||
conn.execute(
|
||||
sa.text("DELETE FROM ai_capability_configs WHERE capability_key IN ('image_generation', 'video_generation', 'embedding')")
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"DELETE FROM ai_models WHERE model_key IN ('qwen3.7-plus', 'doubao-seedream-5-0-flash-260915', 'doubao-seedance-2-5-260628', 'doubao-embedding-vision-251215', 'wan3.0-video', 'doubao-seed-2-1-pro-260915') AND deleted_at IS NULL"
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,107 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""100: 修正已有 capability 的模型绑定.
|
||||
|
||||
幂等:仅当 primary_model_id 当前绑定到旧模型 (doubao-seed-1-6) 时才更新,
|
||||
避免覆盖用户在后台的自定义配置。
|
||||
|
||||
- 更新 5 个 LLM capability (intent_parsing, copy_fusion, storyboard, copy_review, asset_classify)
|
||||
的 primary_model_id 从 doubao-seed-1-6 改为 doubao-seed-2-1-pro-260915
|
||||
- 更新 image_analysis 的 primary/lite/fallback 模型绑定
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "100_fix_capability_model_bindings"
|
||||
down_revision = "099_ai_model_router_seed"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# Check tables exist
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
return
|
||||
|
||||
config_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not config_table_check:
|
||||
return
|
||||
|
||||
# Look up model IDs by model_key (not hardcoded UUIDs)
|
||||
pro_model_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'doubao-seed-2-1-pro-260915' AND deleted_at IS NULL LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
if not pro_model_row:
|
||||
return
|
||||
pro_model_id = pro_model_row[0]
|
||||
|
||||
old_model_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'doubao-seed-1-6-250615' LIMIT 1")
|
||||
).first()
|
||||
old_model_id = old_model_row[0] if old_model_row else None
|
||||
|
||||
llm_capabilities = [
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"copy_review",
|
||||
"asset_classify",
|
||||
]
|
||||
|
||||
for cap_key in llm_capabilities:
|
||||
if old_model_id:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET primary_model_id = :new_id, updated_at = NOW() "
|
||||
"WHERE capability_key = :cap_key AND primary_model_id = :old_id"
|
||||
),
|
||||
{"new_id": pro_model_id, "old_id": old_model_id, "cap_key": cap_key},
|
||||
)
|
||||
|
||||
# Update image_analysis
|
||||
qwen38_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.8-flash' AND deleted_at IS NULL LIMIT 1")
|
||||
).first()
|
||||
qwen37_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' AND deleted_at IS NULL LIMIT 1")
|
||||
).first()
|
||||
|
||||
if qwen38_row and qwen37_row:
|
||||
qwen38_id = qwen38_row[0]
|
||||
qwen37_id = qwen37_row[0]
|
||||
|
||||
current_ia = conn.execute(
|
||||
sa.text(
|
||||
"SELECT primary_model_id, lite_model_id, fallback_model_id "
|
||||
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
).first()
|
||||
|
||||
if current_ia:
|
||||
current_primary, current_lite, current_fallback = current_ia
|
||||
updates = {}
|
||||
if current_primary != qwen38_id:
|
||||
updates["primary_model_id"] = qwen38_id
|
||||
if current_lite != qwen38_id:
|
||||
updates["lite_model_id"] = qwen38_id
|
||||
if current_fallback != qwen37_id:
|
||||
updates["fallback_model_id"] = qwen37_id
|
||||
|
||||
if updates:
|
||||
set_clause = ", ".join([f"{k} = :{k}" for k in updates.keys()])
|
||||
set_clause += ", updated_at = NOW()"
|
||||
updates["cap_key"] = "image_analysis"
|
||||
conn.execute(
|
||||
sa.text(f"UPDATE ai_capability_configs SET {set_clause} WHERE capability_key = :cap_key"),
|
||||
updates,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -0,0 +1,153 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""101: 补齐 qwen-vl-plus 视觉模型并修正 image_analysis 绑定与 max_tokens.
|
||||
|
||||
背景:
|
||||
- qwen-vl-plus 做图片识别时返回 JSON 约 500-600 tokens,旧硬编码
|
||||
max_tokens=350 导致 JSON 被截断、解析失败返回"未识别"。
|
||||
- 代码侧已移除硬编码,改由 capability 的 DB 配置决定 max_tokens。
|
||||
|
||||
幂等:
|
||||
- qwen-vl-plus 已存在则不插入;
|
||||
- 仅当 image_analysis 当前 primary_model 不是 qwen-vl-plus 时才更新绑定,
|
||||
避免覆盖后台手动配置。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "101_qwen_vl_plus_and_max_tokens"
|
||||
down_revision = "100_fix_capability_model_bindings"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
models_table = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not models_table:
|
||||
return
|
||||
|
||||
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not caps_table:
|
||||
return
|
||||
|
||||
# ── c. 补全其他 capability 的 max_tokens 默认值(幂等)──────────────────
|
||||
# 放在 image_analysis 特定逻辑之前,确保任何分支 return 都不会跳过本段。
|
||||
# 仅在当前值为 NULL 或过小 (<100) 时更新,不覆盖已有合理配置。
|
||||
# embedding / tts / voice_clone 不走 chat 接口,无需设置。
|
||||
default_max_tokens = {
|
||||
"intent_parsing": 500,
|
||||
"copy_fusion": 2500,
|
||||
"storyboard": 4000,
|
||||
"copy_review": 1000,
|
||||
"asset_classify": 500,
|
||||
"image_generation": 500,
|
||||
"video_generation": 500,
|
||||
}
|
||||
for cap_key, mt in default_max_tokens.items():
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs "
|
||||
"SET max_tokens = :mt, updated_at = now() "
|
||||
"WHERE capability_key = :key "
|
||||
"AND (max_tokens IS NULL OR max_tokens < 100)"
|
||||
),
|
||||
{"mt": mt, "key": cap_key},
|
||||
)
|
||||
|
||||
# ── a. 确保 qwen-vl-plus 模型存在 ────────────────────────────────────────
|
||||
conn.execute(sa.text("""
|
||||
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base,
|
||||
description, status, is_default, usage_today,
|
||||
created_at, updated_at)
|
||||
SELECT gen_random_uuid()::text,
|
||||
'通义千问VL Plus',
|
||||
'dashscope',
|
||||
'qwen-vl-plus',
|
||||
COALESCE(
|
||||
(SELECT api_key FROM ai_models
|
||||
WHERE provider = 'dashscope' AND deleted_at IS NULL
|
||||
AND api_key IS NOT NULL AND api_key != ''
|
||||
LIMIT 1),
|
||||
''
|
||||
),
|
||||
'https://dashscope.aliyuncs.com/compatible-mode/v1',
|
||||
'阿里云视觉理解模型(图片识别/分析)',
|
||||
'active', false, 0, now(), now()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM ai_models
|
||||
WHERE model_key = 'qwen-vl-plus' AND deleted_at IS NULL
|
||||
)
|
||||
"""))
|
||||
|
||||
qwen_vl_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'qwen-vl-plus' "
|
||||
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
if not qwen_vl_row:
|
||||
return
|
||||
qwen_vl_id = qwen_vl_row[0]
|
||||
|
||||
qwen37_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' "
|
||||
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
qwen37_id = qwen37_row[0] if qwen37_row else None
|
||||
|
||||
# ── b. 仅当当前 primary 不是 qwen-vl-plus 时修正绑定与 max_tokens ───────
|
||||
current = conn.execute(
|
||||
sa.text(
|
||||
"SELECT primary_model_id, lite_model_id, fallback_model_id, max_tokens "
|
||||
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
).first()
|
||||
|
||||
if current is None:
|
||||
# capability 不存在则创建
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO ai_capability_configs
|
||||
(id, capability_key, capability_name, primary_model_id,
|
||||
lite_model_id, fallback_model_id, timeout_seconds,
|
||||
max_retries, max_tokens, concurrency, extra_params,
|
||||
is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, 'image_analysis', '图片分析',
|
||||
:primary, :primary, :fallback, 30, 1, 1000, 2,
|
||||
'{}'::jsonb, true, now(), now())
|
||||
"""),
|
||||
{"primary": qwen_vl_id, "fallback": qwen37_id},
|
||||
)
|
||||
return
|
||||
|
||||
current_primary = current[0]
|
||||
if current_primary == qwen_vl_id:
|
||||
# 已经绑定 qwen-vl-plus:视为后台/数据迁移已处理,不覆盖任何配置
|
||||
return
|
||||
|
||||
set_parts = [
|
||||
"primary_model_id = :vl_id",
|
||||
"lite_model_id = :vl_id",
|
||||
"max_tokens = 1000",
|
||||
"updated_at = now()",
|
||||
]
|
||||
params: dict = {"vl_id": qwen_vl_id}
|
||||
if qwen37_id is not None:
|
||||
set_parts.insert(2, "fallback_model_id = :qwen37_id")
|
||||
params["qwen37_id"] = qwen37_id
|
||||
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET " + ", ".join(set_parts) + " WHERE capability_key = 'image_analysis'"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -0,0 +1,36 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""102: image_analysis max_tokens 1200 -> 1500.
|
||||
|
||||
v6 prompt 更长、字段更多,旧 max_tokens 容易截断 JSON。
|
||||
仅在 image_analysis 当前 max_tokens < 1500 时更新(幂等,不覆盖后台已调到 >=1500 的配置)。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "102_image_analysis_max_tokens_1500"
|
||||
down_revision = "101_qwen_vl_plus_and_max_tokens"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not caps_table:
|
||||
return
|
||||
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs "
|
||||
"SET max_tokens = 1500, updated_at = now() "
|
||||
"WHERE capability_key = 'image_analysis' "
|
||||
"AND (max_tokens IS NULL OR max_tokens < 1500)"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -0,0 +1,194 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""image_analysis v7 prompt + max_tokens 3000 + max_retries 3
|
||||
|
||||
Revision ID: 103_v7_prompt_and_tokens_3000
|
||||
Revises: 102_image_analysis_max_tokens_1500
|
||||
Create Date: 2026-10-07
|
||||
|
||||
变更:
|
||||
1. 插入v7精简prompt(~1KB,v6 ~4.5KB,删除few-shot/冗长规则,减少输出token占用),设为active
|
||||
2. v6停用(is_active=False),保留历史
|
||||
3. image_analysis capability: max_tokens 1500→3000,max_retries 1→3
|
||||
|
||||
ai_capability_configs 由应用 create_all 创建,全新 alembic-only 库可能不存在,
|
||||
故第3步做 to_regclass 守卫(同 102)。
|
||||
"""
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "103_v7_prompt_and_tokens_3000"
|
||||
down_revision = "102_image_analysis_max_tokens_1500"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
V7_SYSTEM = """# 角色
|
||||
你是一位专业的图片分析师,擅长准确识别图片中的场景、人物、物体、文字、氛围。
|
||||
|
||||
# 任务
|
||||
对用户上传的图片逐张分析,描述你看到的内容,输出JSON格式。
|
||||
|
||||
## 技能
|
||||
|
||||
### 技能1:判断图片类型
|
||||
判断图片属于哪种类型,type字段填对应的英文值:
|
||||
- 商品图(product):单个或多个商品、产品包装
|
||||
- 门店场景图(store):店铺内部、门头招牌、货架陈列
|
||||
- 人物图(person):人物形象、穿搭造型、肖像照片
|
||||
- 风景图(scene):风景、动物、美食、街景
|
||||
- 其他(other):以上都不是
|
||||
|
||||
### 技能2:描述通用信息
|
||||
不管什么图都要描述:
|
||||
- type:图片类型,填product/store/person/scene/other其中一个
|
||||
- scene:一句话描述场景,例如"理疗养生店内部,摆着多张理疗床和产品货架"
|
||||
- mood:整体氛围,2-4个词,例如"整洁专业"、"热闹温馨"
|
||||
- colors:主要颜色,最多5个,写具体颜色名(亮红色/米白色/深蓝色,不写笼统的红色蓝色)
|
||||
- visible_text:图片里看到的文字,说明什么字、在什么位置,最多5条;没看到就空数组
|
||||
- lighting:光线情况,例如"明亮柔光"、"自然光"、"室内暖黄灯"
|
||||
- composition:怎么拍的,例如"居中特写"、"中景平视"、"俯拍"
|
||||
- has_person:有没有人,true或false
|
||||
|
||||
### 技能3:描述门店场景
|
||||
如果是门店场景图(type="store"),还要描述:
|
||||
- store_type:什么类型的店,例如"养生馆"、"便利店"、"餐饮店"、"母婴店"
|
||||
- brand_signage:招牌上写了什么字、有什么品牌标识
|
||||
- visual_elements:看到哪些显眼的东西(招牌样式、灯光、货架、商品陈列、海报、收银台等),最多8个
|
||||
- product_categories:看到哪些品类的商品,例如"饮料零食"、"养生产品"
|
||||
- promotion_elements:有没有促销活动(打折海报、满减吊旗等),没有就空数组
|
||||
- atmosphere:店内什么氛围,例如"亲民生活化"、"老字号专业感"
|
||||
- cleanliness:店内干净程度,例如"干净整洁"、"货架整齐"
|
||||
- 看到顾客或店员要描述他们在做什么,has_person填true
|
||||
|
||||
### 技能4:描述商品
|
||||
如果是商品图(type="product"),逐个商品描述:
|
||||
- product_name:商品名称,尽量具体,例如"OMO奥妙除菌除螨洗衣液";看不出来填null
|
||||
- brand:什么牌子,看不出来填null
|
||||
- category:类目,从以下选一个:服饰鞋包/美妆/数码/食品/家居清洁/母婴/配饰/其他
|
||||
- package_type:什么包装,例如"瓶装"、"盒装"、"罐装"、"袋装"、"多瓶装"
|
||||
- package_color:包装主要颜色,写具体色(亮红色不写红色)
|
||||
- body_shape:瓶身或包装形状,例如"圆润胖瓶"、"竖款带把手瓶身"
|
||||
- label_design:标签设计,例如"红色标签印白色品牌logo"
|
||||
- key_text_on_package:包装上最显眼的文字(品牌名、功能词、卖点词),最多5个
|
||||
- product_features:包装特征,3-6个短语,包含颜色、瓶盖、形状、标签图案
|
||||
- key_selling_points:核心卖点,1-3个短语
|
||||
|
||||
### 技能5:描述人物
|
||||
如果是人物图(type="person"),描述:
|
||||
- person_count:几个人
|
||||
- gender:性别(男/女/无法判断)
|
||||
- age_range:年龄段(儿童/青少年/青年/中年/老年/无法判断)
|
||||
- outfit_style:穿搭风格,例如"休闲日常"、"通勤商务"、"街头潮流"
|
||||
- upper_wear:上装(颜色+款式+材质),穿裙装不填
|
||||
- lower_wear:下装(颜色+款式+版型),穿裙装不填
|
||||
- dress_wear:裙装描述,穿上下装不填
|
||||
- outerwear:外套
|
||||
- shoes:鞋子
|
||||
- bag:包袋,没有填null
|
||||
- accessories:配饰(眼镜/帽子/项链/耳环/手表/手链/围巾/腰带等),没有填空数组
|
||||
- hairstyle:发型
|
||||
- makeup:妆容,男生或看不出填null
|
||||
- expression:表情,例如"微笑看镜头"、"冷酷无表情"
|
||||
- pose:姿势动作,例如"身直立正对镜头"、"单手撩发"
|
||||
- body_type:身材,例如"纤细苗条"、"高挑身材"、"丰满匀称"
|
||||
- portrait_prompt:80-150字详细描述人物形象(后面用来AI生成肖像图),要写清年龄段、穿搭完整细节、发型发色、妆容、表情、姿势、场景、光线、风格感觉,语言要有画面感
|
||||
|
||||
### 技能6:描述风景
|
||||
如果是风景图(type="scene"),描述:
|
||||
- scene_type:什么场景,例如"自然风景"、"城市街景"、"动物"、"美食"
|
||||
- main_subject:画面主体是什么
|
||||
- key_elements:关键元素,最多8个
|
||||
- environment_objects:周围环境物体,最多8个
|
||||
- atmosphere:整体氛围,例如"秋日慵懒氛围感"、"清新自然氧气感"
|
||||
- 有人物就描述人物特征
|
||||
|
||||
## 限制
|
||||
- 只输出JSON,不要任何解释文字,不要markdown代码块包裹,不要写"好的""以下是分析结果"这种废话
|
||||
- 颜色写具体色调(亮红色/米白色/深蓝色/翠绿色),不写笼统词汇
|
||||
- 瓶身、包装、招牌上的文字尽量识别出来(品牌名、功能词、卖点词)
|
||||
- 多个商品、多个人物分开描述,不要合并
|
||||
- 看不出来、不确定的字段填null或空数组,布尔值填true/false,绝对不要瞎编
|
||||
- 确保JSON格式合法,所有大括号、中括号、引号正确闭合
|
||||
- 数组字段控制数量:colors最多5个,visible_text最多5条,visual_elements最多8个,accessories最多10个"""
|
||||
V7_USER = "请分析这张图片,按系统消息的JSON结构输出。"
|
||||
|
||||
|
||||
def _capability_table_exists(bind) -> bool:
|
||||
return bool(bind.execute(text("SELECT to_regclass('public.ai_capability_configs')")).scalar())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# 1. 停用旧的active image_analysis prompt(含v6)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
|
||||
"WHERE prompt_type = 'image_analysis' AND is_active = TRUE"
|
||||
)
|
||||
)
|
||||
# 2. 幂等插入v7(存在则更新并重新激活)
|
||||
existing = bind.execute(
|
||||
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
|
||||
).fetchone()
|
||||
if existing:
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
|
||||
"system_prompt = :sys, user_prompt_template = :usr, "
|
||||
"name = 'v7 精简结构化分析', updated_at = NOW() "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 7"
|
||||
),
|
||||
{"sys": V7_SYSTEM, "usr": V7_USER},
|
||||
)
|
||||
else:
|
||||
bind.execute(
|
||||
text(
|
||||
"INSERT INTO viral_video_prompt_templates "
|
||||
"(prompt_type, version, name, system_prompt, user_prompt_template, "
|
||||
"is_active, created_at, updated_at) "
|
||||
"VALUES ('image_analysis', 7, 'v7 精简结构化分析', "
|
||||
":sys, :usr, TRUE, NOW(), NOW())"
|
||||
),
|
||||
{"sys": V7_SYSTEM, "usr": V7_USER},
|
||||
)
|
||||
# 3. capability max_tokens=3000、max_retries=3(表不存在则跳过)
|
||||
if _capability_table_exists(bind):
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_tokens = 3000, "
|
||||
"updated_at = NOW() "
|
||||
"WHERE capability_key = 'image_analysis' AND "
|
||||
"(max_tokens IS NULL OR max_tokens < 3000)"
|
||||
)
|
||||
)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_retries = 3, updated_at = NOW() "
|
||||
"WHERE capability_key = 'image_analysis' AND "
|
||||
"(max_retries IS NULL OR max_retries < 3)"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# 删除v7
|
||||
bind.execute(
|
||||
text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
|
||||
)
|
||||
# 恢复v6为active
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 6"
|
||||
)
|
||||
)
|
||||
# tokens/retries回退
|
||||
if _capability_table_exists(bind):
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_tokens = 1500, max_retries = 1, "
|
||||
"updated_at = NOW() WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
)
|
||||
@@ -280,11 +280,18 @@ def generate_copy(
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED):
|
||||
# 允许首次进入(IMAGE_ANALYZED/PENDING)、失败重试(FAILED)、文案重新生成(COPY_GENERATED/COMPLETED)
|
||||
if job.status not in (
|
||||
ViralVideoStatus.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.PENDING,
|
||||
ViralVideoStatus.FAILED,
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
):
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案")
|
||||
|
||||
# 允许失败任务重试:重置
|
||||
if job.status == ViralVideoStatus.FAILED:
|
||||
# 失败重试 / 重新生成:retry_count 自增
|
||||
if job.status in (ViralVideoStatus.FAILED, ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
|
||||
job.retry_count += 1
|
||||
job.error_msg = ""
|
||||
|
||||
@@ -341,6 +348,9 @@ def confirm_copy(
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status != ViralVideoStatus.COPY_GENERATED:
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
|
||||
# #2218: 额外校验 copy_result 完整性,防止孤儿/脏数据进入渲染
|
||||
if not isinstance(job.copy_result, dict) or not job.copy_result:
|
||||
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
|
||||
|
||||
# 积分预扣(已扣过/重试任务跳过)
|
||||
from app.config import settings as _settings
|
||||
|
||||
@@ -1050,10 +1050,13 @@
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: 360px;
|
||||
padding: 28px 16px;
|
||||
gap: 10px;
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 10px;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-copy-loading .vv-spinner {
|
||||
width: 28px;
|
||||
@@ -1082,18 +1085,38 @@
|
||||
|
||||
/* ── Storyboard (linear doc style) ── */
|
||||
.vv-storyboard {
|
||||
padding: 6px 2px;
|
||||
background: transparent;
|
||||
border: none;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
height: 360px;
|
||||
padding: 10px 12px;
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 10px;
|
||||
margin-top: 8px;
|
||||
overflow: hidden;
|
||||
}
|
||||
|
||||
.vv-sb-doc {
|
||||
flex: 1 1 auto;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 3px;
|
||||
color: #1f2937;
|
||||
font-size: 13px;
|
||||
line-height: 1.55;
|
||||
overflow-y: auto;
|
||||
padding-right: 4px;
|
||||
margin-right: -4px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar-thumb {
|
||||
background: #d8c4ff;
|
||||
border-radius: 3px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar-track {
|
||||
background: transparent;
|
||||
}
|
||||
.vv-sb-h {
|
||||
margin: 6px 0 2px;
|
||||
@@ -1383,12 +1406,14 @@
|
||||
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
|
||||
|
||||
.vv-sb-actions {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
margin-top: 8px;
|
||||
padding-top: 8px;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
}
|
||||
.vv-sb-actions .vv-btn-ghost {
|
||||
padding: 6px 14px;
|
||||
|
||||
@@ -92,22 +92,30 @@ def _save_job(repo, job, session):
|
||||
session.commit()
|
||||
|
||||
|
||||
def _start_trust_chain_preheat(job_id: str, portrait_descriptions: list[str]) -> None:
|
||||
"""#2172/#2174 后台启动信任链预热(Seedream t2i 文生图人像),不阻塞调用方。
|
||||
def _start_trust_chain_preheat(job_id: str, products: list[dict]) -> None:
|
||||
"""#2172/#2174/#2220 后台启动信任链预热(Seedream t2i 文生图人像),不阻塞调用方。
|
||||
|
||||
#2174 重要:改为 t2i 文生图模式——用 VLM 分析出的人物外貌描述做 prompt,不传 reference_images,
|
||||
产物是方舟信任模型输出,Seedance 直接放行不触发肖像审核。
|
||||
i2i(传用户照片做 reference)产物不被信任,实测仍被 400 portrait_intercept 拦截。
|
||||
#2220 修复:只对 has_person=True 的图(真人照片)生成 AI 人像替换,
|
||||
场景图/商品图/门店图保持原图不变,传给 Seedance 作为 reference_image 直接使用。
|
||||
|
||||
预热成功后把结果写入 job.pre_trusted_images,阶段3 渲染直接使用,省掉串行等待。
|
||||
预热失败静默(pre_trusted_images 保持 None),阶段3 会走 #2166 自动降级纯 t2v。
|
||||
预热结果写入 job.pre_trusted_images:与 products 等长的稀疏列表,
|
||||
人像位是 AI 图 URL,非人像位是 None(表示保留原图)。
|
||||
"""
|
||||
# 过滤有效描述:非空且不是"无人像"
|
||||
_valid = [
|
||||
d for d in (portrait_descriptions or []) if d and isinstance(d, str) and "无人像" not in d and len(d) >= 10
|
||||
]
|
||||
# 构建人像位索引映射:person_indices[k] = products中第k个人像的位置
|
||||
_person_indices: list[int] = []
|
||||
_valid: list[str] = []
|
||||
for _i, _p in enumerate(products or []):
|
||||
if not isinstance(_p, dict):
|
||||
continue
|
||||
if not _p.get("has_person", False):
|
||||
continue
|
||||
_d = (_p.get("portrait_prompt") or "").strip()
|
||||
if not _d or "无人像" in _d or len(_d) < 10:
|
||||
continue
|
||||
_person_indices.append(_i)
|
||||
_valid.append(_d)
|
||||
if not _valid:
|
||||
logger.info("[trust-chain][preheat] 无有效人物描述(可能是纯商品图),跳过预热 job=%s", job_id)
|
||||
logger.info("[trust-chain][preheat] 无有效人物描述(可能是纯商品/场景图),跳过预热 job=%s", job_id)
|
||||
return
|
||||
# 判断是否是 doubao provider(DashScope/Wan 不需要信任链)
|
||||
try:
|
||||
@@ -133,21 +141,33 @@ def _start_trust_chain_preheat(job_id: str, portrait_descriptions: list[str]) ->
|
||||
|
||||
logger.info("[trust-chain][preheat] 后台t2i预热启动 job=%s n=%d", job_id, len(_valid))
|
||||
result = preheat_trust_chain(_valid, timeout=120)
|
||||
if result and len(result) >= 1:
|
||||
if result and len(result) == len(_valid):
|
||||
sess2, repo2, job2 = _get_repo_and_job(job_id)
|
||||
try:
|
||||
job2.pre_trusted_images = result
|
||||
# #2220: 构建与 products 等长的稀疏列表,人像位放AI图URL,非人像位None
|
||||
_imgs2 = job2.images or []
|
||||
_sparse: list[str | None] = [None] * max(len(_imgs2), len(products or []))
|
||||
for _k, _url in enumerate(result):
|
||||
if _k < len(_person_indices):
|
||||
_sparse[_person_indices[_k]] = _url
|
||||
job2.pre_trusted_images = _sparse
|
||||
repo2.update(job2)
|
||||
sess2.commit()
|
||||
logger.info(
|
||||
"[trust-chain][preheat] t2i预热完成并持久化 job=%s n=%d",
|
||||
"[trust-chain][preheat] t2i预热完成并持久化 job=%s n_person=%d total=%d",
|
||||
job_id,
|
||||
len(result),
|
||||
len(_sparse),
|
||||
)
|
||||
finally:
|
||||
sess2.close()
|
||||
else:
|
||||
logger.info("[trust-chain][preheat] 预热失败 job=%s,阶段3现场跑兜底", job_id)
|
||||
logger.info(
|
||||
"[trust-chain][preheat] 预热失败或数量不匹配 job=%s got=%s expect=%d,阶段3现场跑兜底",
|
||||
job_id,
|
||||
len(result) if result else 0,
|
||||
len(_valid),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[trust-chain][preheat] 预热异常 job=%s err=%s", job_id, e, exc_info=True)
|
||||
|
||||
@@ -244,8 +264,8 @@ def _recover_stale_jobs() -> int:
|
||||
"SET status='failed', error_msg='任务执行超时,请重试', updated_at=:now "
|
||||
"WHERE status='running' "
|
||||
" AND started_at IS NOT NULL AND started_at < :cutoff_start "
|
||||
" AND (heartbeat_at IS NULL OR heartbeat_at < :cutoff_beat) "
|
||||
" AND (heartbeat_at IS NOT NULL OR updated_at < :cutoff_beat)"
|
||||
" AND updated_at < :cutoff_beat "
|
||||
" AND (heartbeat_at IS NULL OR heartbeat_at < :cutoff_beat)"
|
||||
)
|
||||
result = ssn.execute(sql, {"now": now, "cutoff_start": cutoff_start, "cutoff_beat": cutoff_beat})
|
||||
ssn.commit()
|
||||
@@ -268,6 +288,7 @@ def _recover_stale_jobs() -> int:
|
||||
|
||||
_DEFAULT_HARD_CONSTRAINTS = [
|
||||
"无字幕、无水印、无任何自动生成文字、无 logo",
|
||||
"严格还原参考图片中的真实场景、门店环境、商品陈列、人物外貌服装特征,不得凭空生成与参考图无关的人物、场景或物品",
|
||||
"同一人物全程保持一致的五官、发型、服装、身材,不得换脸或变形",
|
||||
"口播语音必须在指定时长内自然念完,语速自然,口型与语音同步",
|
||||
"画面流畅无闪烁、无多余肢体、无扭曲变形、无穿模",
|
||||
@@ -295,9 +316,9 @@ _DEFAULT_NEGATIVE_PROMPTS = [
|
||||
]
|
||||
|
||||
|
||||
def _empty_copy_result(duration: int = 15, ratio: str = "9:16") -> dict:
|
||||
def _empty_copy_result(duration: int = 15, ratio: str = "9:16", theme: str = "") -> dict:
|
||||
return {
|
||||
"overview": {"theme": "好物推荐", "total_duration": duration, "aspect_ratio": ratio},
|
||||
"overview": {"theme": theme or "好物推荐", "total_duration": duration, "aspect_ratio": ratio},
|
||||
"scene_and_lighting": "简洁明亮的室内场景,柔和自然光,产品主体清晰",
|
||||
"shots": [],
|
||||
"hard_constraints": list(_DEFAULT_HARD_CONSTRAINTS),
|
||||
@@ -323,6 +344,7 @@ def _vision_fallback(idx: int, reason: str, extra: dict | None = None) -> dict:
|
||||
"scene": "通用",
|
||||
"portrait_prompt": "无人像",
|
||||
"summary": "",
|
||||
"has_person": False,
|
||||
"_source": reason,
|
||||
}
|
||||
if extra:
|
||||
@@ -421,10 +443,14 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.shared.ai_service import call_llm
|
||||
from packages.shared.ai_router import ai_router
|
||||
except ImportError:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
_llm_client = ai_router.get_llm_client("intent_parsing")
|
||||
if not _llm_client.is_available:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
products_summary = ""
|
||||
products = (image_analysis or {}).get("products", []) or []
|
||||
for p in products:
|
||||
@@ -445,11 +471,15 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
|
||||
template = get_template("intent_parsing")
|
||||
system = render_system_prompt(template)
|
||||
marketing_purpose = getattr(job, "marketing_purpose", "") or "未指定"
|
||||
image_category_hint = _determine_theme(image_analysis, marketing_purpose)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
user_copy_text=job.user_copy_text or "(未提供,全由 AI 创作)",
|
||||
industry=job.industry or "未指定",
|
||||
image_analysis=products_summary or "- (无图片分析结果)",
|
||||
marketing_purpose=marketing_purpose,
|
||||
image_category_hint=image_category_hint,
|
||||
)
|
||||
|
||||
def _parse(raw: str) -> dict:
|
||||
@@ -473,19 +503,31 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
"suggested_title": "",
|
||||
}
|
||||
|
||||
_s = get_shared_settings()
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = _s.doubao_model
|
||||
for _m, _lbl in [(_fast, "fast"), (_pro, "pro-fallback")]:
|
||||
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
_client_fast = ai_router.get_llm_client("intent_parsing", variant="primary")
|
||||
_client_pro = ai_router.get_llm_client("intent_parsing", variant="lite")
|
||||
_intent_deadline = time.time() + 60
|
||||
_seen_models: set[str] = set()
|
||||
for _client, _lbl in [(_client_fast, "fast"), (_client_pro, "pro-fallback")]:
|
||||
if not _client or not _client.is_available:
|
||||
continue
|
||||
if _client.model in _seen_models:
|
||||
logger.info("[爆款视频] 意图解析跳过重复模型 %s label=%s", _client.model, _lbl)
|
||||
continue
|
||||
_seen_models.add(_client.model)
|
||||
if time.time() > _intent_deadline:
|
||||
logger.warning("[爆款视频] 意图解析超过60s总预算,跳过 label=%s", _lbl)
|
||||
break
|
||||
try:
|
||||
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
|
||||
raw = call_llm(
|
||||
logger.info("[爆款视频] 意图解析 model=%s label=%s", _client.model, _lbl)
|
||||
raw = _client.chat_completion(
|
||||
[{"role": "system", "content": system}, {"role": "user", "content": user}],
|
||||
temperature=0.4,
|
||||
max_tokens=1024,
|
||||
model=_m,
|
||||
timeout=60,
|
||||
) # #2180: 意图解析 LLM 实测需更长响应,原25s太紧
|
||||
timeout=20,
|
||||
)
|
||||
if not raw:
|
||||
continue
|
||||
parsed = _parse(raw)
|
||||
@@ -521,6 +563,70 @@ def _persona_style_hint(persona_id: str) -> str:
|
||||
return "【人设风格:未指定】亲切自然、像朋友分享好物"
|
||||
|
||||
|
||||
def _determine_theme(image_analysis: dict | None, marketing_purpose: str = "") -> str:
|
||||
"""根据图片分析结果和营销目的,智能推断默认主题。
|
||||
|
||||
门店类→门店探店/到店体验;商品图→好物分享/产品种草;
|
||||
人物图→穿搭/人物故事;场景图→场景氛围/空间体验。
|
||||
"""
|
||||
products = (image_analysis or {}).get("products", []) or []
|
||||
type_counts: dict[str, int] = {}
|
||||
for p in products:
|
||||
if not isinstance(p, dict):
|
||||
continue
|
||||
cat = (p.get("category") or "").strip()
|
||||
if any(
|
||||
k in cat
|
||||
for k in (
|
||||
"门店",
|
||||
"店铺",
|
||||
"餐饮",
|
||||
"美容",
|
||||
"美发",
|
||||
"养生",
|
||||
"健身",
|
||||
"酒店",
|
||||
"咖啡",
|
||||
"奶茶",
|
||||
"餐厅",
|
||||
"颈肩",
|
||||
"调理",
|
||||
)
|
||||
):
|
||||
type_counts["store"] = type_counts.get("store", 0) + 1
|
||||
elif any(k in cat for k in ("人物", "穿搭", "人像", "服装")):
|
||||
type_counts["person"] = type_counts.get("person", 0) + 1
|
||||
elif any(k in cat for k in ("场景", "空间", "环境", "非产品")):
|
||||
type_counts["scene"] = type_counts.get("scene", 0) + 1
|
||||
elif cat and cat not in ("无法判断", "非产品图", ""):
|
||||
type_counts["product"] = type_counts.get("product", 0) + 1
|
||||
src = p.get("_source") or ""
|
||||
if "store" in src:
|
||||
type_counts["store"] = type_counts.get("store", 0) + 1
|
||||
elif "person" in src:
|
||||
type_counts["person"] = type_counts.get("person", 0) + 1
|
||||
|
||||
dominant = max(type_counts, key=type_counts.get) if type_counts else "product"
|
||||
mp = (marketing_purpose or "").strip()
|
||||
|
||||
if any(k in mp for k in ("获客", "引流", "到店")):
|
||||
if dominant == "store":
|
||||
return "门店探店·到店体验"
|
||||
return "门店探店·到店体验"
|
||||
if any(k in mp for k in ("品牌", "宣传")):
|
||||
return "品牌故事·门店体验" if dominant == "store" else "品牌故事·产品展示"
|
||||
if any(k in mp for k in ("种草", "推荐")):
|
||||
return "穿搭分享·人物种草" if dominant == "person" else "好物分享·产品种草"
|
||||
|
||||
theme_map = {
|
||||
"store": "门店探店·到店体验",
|
||||
"person": "穿搭分享·人物故事",
|
||||
"scene": "空间体验·场景氛围",
|
||||
"product": "好物分享·产品种草",
|
||||
}
|
||||
return theme_map.get(dominant, "好物分享·产品种草")
|
||||
|
||||
|
||||
def _build_products_summary(image_analysis: dict) -> str:
|
||||
"""把 VLM 返回的商品分析结果拼给文案/分镜生成 prompt 用。
|
||||
优先用 summary(自然段落);没有时用结构化字段兜底拼一段。"""
|
||||
@@ -637,14 +743,36 @@ def _fallback_script(job: ViralVideoJob) -> dict:
|
||||
"""脚本生成失败时的兜底脚本(极简但可用)。"""
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
ratio = getattr(job, "video_ratio", None) or "9:16"
|
||||
base = _empty_copy_result(dur, ratio)
|
||||
voiceover = job.user_copy_text or "你好,给大家分享一款我最近在用的好物,真的很不错,推荐你们也试试。"
|
||||
_ia = getattr(job, "image_analysis", None) or {}
|
||||
_mp = getattr(job, "marketing_purpose", "") or ""
|
||||
default_theme = _determine_theme(_ia, _mp)
|
||||
base = _empty_copy_result(dur, ratio, theme=default_theme)
|
||||
_voiceover_map = {
|
||||
"store": "带你探店!今天来到这家店,环境真的超棒,服务也很到位,推荐大家来体验一下。",
|
||||
"person": "哈喽,今天给大家分享我的日常穿搭,简单舒适又好看,你们觉得怎么样?",
|
||||
"scene": "带大家感受一下这个空间,氛围感拉满,真的很适合打卡体验。",
|
||||
"product": "你好,给大家分享一款我最近在用的好物,真的很不错,推荐你们也试试。",
|
||||
}
|
||||
_products = (_ia or {}).get("products", []) or []
|
||||
_dominant = "product"
|
||||
for p in _products:
|
||||
if not isinstance(p, dict):
|
||||
continue
|
||||
src = p.get("_source") or ""
|
||||
cat = p.get("category") or ""
|
||||
if "store" in src or any(k in cat for k in ("门店", "店铺", "餐饮", "美容", "颈肩", "调理")):
|
||||
_dominant = "store"
|
||||
break
|
||||
elif "person" in src or any(k in cat for k in ("人物", "穿搭", "人像")):
|
||||
_dominant = "person"
|
||||
break
|
||||
voiceover = job.user_copy_text or _voiceover_map.get(_dominant, _voiceover_map["product"])
|
||||
shots = [
|
||||
{
|
||||
"time_range": f"0-{dur}秒",
|
||||
"shot_type_angle_movement": "中景平视,缓慢推镜",
|
||||
"scene_and_dialogue": "明亮室内,人物自然出镜,微笑着看向镜头。" + voiceover,
|
||||
"action_details": "人物手持产品自然展示,表情亲切,动作流畅",
|
||||
"scene_and_dialogue": voiceover,
|
||||
"action_details": "自然展示,表情亲切,动作流畅",
|
||||
"audio_bgm": "轻快流行BGM",
|
||||
"transition": "结束",
|
||||
"reference_image_index": 0 if job.images else None,
|
||||
@@ -654,7 +782,7 @@ def _fallback_script(job: ViralVideoJob) -> dict:
|
||||
base["voiceover_script"] = voiceover
|
||||
base["final_copy"] = voiceover
|
||||
base["suggested_copy"] = voiceover
|
||||
base["title"] = "好物分享"
|
||||
base["title"] = default_theme
|
||||
return base
|
||||
|
||||
|
||||
@@ -672,12 +800,18 @@ def _validate_and_normalize_script(raw, job: ViralVideoJob) -> dict:
|
||||
ov = raw.get("overview")
|
||||
if isinstance(ov, dict):
|
||||
base["overview"] = {
|
||||
"theme": str(ov.get("theme") or "好物分享"),
|
||||
"theme": str(
|
||||
ov.get("theme")
|
||||
or _determine_theme(getattr(job, "image_analysis", None), getattr(job, "marketing_purpose", ""))
|
||||
),
|
||||
"total_duration": int(ov.get("total_duration") or dur),
|
||||
"aspect_ratio": str(ov.get("aspect_ratio") or ratio),
|
||||
}
|
||||
else:
|
||||
base["overview"]["theme"] = str(raw.get("title") or "好物分享")
|
||||
base["overview"]["theme"] = str(
|
||||
raw.get("title")
|
||||
or _determine_theme(getattr(job, "image_analysis", None), getattr(job, "marketing_purpose", ""))
|
||||
)
|
||||
|
||||
base["scene_and_lighting"] = str(raw.get("scene_and_lighting") or base["scene_and_lighting"])
|
||||
|
||||
@@ -764,7 +898,11 @@ def _script_from_xml(raw: str, job: ViralVideoJob) -> dict | None:
|
||||
base = _empty_copy_result(dur, ratio)
|
||||
if not raw:
|
||||
return None
|
||||
base["overview"]["theme"] = xp.text_of(raw, "overview_theme") or xp.text_of(raw, "title") or "好物分享"
|
||||
base["overview"]["theme"] = (
|
||||
xp.text_of(raw, "overview_theme")
|
||||
or xp.text_of(raw, "title")
|
||||
or _determine_theme(getattr(job, "image_analysis", None), getattr(job, "marketing_purpose", ""))
|
||||
)
|
||||
est = xp.attr_int(xp.text_of(raw, "estimated_duration"), 0)
|
||||
if est:
|
||||
base["overview"]["total_duration"] = est
|
||||
@@ -825,10 +963,14 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
GLOBAL_CONSTRAINTS,
|
||||
NEGATIVE_RULES,
|
||||
)
|
||||
from packages.shared.ai_service import call_llm
|
||||
from packages.shared.ai_router import ai_router
|
||||
except ImportError:
|
||||
return _fallback_script(job)
|
||||
|
||||
_llm_client2 = ai_router.get_llm_client("storyboard")
|
||||
if not _llm_client2.is_available:
|
||||
return _fallback_script(job)
|
||||
|
||||
products_summary = _build_products_summary(image_analysis)
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
|
||||
@@ -850,8 +992,11 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
system_tpl = system_tpl.replace("{global_constraints}", GLOBAL_CONSTRAINTS)
|
||||
system_tpl = system_tpl.replace("{negative_rules}", NEGATIVE_RULES)
|
||||
|
||||
marketing_purpose = getattr(job, "marketing_purpose", "") or "未指定"
|
||||
image_category_hint = _determine_theme(image_analysis, marketing_purpose)
|
||||
fusion_brief = (
|
||||
f"意图:{intent_str}\n关键信息:{key_msgs}\n调性:{tone}\n"
|
||||
f"营销目的:{marketing_purpose}\n建议主题方向:{image_category_hint}\n"
|
||||
f"用户原文:{job.user_copy_text or '(未提供)'}\n创作模式:{fusion_level}"
|
||||
)
|
||||
user = render_user_prompt(
|
||||
@@ -862,13 +1007,14 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
image_analysis=products_summary,
|
||||
)
|
||||
|
||||
def _try_gen(model: str, temp: float, max_tok: int, label: str, tmo: int = 25):
|
||||
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", model, label, tmo)
|
||||
raw = call_llm(
|
||||
def _try_gen(client, temp: float, max_tok: int, label: str, tmo: int = 25):
|
||||
if not client or not client.is_available:
|
||||
return None
|
||||
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", client.model, label, tmo)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
|
||||
temperature=temp,
|
||||
max_tokens=max_tok,
|
||||
model=model,
|
||||
timeout=tmo,
|
||||
)
|
||||
if not raw:
|
||||
@@ -899,24 +1045,35 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
)
|
||||
return None if is_fallback else normalized
|
||||
|
||||
_s = get_shared_settings()
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = getattr(_s, "doubao_model", None) or _fast
|
||||
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
|
||||
_client_fast = ai_router.get_llm_client("storyboard", variant="primary")
|
||||
_client_pro = ai_router.get_llm_client("storyboard", variant="lite")
|
||||
_script_fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "90"))
|
||||
_script_pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "60"))
|
||||
_script_deadline = time.time() + 180
|
||||
try:
|
||||
# 第一次:快模型 25s
|
||||
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=90)
|
||||
# #2233: fast_timeout=90s, pro_timeout=60s,总deadline 180s
|
||||
normalized = _try_gen(_client_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
# #2183: 实测pro 1500tok输出需75.8s,单次timeout提到90s
|
||||
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=90)
|
||||
if time.time() > _script_deadline:
|
||||
logger.warning("[爆款视频] 编导脚本超过180s总预算,使用兜底脚本")
|
||||
return _fallback_script(job)
|
||||
normalized = _try_gen(_client_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
# 第三次:用主力模型兜底,给 120s
|
||||
if _pro and _pro != _fast:
|
||||
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=120)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
|
||||
# 第三次:用 lite/pro 模型兜底,跳过与 primary 相同的模型
|
||||
if _client_pro and _client_pro.is_available:
|
||||
if _client_pro.model != _client_fast.model:
|
||||
if time.time() <= _script_deadline:
|
||||
normalized = _try_gen(_client_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
else:
|
||||
logger.warning("[爆款视频] 编导脚本超过180s总预算,跳过pro-fallback")
|
||||
else:
|
||||
logger.info("[爆款视频] pro-fallback模型与primary相同(%s),跳过重复调用", _client_pro.model)
|
||||
logger.warning("[爆款视频] 编导脚本均未生成合格结果,使用兜底脚本")
|
||||
return _fallback_script(job)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 编导脚本生成异常: %s,使用兜底脚本", e, exc_info=True)
|
||||
@@ -1236,9 +1393,12 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
if u not in all_portrait_urls:
|
||||
all_portrait_urls.append(u)
|
||||
pti = getattr(job, "pre_trusted_images", None)
|
||||
if pti and len(pti) == len(all_portrait_urls):
|
||||
pre_trusted = list(pti)
|
||||
logger.info("[爆款视频] 使用信任链预热结果 n=%d,跳过现场 Seedream AI 化", len(pre_trusted))
|
||||
# #2220: 稀疏列表模式(与images等长,None表示该位置保留原图)
|
||||
_n_total = len(all_portrait_urls)
|
||||
if pti and isinstance(pti, list) and len(pti) >= _n_total and _n_total > 0:
|
||||
pre_trusted = list(pti[:_n_total])
|
||||
_n_trusted = sum(1 for _x in pre_trusted if _x)
|
||||
logger.info("[爆款视频] 使用信任链预热结果 person=%d total=%d,跳过现场Seedream AI化", _n_trusted, _n_total)
|
||||
elif all_portrait_urls and _mcfg.get("provider", "doubao") == "doubao":
|
||||
# #2183: 真·现场跑信任链——同步调用 Seedream t2i,拿到 AI 人像 URL 后再传 Seedance
|
||||
logger.info(
|
||||
@@ -1251,32 +1411,39 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
|
||||
_ia = getattr(job, "image_analysis", None) or {}
|
||||
_prods = (_ia.get("products") if isinstance(_ia, dict) else None) or []
|
||||
_pdescs = []
|
||||
if _prods:
|
||||
_pdescs = [(pp.get("portrait_prompt") or "无人像") for pp in _prods]
|
||||
elif isinstance(_ia, dict):
|
||||
_pp0 = _ia.get("portrait_prompt") or "无人像"
|
||||
if _pp0 and _pp0 != "无人像":
|
||||
_pdescs = [_pp0]
|
||||
_valid = [d for d in _pdescs if d and isinstance(d, str) and "无人像" not in d and len(d) >= 10]
|
||||
if _valid:
|
||||
_live_person_idx: list[int] = []
|
||||
_live_pdescs: list[str] = []
|
||||
for _i, _pp in enumerate(_prods):
|
||||
if not isinstance(_pp, dict) or not _pp.get("has_person", False):
|
||||
continue
|
||||
_d = (_pp.get("portrait_prompt") or "").strip()
|
||||
if not _d or "无人像" in _d or len(_d) < 10:
|
||||
continue
|
||||
_live_person_idx.append(_i)
|
||||
_live_pdescs.append(_d)
|
||||
if _live_pdescs:
|
||||
_t0 = time.time()
|
||||
_live_urls = preheat_trust_chain(_valid, timeout=120)
|
||||
if _live_urls and len(_live_urls) == len(all_portrait_urls):
|
||||
pre_trusted = list(_live_urls)
|
||||
_live_urls = preheat_trust_chain(_live_pdescs, timeout=120)
|
||||
if _live_urls and len(_live_urls) == len(_live_pdescs):
|
||||
# #2220: 构建稀疏列表,人像位替换AI图,非人像位保留None(ai_client里用原图)
|
||||
pre_trusted = [None] * len(all_portrait_urls)
|
||||
for _k, _u in enumerate(_live_urls):
|
||||
if _k < len(_live_person_idx):
|
||||
pre_trusted[_live_person_idx[_k]] = _u
|
||||
logger.info(
|
||||
"[爆款视频] 现场信任链t2i完成 %d张 耗时%.1fs,将用AI人像传Seedance",
|
||||
len(pre_trusted),
|
||||
"[爆款视频] 现场信任链t2i完成 %d张人像AI化 耗时%.1fs(共%d张图,其余保留原图)",
|
||||
len(_live_urls),
|
||||
time.time() - _t0,
|
||||
len(all_portrait_urls),
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"[爆款视频] 现场信任链t2i返回不匹配 urls=%s n_portraits=%d,回退原图+400降级纯t2v",
|
||||
"[爆款视频] 现场信任链t2i返回不匹配 urls=%s n_person=%d,人像位原图传Seedance(可能触发400拦截)",
|
||||
_live_urls,
|
||||
len(all_portrait_urls),
|
||||
len(_live_pdescs),
|
||||
)
|
||||
else:
|
||||
logger.info("[爆款视频] 无有效人物描述(可能是商品图),无需现场跑信任链")
|
||||
logger.info("[爆款视频] 无有效人物描述(商品/场景图),无需AI化,直接传原图给Seedance")
|
||||
except Exception as _te:
|
||||
logger.warning("[爆款视频] 现场跑信任链异常: %s,回退原图+400降级纯t2v", _te, exc_info=True)
|
||||
|
||||
@@ -1399,13 +1566,8 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
|
||||
if job.images:
|
||||
try:
|
||||
_products = (image_analysis or {}).get("products", []) or []
|
||||
_portrait_descs = [(p.get("portrait_prompt") or "无人像") for p in _products] if _products else []
|
||||
# 兼容单图结果格式(非products列表)
|
||||
if not _portrait_descs and isinstance(image_analysis, dict):
|
||||
_pp = image_analysis.get("portrait_prompt") or "无人像"
|
||||
if _pp and _pp != "无人像":
|
||||
_portrait_descs = [_pp]
|
||||
_start_trust_chain_preheat(job.id, _portrait_descs)
|
||||
# #2220: 直接传 products 列表,由 _start_trust_chain_preheat 内部按 has_person 筛选
|
||||
_start_trust_chain_preheat(job.id, _products)
|
||||
except Exception as _e:
|
||||
logger.warning("[爆款视频][阶段1] 启动信任链t2i预热失败: %s", _e)
|
||||
_save_job(repo, job, session)
|
||||
@@ -1577,12 +1739,8 @@ def run_viral_video_analyze(self: Task, job_id: str) -> dict:
|
||||
if job.images:
|
||||
try:
|
||||
_products = (image_analysis or {}).get("products", []) or []
|
||||
_portrait_descs = [(p.get("portrait_prompt") or "无人像") for p in _products] if _products else []
|
||||
if not _portrait_descs and isinstance(image_analysis, dict):
|
||||
_pp = image_analysis.get("portrait_prompt") or "无人像"
|
||||
if _pp and _pp != "无人像":
|
||||
_portrait_descs = [_pp]
|
||||
_start_trust_chain_preheat(job.id, _portrait_descs)
|
||||
# #2220: 直接传 products 列表,由 _start_trust_chain_preheat 内部按 has_person 筛选
|
||||
_start_trust_chain_preheat(job.id, _products)
|
||||
except Exception as _e:
|
||||
logger.warning("[爆款视频][阶段1] 启动信任链t2i预热失败: %s", _e)
|
||||
_save_job(repo, job, session)
|
||||
@@ -1670,7 +1828,7 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
intent_result = _step_intent_parsing(job, image_analysis)
|
||||
job.intent_result = intent_result
|
||||
_save_job(repo, job, session)
|
||||
# #2218: 不在意图解析后单独落库,等 copy_result 生成后与 mark_copy_generated 一起原子写入
|
||||
_emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 35.0, "意图解析完成")
|
||||
|
||||
# 阶段:编导脚本生成(核心耗时环节,已用快模型)
|
||||
@@ -1728,7 +1886,8 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
|
||||
except Retry:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段2] 异常: %s", e, exc_info=True)
|
||||
logger.error("[爆款视频][阶段2] 异常 job_id=%s: %s", job_id, e, exc_info=True)
|
||||
# #2218: 阶段2任何异常都标记为 failed(由 _mark_failed_and_notify 处理),前端提示重试
|
||||
_mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.SCRIPT_GENERATION)
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
finally:
|
||||
@@ -1887,16 +2046,26 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
阶段2 generate-copy 已把 LLM 深度审核后置,这里在 TTS 前做最终审核(不通过则自动重写1次)。
|
||||
所有阶段通过 _set_stage 持久化 current_stage/phase_message。
|
||||
"""
|
||||
image_analysis = job.image_analysis or {"products": []}
|
||||
|
||||
# 如果没有 copy_result(旧数据/失败重试),现场补生成(意图+脚本,不走 LLM 审核,出片前会统一做)
|
||||
# #2218: render 流程严禁补生成意图+编导脚本。copy_result 必须由 generate-copy 提前准备好;
|
||||
# 若缺失说明 generate-copy 未完成或数据丢失,直接报错让用户重新点「生成文案」。
|
||||
copy_result = job.copy_result
|
||||
_copy_src = "db"
|
||||
if not isinstance(copy_result, dict) or not copy_result:
|
||||
_set_stage(job, repo, session, ViralVideoStage.SCRIPT_GENERATION, "正在补生成编导脚本...")
|
||||
intent = job.intent_result or _step_intent_parsing(job, image_analysis)
|
||||
copy_result = _step_script_generation(job, intent, image_analysis)
|
||||
job.mark_copy_generated(copy_result)
|
||||
_save_job(repo, job, session)
|
||||
logger.error(
|
||||
"[爆款视频][阶段3] copy_result 为空或无效,无法进入渲染流程。job_id=%s status=%s intent_len=%d,请重新触发「生成文案」",
|
||||
job_id,
|
||||
job.status,
|
||||
len((job.intent_result or {}) if isinstance(job.intent_result, dict) else {}),
|
||||
)
|
||||
raise ValueError("文案数据缺失,请先点击「生成文案」完成文案生成后再生成视频")
|
||||
logger.info(
|
||||
"[爆款视频][阶段3] 进入渲染流程 job_id=%s copy_result_shots=%d copy_result_len=%d source=%s",
|
||||
job_id,
|
||||
len((copy_result.get("shots") or [])),
|
||||
len(str(copy_result)),
|
||||
_copy_src,
|
||||
)
|
||||
|
||||
# 出片前 LLM 深度合规审核(#2134 问题7:审核从阶段2后置到这里,不阻塞前端预览脚本)
|
||||
_set_stage(job, repo, session, ViralVideoStage.REVIEW, "正在进行出片前合规审核...")
|
||||
@@ -1909,12 +2078,18 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
if isinstance(rewritten, dict) and rewritten:
|
||||
copy_result = rewritten
|
||||
else:
|
||||
intent = job.intent_result or _step_intent_parsing(job, image_analysis)
|
||||
copy_result = _step_script_generation(job, intent, image_analysis)
|
||||
_step_review(job, copy_result)
|
||||
# #2218: 审核重写失败不再从意图解析重跑,直接报错让用户重新生成文案
|
||||
logger.error(
|
||||
"[爆款视频][阶段3] 合规审核未通过且自动重写失败 job_id=%s,终止渲染",
|
||||
job_id,
|
||||
)
|
||||
raise ValueError("文案合规审核未通过,请修改文案后重试或重新生成文案")
|
||||
job.copy_result = copy_result
|
||||
job.generated_copy_text = copy_result.get("voiceover_script", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
except ValueError:
|
||||
# #2218: 审核未通过/文案缺失的业务异常,不继续出片,向上抛出
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频][阶段3] 合规审核异常,继续出片: %s", e)
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
|
||||
|
||||
@@ -3,8 +3,9 @@
|
||||
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到纯硬编码 JSON schema prompt。
|
||||
|
||||
规则(简单直接,不做字符串匹配判断):
|
||||
- DB 有 is_active=true 的 image_analysis 记录(含种子默认XML和用户修改后的版本):
|
||||
* system = DB.system_prompt + JSON_SCHEMA_APPEND(追加完整JSON字段schema,覆盖XML等其他输出格式要求)
|
||||
- DB 有 is_active=true 的 image_analysis 记录(含种子版本和用户修改后的版本):
|
||||
* system = DB.system_prompt(DB prompt 自带完整输出格式,不追加硬编码 schema,
|
||||
避免 DB 写 XML、调用强制 json_object 造成的格式冲突)
|
||||
* user = DB.user_prompt_template 渲染后使用;渲染后为空则用硬编码默认
|
||||
- DB 无记录/连接异常/返回空:system/user 全部用纯硬编码 JSON schema prompt
|
||||
"""
|
||||
@@ -73,7 +74,8 @@ _PRO_JSON_SCHEMA = (
|
||||
)
|
||||
DEFAULT_PRO_USER = "分析这张图片,返回符合schema的JSON。"
|
||||
|
||||
# DB 配置存在时,追加在用户 system_prompt 末尾的JSON schema约束
|
||||
# 保留旧 JSON schema 追加文本作为常量(DB prompt 完全控制输出格式后不再使用,
|
||||
# 保留以便排查历史行为)。
|
||||
_FAST_JSON_APPEND = (
|
||||
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
|
||||
"严格包含以下字段(字段值不确定时填null或空数组):\n"
|
||||
@@ -173,7 +175,6 @@ def _resolve(kind: str) -> tuple[str, str]:
|
||||
|
||||
default_sys = _FAST_JSON_SCHEMA if kind == "fast" else _PRO_JSON_SCHEMA
|
||||
default_user = DEFAULT_FAST_USER if kind == "fast" else DEFAULT_PRO_USER
|
||||
append = _FAST_JSON_APPEND if kind == "fast" else _PRO_JSON_APPEND
|
||||
|
||||
sys_prompt = default_sys
|
||||
usr_prompt = default_user
|
||||
@@ -182,7 +183,7 @@ def _resolve(kind: str) -> tuple[str, str]:
|
||||
if tpl is not None:
|
||||
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
|
||||
if db_sys:
|
||||
sys_prompt = db_sys + append
|
||||
sys_prompt = db_sys # DB prompt自带完整输出格式,不追加硬编码schema避免冲突
|
||||
usr_prompt = _render_user(tpl, default_user)
|
||||
logger.info(
|
||||
"[vision.v2] 使用DB image_analysis prompt (kind=%s version=%s sys_len=%d)",
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""把 fast_json VLM 输出 + OCR 文本组装为与旧 _normalize() 完全一致的 dict。
|
||||
"""把 fast_json VLM 输出 + OCR 文本组装为下游兼容的 product dict。
|
||||
|
||||
v4 schema: DB prompt完全控制输出格式,可能是v4嵌套schema(type/products/people/store_info)
|
||||
或旧扁平schema(has_person/upper_wear/product_name/brand等)。assembler兼容两种格式。
|
||||
|
||||
目标:下游(信任链t2i/intent_parsing/script_generation)零改动。
|
||||
必出字段:name, brand, category, appearance, packaging, text_on_package,
|
||||
@@ -10,29 +13,16 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# ---------- portrait_prompt 模板 ----------
|
||||
# 目标:60-100 字的人物穿搭描述,用于 Seedream 纯文生图。要求具体、风格化、视觉细节丰富。
|
||||
# 旧 VLM 输出格式参考:"一位25岁左右的亚洲女性,身穿白色V领短袖T恤,黑色高腰阔腿裤,
|
||||
# 搭配银色项链,长发披肩,表情自信,街拍风格,阳光明媚的城市街头"
|
||||
|
||||
|
||||
def _join_parts(*parts: str | None) -> str:
|
||||
return "".join(p for p in parts if p)
|
||||
|
||||
|
||||
_AGE_PREFIX = {
|
||||
"青年": "年轻",
|
||||
"中年": "中年",
|
||||
"老年": "老年",
|
||||
}
|
||||
# gender 后缀
|
||||
_AGE_PREFIX = {"青年": "年轻", "中年": "中年", "老年": "老年"}
|
||||
_GENDER_WORD = {"男": "男性", "女": "女性"}
|
||||
|
||||
|
||||
def _person_subject(fj: dict[str, Any]) -> str:
|
||||
"""人物主语:年轻女性 / 中年男性 / 少女 / 小男孩 / 人物 等。"""
|
||||
gender = fj.get("gender") or ""
|
||||
age = fj.get("age_range") or ""
|
||||
def _person_subject(gender: str, age: str) -> str:
|
||||
gw = _GENDER_WORD.get(gender, "")
|
||||
if age == "儿童":
|
||||
if gender == "女":
|
||||
@@ -52,8 +42,147 @@ def _person_subject(fj: dict[str, Any]) -> str:
|
||||
return f"{prefix}人物" if prefix else "人物"
|
||||
|
||||
|
||||
def _build_wear_sentence(fj: dict[str, Any]) -> str:
|
||||
"""穿搭段:上装+下装/连衣裙,带颜色+材质+图案。"""
|
||||
def _build_wear_from_v4(p: dict) -> str:
|
||||
"""v4 person schema: upper_wear/upper_color/lower_wear/lower_color/dress_color"""
|
||||
upper = p.get("upper_wear") or ""
|
||||
upper_color = p.get("upper_color") or ""
|
||||
lower = p.get("lower_wear") or ""
|
||||
lower_color = p.get("lower_color") or ""
|
||||
dress_color = p.get("dress_color") or ""
|
||||
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
|
||||
if is_dress:
|
||||
c = dress_color or upper_color
|
||||
return f"身穿{c}{upper}" if c else f"身穿{upper}"
|
||||
parts = []
|
||||
if upper:
|
||||
up = f"{upper_color}{upper}" if upper_color else upper
|
||||
parts.append(f"上身{up}")
|
||||
if lower:
|
||||
lo = f"{lower_color}{lower}" if lower_color else lower
|
||||
parts.append(f"下身{lo}")
|
||||
return ",".join(parts)
|
||||
|
||||
|
||||
def _build_portrait_prompt_from_v4(p: dict) -> str:
|
||||
"""v4 person: 直接用portrait_prompt字段;没有就拼"""
|
||||
direct = p.get("portrait_prompt")
|
||||
if direct and len(direct) >= 10:
|
||||
return direct
|
||||
subject = _person_subject(p.get("gender", ""), p.get("age_range", ""))
|
||||
wear = _build_wear_from_v4(p)
|
||||
acc = p.get("accessories") or []
|
||||
if isinstance(acc, str):
|
||||
acc = [acc]
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in acc if a) if acc else ""
|
||||
hair = p.get("hairstyle") or ""
|
||||
expr = p.get("expression") or ""
|
||||
pose = p.get("pose") or ""
|
||||
style = p.get("outfit_style") or p.get("style") or ""
|
||||
scene = p.get("scene") or ""
|
||||
mood = p.get("mood") or ""
|
||||
details = []
|
||||
if hair:
|
||||
details.append(hair)
|
||||
if expr and expr not in ("自然", "平静"):
|
||||
details.append(f"神情{expr}")
|
||||
if pose and pose not in ("站立",):
|
||||
details.append(pose)
|
||||
style_parts = []
|
||||
if style:
|
||||
style_parts.append(style)
|
||||
if mood:
|
||||
style_parts.append(mood)
|
||||
if scene and scene not in ("通用",):
|
||||
style_parts.append(scene)
|
||||
pieces = [f"一位{subject}"]
|
||||
if wear:
|
||||
pieces.append(wear)
|
||||
if acc_str:
|
||||
pieces.append(acc_str.lstrip(","))
|
||||
if details:
|
||||
pieces.append(",".join(details))
|
||||
pieces.append(("".join(style_parts) + "风格") if style_parts else "人像写真")
|
||||
full = ",".join(p for p in pieces if p)
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
|
||||
|
||||
def _build_product_prompt_from_v4(prod: dict, top: dict) -> str:
|
||||
"""v4 product: 拼商品视觉描述prompt(用于AI生图参考)"""
|
||||
name = prod.get("product_name") or "商品"
|
||||
brand = prod.get("brand") or ""
|
||||
lead = f"{brand} {name}" if brand and brand not in name else name
|
||||
pkg_color = prod.get("package_color") or ""
|
||||
pkg_type = prod.get("package_type") or ""
|
||||
cap = prod.get("cap_type") or ""
|
||||
body = prod.get("body_shape") or ""
|
||||
features = prod.get("product_features") or []
|
||||
sell = prod.get("key_selling_points") or []
|
||||
colors = top.get("colors") or []
|
||||
style = top.get("style") or ""
|
||||
scene = top.get("scene") or ""
|
||||
mood = top.get("mood") or ""
|
||||
|
||||
parts = [lead]
|
||||
desc = []
|
||||
if pkg_color:
|
||||
desc.append(pkg_color)
|
||||
if pkg_type:
|
||||
desc.append(pkg_type)
|
||||
if cap and len(desc) < 3:
|
||||
desc.append(f"配{cap}")
|
||||
if body and len(desc) < 3:
|
||||
desc.append(body)
|
||||
if desc:
|
||||
parts.append(",".join(desc))
|
||||
if features:
|
||||
core = [str(f) for f in features[:3] if f and len(str(f)) <= 25]
|
||||
if core:
|
||||
parts.append(";".join(core))
|
||||
if sell:
|
||||
s = [str(x) for x in sell[:2] if x]
|
||||
if s:
|
||||
parts.append("突出" + "、".join(s))
|
||||
cnames = []
|
||||
for cc in colors:
|
||||
if isinstance(cc, dict) and cc.get("name"):
|
||||
cnames.append(cc["name"])
|
||||
elif isinstance(cc, str):
|
||||
cnames.append(cc)
|
||||
cnames = cnames[:3]
|
||||
if cnames:
|
||||
parts.append("、".join(cnames) + "主色")
|
||||
if style:
|
||||
parts.append(style)
|
||||
if mood:
|
||||
parts.append(mood)
|
||||
if scene and not any(k in scene for k in ("白色背景", "纯色", "通用")):
|
||||
parts.append(scene)
|
||||
parts.append("产品特写,画面清晰")
|
||||
prompt = ",".join(p for p in parts if p)
|
||||
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
|
||||
|
||||
|
||||
def _is_v4_schema(fj: dict) -> bool:
|
||||
"""判断是v4嵌套schema还是旧扁平schema"""
|
||||
return (
|
||||
isinstance(fj.get("products"), list)
|
||||
or fj.get("type") in ("product", "store", "person", "scene", "other")
|
||||
or isinstance(fj.get("people"), dict)
|
||||
)
|
||||
|
||||
|
||||
# ---------- 旧扁平schema兼容(保留原逻辑) ----------
|
||||
|
||||
|
||||
def _person_subject_old(fj: dict) -> str:
|
||||
return _person_subject(fj.get("gender", ""), fj.get("age_range", ""))
|
||||
|
||||
|
||||
def _build_wear_sentence_old(fj: dict) -> str:
|
||||
upper = fj.get("upper_wear") or ""
|
||||
upper_color = fj.get("upper_color") or ""
|
||||
lower = fj.get("lower_wear") or ""
|
||||
@@ -61,7 +190,6 @@ def _build_wear_sentence(fj: dict[str, Any]) -> str:
|
||||
dress_color = fj.get("dress_color") or ""
|
||||
material = fj.get("material") or ""
|
||||
pattern = fj.get("pattern") or ""
|
||||
|
||||
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
|
||||
if is_dress:
|
||||
c = dress_color or upper_color
|
||||
@@ -71,25 +199,22 @@ def _build_wear_sentence(fj: dict[str, Any]) -> str:
|
||||
if pattern and pattern not in wear and pattern != "纯色":
|
||||
wear += f",{pattern}图案"
|
||||
return f"身穿{wear}"
|
||||
|
||||
parts: list[str] = []
|
||||
parts = []
|
||||
if upper:
|
||||
up = f"{upper_color}{upper}" if upper_color else upper
|
||||
if material and material not in up:
|
||||
up = f"{material}{up}"
|
||||
if pattern and pattern != "纯色" and pattern not in up:
|
||||
up += f"({pattern})"
|
||||
parts.append(f"上身{up}" if up else "")
|
||||
parts.append(f"上身{up}")
|
||||
if lower:
|
||||
lo = f"{lower_color}{lower}" if lower_color else lower
|
||||
parts.append(f"下身{lo}" if lo else "")
|
||||
parts.append(f"下身{lo}")
|
||||
return ",".join(p for p in parts if p)
|
||||
|
||||
|
||||
def _build_portrait_prompt(fj: dict[str, Any]) -> str:
|
||||
"""组装最终 portrait_prompt(目标 60-100 字,用于 Seedream 纯文生图)。"""
|
||||
def _build_portrait_prompt_old(fj: dict) -> str:
|
||||
if not fj.get("has_person"):
|
||||
# 非人像:用商品+场景+mood 拼一段
|
||||
name = fj.get("product_name") or "商品"
|
||||
brand = fj.get("brand") or ""
|
||||
colors = fj.get("colors") or []
|
||||
@@ -101,7 +226,15 @@ def _build_portrait_prompt(fj: dict[str, Any]) -> str:
|
||||
pieces.append(brand)
|
||||
pieces.append(name)
|
||||
if colors:
|
||||
pieces.append("、".join(colors[:3]) + "配色")
|
||||
cnames = []
|
||||
for c in colors:
|
||||
if isinstance(c, dict):
|
||||
cnames.append(c.get("name", ""))
|
||||
elif isinstance(c, str):
|
||||
cnames.append(c)
|
||||
cnames = [c for c in cnames if c][:3]
|
||||
if cnames:
|
||||
pieces.append("、".join(cnames) + "配色")
|
||||
if style:
|
||||
pieces.append(style + "风格")
|
||||
if mood:
|
||||
@@ -111,40 +244,32 @@ def _build_portrait_prompt(fj: dict[str, Any]) -> str:
|
||||
pieces.append("产品特写")
|
||||
prompt = ",".join(p for p in pieces if p)
|
||||
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
|
||||
|
||||
subject = _person_subject(fj)
|
||||
wear = _build_wear_sentence(fj)
|
||||
|
||||
subject = _person_subject_old(fj)
|
||||
wear = _build_wear_sentence_old(fj)
|
||||
accessories = fj.get("accessories") or []
|
||||
if isinstance(accessories, str):
|
||||
accessories = [accessories]
|
||||
acc_str = ""
|
||||
if accessories:
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a)
|
||||
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a) if accessories else ""
|
||||
hairstyle = fj.get("hairstyle") or ""
|
||||
expression = fj.get("expression") or ""
|
||||
pose = fj.get("pose") or ""
|
||||
style = fj.get("style") or ""
|
||||
scene = fj.get("scene") or ""
|
||||
mood = fj.get("mood") or ""
|
||||
|
||||
detail_parts: list[str] = []
|
||||
detail_parts = []
|
||||
if hairstyle:
|
||||
detail_parts.append(hairstyle)
|
||||
if expression and expression not in ("自然", "平静"):
|
||||
detail_parts.append(f"神情{expression}")
|
||||
if pose and pose not in ("站立",):
|
||||
detail_parts.append(pose)
|
||||
|
||||
style_parts: list[str] = []
|
||||
style_parts = []
|
||||
if style:
|
||||
style_parts.append(style)
|
||||
if mood:
|
||||
style_parts.append(mood)
|
||||
if scene and scene not in ("通用",):
|
||||
style_parts.append(scene)
|
||||
|
||||
pieces = [f"一位{subject}"]
|
||||
if wear:
|
||||
pieces.append(wear)
|
||||
@@ -152,53 +277,40 @@ def _build_portrait_prompt(fj: dict[str, Any]) -> str:
|
||||
pieces.append(acc_str.lstrip(","))
|
||||
if detail_parts:
|
||||
pieces.append(",".join(detail_parts))
|
||||
if style_parts:
|
||||
# 风格词之间不用逗号,用空格紧凑
|
||||
pieces.append("".join(style_parts) + "风格")
|
||||
else:
|
||||
pieces.append("人像写真")
|
||||
|
||||
pieces.append("".join(style_parts) + "风格" if style_parts else "人像写真")
|
||||
full = ",".join(p for p in pieces if p)
|
||||
# 过短补充镜头词
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
# 过长截断
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
|
||||
|
||||
# ---------- 商品字段 ----------
|
||||
|
||||
|
||||
def _infer_name(fj: dict[str, Any], ocr_texts: list[str]) -> str:
|
||||
def _infer_name_old(fj: dict, ocr_texts: list[str]) -> str:
|
||||
pname = fj.get("product_name")
|
||||
if pname and pname != "未识别":
|
||||
return str(pname)
|
||||
# 人物图 → name 用穿搭主件
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or ""
|
||||
if "连衣裙" in up:
|
||||
return up
|
||||
return up or "人物穿搭"
|
||||
if ocr_texts:
|
||||
# 商品名可能是 OCR 最长的一行(品牌/产品名)
|
||||
return max(ocr_texts, key=len)
|
||||
return "未识别"
|
||||
|
||||
|
||||
def _infer_brand(fj: dict[str, Any], ocr_texts: list[str]) -> str:
|
||||
def _infer_brand_old(fj: dict, ocr_texts: list[str]) -> str:
|
||||
brand = fj.get("brand")
|
||||
if brand:
|
||||
return str(brand)
|
||||
# OCR 里短的、纯字母/汉字短串可能是 brand
|
||||
for t in ocr_texts:
|
||||
if 1 < len(t) <= 12:
|
||||
return t
|
||||
return "无法判断"
|
||||
|
||||
|
||||
def _infer_category(fj: dict[str, Any]) -> str:
|
||||
def _infer_category_old(fj: dict) -> str:
|
||||
cat = fj.get("category")
|
||||
if cat:
|
||||
return str(cat)
|
||||
@@ -207,27 +319,19 @@ def _infer_category(fj: dict[str, Any]) -> str:
|
||||
return "非产品图"
|
||||
|
||||
|
||||
def _build_appearance(fj: dict[str, Any]) -> str:
|
||||
"""外观描述:颜色+款式+材质+图案 拼成一段。"""
|
||||
parts: list[str] = []
|
||||
for key, _label in [
|
||||
("upper_color", "主色"),
|
||||
("upper_wear", "款式"),
|
||||
("material", "材质"),
|
||||
("pattern", "图案"),
|
||||
]:
|
||||
def _build_appearance_old(fj: dict) -> str:
|
||||
parts = []
|
||||
for key in ("upper_color", "upper_wear", "material", "pattern"):
|
||||
v = fj.get(key)
|
||||
if v and v not in ("无法判断", "未知", "纯色"):
|
||||
parts.append(str(v))
|
||||
if not parts:
|
||||
if fj.get("has_person"):
|
||||
return "人像穿搭整体造型"
|
||||
return "无法判断"
|
||||
return "人像穿搭整体造型" if fj.get("has_person") else "无法判断"
|
||||
return "、".join(parts)
|
||||
|
||||
|
||||
def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
|
||||
feats: list[str] = []
|
||||
def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
|
||||
feats = []
|
||||
for key in (
|
||||
"upper_wear",
|
||||
"lower_wear",
|
||||
@@ -248,9 +352,7 @@ def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
|
||||
feats.append(v)
|
||||
if ocr_texts:
|
||||
feats.append(f"画面文字: {'/'.join(ocr_texts[:3])}")
|
||||
# 去重
|
||||
out: list[str] = []
|
||||
seen: set[str] = set()
|
||||
out, seen = [], set()
|
||||
for f in feats:
|
||||
f = f.strip()
|
||||
if f and f not in seen and len(f) <= 30:
|
||||
@@ -259,27 +361,552 @@ def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
|
||||
return out[:6] if out else ["无法判断"]
|
||||
|
||||
|
||||
def assemble_result(
|
||||
idx: int,
|
||||
fast_json: dict[str, Any] | None,
|
||||
ocr_texts: list[str],
|
||||
) -> dict[str, Any]:
|
||||
"""把 fast_json 结果 + OCR 文本组装成下游兼容的 product dict。"""
|
||||
def _flatten_colors(c) -> list[str]:
|
||||
"""colors可能是字符串数组或[{hex,name,coverage}],统一返回名字数组"""
|
||||
if not c:
|
||||
return []
|
||||
out = []
|
||||
for item in c:
|
||||
if isinstance(item, dict):
|
||||
n = item.get("name")
|
||||
if n:
|
||||
out.append(n)
|
||||
elif isinstance(item, str):
|
||||
out.append(item)
|
||||
return out
|
||||
|
||||
|
||||
def assemble_result(idx: int, fast_json: dict | None, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
fj = fast_json or {}
|
||||
ocr_texts = ocr_texts or []
|
||||
|
||||
portrait_prompt = _build_portrait_prompt(fj)
|
||||
name = _infer_name(fj, ocr_texts)
|
||||
brand = _infer_brand(fj, ocr_texts)
|
||||
category = _infer_category(fj)
|
||||
appearance = _build_appearance(fj)
|
||||
key_features = _build_key_features(fj, ocr_texts)
|
||||
if _is_v4_schema(fj):
|
||||
result = _assemble_v4(idx, fj, ocr_texts)
|
||||
else:
|
||||
result = _assemble_old(idx, fj, ocr_texts)
|
||||
return _apply_partial_fallback(result, fj)
|
||||
|
||||
|
||||
def _apply_partial_fallback(result: dict[str, Any], fj: dict) -> dict[str, Any]:
|
||||
"""partial(截断修复)产物的字段兜底:用已有碎片填充空字段,
|
||||
避免"无法判断"直接透传给下游。非partial产物原样返回。"""
|
||||
if not fj.get("_partial"):
|
||||
return result
|
||||
desc = str(fj.get("description") or "").strip()
|
||||
# 收集所有顶层标量碎片作为兜底素材
|
||||
fragments: list[str] = []
|
||||
for k in ("main_subject", "store_type", "scene_type", "description"):
|
||||
v = fj.get(k)
|
||||
if isinstance(v, str) and v.strip() and v != "无法判断":
|
||||
fragments.append(v.strip())
|
||||
for arr_k in ("environment_objects", "key_elements", "visual_elements"):
|
||||
arr = fj.get(arr_k) or []
|
||||
if isinstance(arr, list):
|
||||
for item in arr[:3]:
|
||||
if isinstance(item, str) and item.strip():
|
||||
fragments.append(item.strip())
|
||||
elif isinstance(item, dict):
|
||||
tv = item.get("text") or item.get("name")
|
||||
if tv:
|
||||
fragments.append(str(tv))
|
||||
frag_text = ";".join(fragments[:3])
|
||||
|
||||
if result.get("name") in ("未识别", "", None) and (desc or frag_text):
|
||||
result["name"] = (desc or fragments[0])[:30]
|
||||
if str(result.get("appearance", "")).startswith("无法判断"):
|
||||
if desc:
|
||||
result["appearance"] = desc[:200]
|
||||
elif frag_text:
|
||||
result["appearance"] = frag_text[:200]
|
||||
if result.get("key_features") in (["无法判断"], []) and (desc or fragments):
|
||||
kf = []
|
||||
if desc:
|
||||
kf.append(desc[:30])
|
||||
for f in fragments[:3]:
|
||||
if f not in kf:
|
||||
kf.append(f[:40])
|
||||
result["key_features"] = kf[:8]
|
||||
if result.get("summary") in ("未识别", "", None) and (desc or fragments):
|
||||
result["summary"] = (desc or fragments[0])[:40]
|
||||
result["_partial"] = True
|
||||
return result
|
||||
|
||||
|
||||
def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"""v4嵌套schema → 下游product dict"""
|
||||
vtype = fj.get("type") or "other"
|
||||
products = fj.get("products") or []
|
||||
scene = fj.get("scene") or "通用"
|
||||
mood = fj.get("mood") or ""
|
||||
packaging = "无法判断" # 包装细节专用API无,保留占位
|
||||
text_on_package = ocr_texts[:8]
|
||||
summary = _build_summary(fj, name, brand, category)
|
||||
colors = fj.get("colors") or []
|
||||
visible_text = fj.get("visible_text") or []
|
||||
color_names = _flatten_colors(colors)
|
||||
|
||||
# 合并OCR文字和visible_text
|
||||
pkg_texts = []
|
||||
for vt in visible_text:
|
||||
if isinstance(vt, dict):
|
||||
t = vt.get("text")
|
||||
if t:
|
||||
pkg_texts.append(str(t))
|
||||
elif isinstance(vt, str):
|
||||
pkg_texts.append(vt)
|
||||
pkg_texts.extend(ocr_texts[:5])
|
||||
# 去重
|
||||
seen_t = set()
|
||||
text_on_package = []
|
||||
for t in pkg_texts:
|
||||
t = str(t).strip()
|
||||
if t and t not in seen_t and len(t) <= 50:
|
||||
seen_t.add(t)
|
||||
text_on_package.append(t)
|
||||
text_on_package = text_on_package[:8]
|
||||
|
||||
has_person = fj.get("has_person", False)
|
||||
|
||||
# 兼容老schema:无type字段或type不在已知枚举时,按has_person/products兜底
|
||||
_KNOWN_V4_TYPES = ("product", "store", "person", "scene", "other")
|
||||
if vtype not in _KNOWN_V4_TYPES:
|
||||
if has_person:
|
||||
vtype = "person"
|
||||
elif products:
|
||||
vtype = "product"
|
||||
else:
|
||||
vtype = "other"
|
||||
|
||||
# 人物信息提取辅助(scene/store分支有人物时追加到key_features)
|
||||
def _extract_person_features(source: dict) -> list[str]:
|
||||
parts = []
|
||||
for pk in ("outfit_style", "upper_wear", "lower_wear", "dress_wear", "outerwear", "pose", "expression"):
|
||||
pv = source.get(pk)
|
||||
if pv and str(pv) not in ("null", None, "无法判断") and len(str(pv)) <= 40:
|
||||
parts.append(f"人物:{pv}")
|
||||
for pk2 in ("shoes", "bag", "hairstyle"):
|
||||
pv2 = source.get(pk2)
|
||||
if pv2 and str(pv2) not in ("null", None) and len(str(pv2)) <= 40:
|
||||
parts.append(f"人物:{pv2}")
|
||||
return parts
|
||||
|
||||
# partial截断保护:声明了product但products数组没来得及输出时,
|
||||
# 按已返回的碎片字段改路由,避免直接掉到other丢信息
|
||||
if fj.get("_partial") and vtype == "product" and not products:
|
||||
if any(fj.get(k) for k in ("signage_details", "store_layout", "brand_signage", "store_type")):
|
||||
vtype = "store"
|
||||
elif any(fj.get(k) for k in ("key_elements", "main_subject", "scene_type", "spatial_layout")):
|
||||
vtype = "scene"
|
||||
else:
|
||||
vtype = "other"
|
||||
|
||||
# ── 人物类 ──
|
||||
if vtype == "person":
|
||||
# 取第一个人物信息(v5 schema人物信息在顶层)
|
||||
person_info = fj
|
||||
# 兼容people嵌套
|
||||
ppl = fj.get("people")
|
||||
if isinstance(ppl, dict) and ppl.get("has_person"):
|
||||
person_info = {**fj, **ppl}
|
||||
has_person = True
|
||||
|
||||
portrait_prompt = _build_portrait_prompt_from_v4(person_info)
|
||||
outfit_style = person_info.get("outfit_style") or ""
|
||||
upper = person_info.get("upper_wear") or ""
|
||||
lower = person_info.get("lower_wear") or ""
|
||||
dress = person_info.get("dress_wear") or ""
|
||||
outer = person_info.get("outerwear") or ""
|
||||
if dress:
|
||||
name = str(dress)[:25]
|
||||
elif outer and upper:
|
||||
name = f"{outer}+{upper}"[:30]
|
||||
elif upper:
|
||||
name = (str(upper) + (f"+{lower}" if lower else ""))[:30]
|
||||
else:
|
||||
name = "人物穿搭"
|
||||
brand = "无法判断"
|
||||
category = "人物穿搭"
|
||||
# appearance: 外套+上衣+下装/裙+鞋+包+发型+妆容
|
||||
app_parts = []
|
||||
for k in ("outerwear", "upper_wear", "lower_wear", "dress_wear", "shoes", "bag", "hairstyle", "makeup"):
|
||||
v = person_info.get(k)
|
||||
if v and v not in ("null", None, "无明显妆容"):
|
||||
app_parts.append(str(v))
|
||||
appearance = ";".join(app_parts) if app_parts else "人像穿搭整体造型"
|
||||
# key_features: 服装+配饰+拍摄信息
|
||||
kf = []
|
||||
for k in (
|
||||
"outfit_style",
|
||||
"upper_wear",
|
||||
"lower_wear",
|
||||
"dress_wear",
|
||||
"outerwear",
|
||||
"shoes",
|
||||
"bag",
|
||||
"hairstyle",
|
||||
"expression",
|
||||
"pose",
|
||||
):
|
||||
v = person_info.get(k)
|
||||
if v and v not in ("null", None, "无法判断"):
|
||||
kf.append(str(v))
|
||||
acc = person_info.get("accessories") or []
|
||||
if isinstance(acc, list):
|
||||
for a in acc:
|
||||
if a and str(a) not in kf:
|
||||
kf.append(str(a))
|
||||
elif isinstance(acc, str) and acc:
|
||||
kf.append(acc)
|
||||
for k in ("shot_type", "camera_angle", "lighting", "atmosphere"):
|
||||
v = person_info.get(k)
|
||||
if v and v not in ("null", None):
|
||||
kf.append(str(v))
|
||||
env_obj = person_info.get("environment_objects") or []
|
||||
if not isinstance(env_obj, list):
|
||||
env_obj = [env_obj]
|
||||
for o in env_obj:
|
||||
if o and str(o) not in kf and len(str(o)) <= 40:
|
||||
kf.append(str(o))
|
||||
if text_on_package:
|
||||
kf.append(f"文字:{'/'.join(text_on_package[:3])}")
|
||||
kf = kf[:8] or ["无法判断"]
|
||||
summary = (outfit_style + " " if outfit_style and outfit_style not in name else "") + name[:25]
|
||||
if not summary.strip():
|
||||
summary = "人物穿搭"
|
||||
return {
|
||||
"name": name[:30],
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance[:400],
|
||||
"packaging": "人物形象无包装",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt[:300],
|
||||
"summary": summary[:50],
|
||||
"_source": "v2_fast_json_v5",
|
||||
"has_person": True,
|
||||
}
|
||||
|
||||
# 商品类
|
||||
if vtype == "product" and products:
|
||||
# 主商品(第一个position=main或第一个)
|
||||
main = products[0]
|
||||
for p in products:
|
||||
if p.get("position") == "main":
|
||||
main = p
|
||||
break
|
||||
name = main.get("product_name") or "未识别"
|
||||
brand = main.get("brand") or "无法判断"
|
||||
category = main.get("category") or "非产品图"
|
||||
# appearance: 包装外观
|
||||
app_parts = []
|
||||
for k in ("package_color", "package_type", "cap_type", "body_shape", "label_design"):
|
||||
v = main.get(k)
|
||||
if v and v not in ("null", None):
|
||||
app_parts.append(str(v))
|
||||
appearance = ";".join(app_parts) if app_parts else "无法判断"
|
||||
# packaging: 包装信息(直接用package_type+package_color)
|
||||
pkg_parts = []
|
||||
if main.get("package_type"):
|
||||
pkg_parts.append(str(main["package_type"]))
|
||||
if main.get("package_color"):
|
||||
pkg_parts.append(str(main["package_color"]))
|
||||
if main.get("cap_type"):
|
||||
pkg_parts.append(f"配{main['cap_type']}")
|
||||
packaging = ",".join(pkg_parts) if pkg_parts else "无法判断"
|
||||
# key_features: product_features字段
|
||||
feats = main.get("product_features") or []
|
||||
if not isinstance(feats, list):
|
||||
feats = [str(feats)]
|
||||
kf = [str(f) for f in feats if f and len(str(f)) <= 40][:6]
|
||||
# 补充卖点
|
||||
sell = main.get("key_selling_points") or []
|
||||
if isinstance(sell, list):
|
||||
for s in sell[:2]:
|
||||
if s and len(str(s)) <= 30 and str(s) not in kf:
|
||||
kf.append(f"卖点:{s}")
|
||||
env_obj = fj.get("environment_objects") or []
|
||||
if not isinstance(env_obj, list):
|
||||
env_obj = [env_obj]
|
||||
for o in env_obj:
|
||||
if o and str(o) not in kf and len(str(o)) <= 40:
|
||||
kf.append(str(o))
|
||||
if text_on_package:
|
||||
kf.append(f"文字: {'/'.join(text_on_package[:3])}")
|
||||
kf = kf[:8] or ["无法判断"]
|
||||
portrait_prompt = _build_product_prompt_from_v4(main, fj)
|
||||
if brand != "无法判断" and brand not in name:
|
||||
summary = f"{brand} {name}"
|
||||
else:
|
||||
summary = name
|
||||
return {
|
||||
"name": str(name)[:50],
|
||||
"brand": str(brand)[:30],
|
||||
"category": str(category)[:20],
|
||||
"appearance": appearance[:200],
|
||||
"packaging": packaging[:100],
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": str(summary)[:60],
|
||||
"has_person": False,
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
# 门店类
|
||||
if vtype == "store":
|
||||
store_type = fj.get("store_type") or "店铺"
|
||||
# brand 多级兜底:brand_signage → visible_text招牌文字 → text_on_package短词
|
||||
brand_raw = fj.get("brand_signage")
|
||||
if not brand_raw or brand_raw in ("无法判断", "", None):
|
||||
brand = None
|
||||
# 从visible_text找招牌文字(通常是位置含招牌/门头/背景的短词)
|
||||
for vt in visible_text:
|
||||
vt_str = vt.get("text") if isinstance(vt, dict) else str(vt)
|
||||
if not vt_str or len(vt_str) < 2 or len(vt_str) > 12:
|
||||
continue
|
||||
loc = (vt.get("location") or "") if isinstance(vt, dict) else ""
|
||||
if any(k in loc for k in ("招牌", "门头", "背景", "招牌墙")):
|
||||
brand = vt_str
|
||||
break
|
||||
# 从text_on_package找2-8字的短词(非描述性)
|
||||
if not brand:
|
||||
_desc_words = {"干净", "整洁", "温馨", "专业", "明亮", "舒适", "宽敞", "现代", "传统", "时尚"}
|
||||
for t in text_on_package:
|
||||
if 2 <= len(t) <= 8 and t not in _desc_words and not any(c in t for c in "的了是在我"):
|
||||
brand = t
|
||||
break
|
||||
if not brand:
|
||||
brand = "无法判断"
|
||||
else:
|
||||
brand = brand_raw
|
||||
# name兜底:store_type为空时用brand
|
||||
name = store_type if store_type != "店铺" else (brand if brand != "无法判断" else store_type)
|
||||
category = "门店场景"
|
||||
# appearance: store_layout + furnishings + 陈设色调
|
||||
appearance_parts = []
|
||||
if fj.get("store_layout"):
|
||||
appearance_parts.append(str(fj["store_layout"]))
|
||||
furnishings = fj.get("furnishings") or []
|
||||
if isinstance(furnishings, dict):
|
||||
_furn_vals = []
|
||||
for fk in ("materials", "furniture", "shelving", "seating"):
|
||||
fv = furnishings.get(fk)
|
||||
if isinstance(fv, list):
|
||||
_furn_vals.extend(str(x) for x in fv if x)
|
||||
elif isinstance(fv, str) and fv:
|
||||
_furn_vals.append(fv)
|
||||
if _furn_vals:
|
||||
appearance_parts.append("陈设:" + "、".join(_furn_vals[:4]))
|
||||
elif isinstance(furnishings, list) and furnishings:
|
||||
appearance_parts.append("陈设:" + "、".join(str(f) for f in furnishings[:4] if f))
|
||||
if fj.get("cleanliness"):
|
||||
appearance_parts.append(str(fj["cleanliness"]))
|
||||
appearance = ";".join(appearance_parts) if appearance_parts else "门店环境"
|
||||
# key_features: 招牌细节+海报+外部物品+视觉元素+陈列商品+环境物件
|
||||
kf = []
|
||||
for arr_key in (
|
||||
"signage_details",
|
||||
"signage_posters",
|
||||
"exterior_items",
|
||||
"visual_elements",
|
||||
"products_on_display",
|
||||
"environment_objects",
|
||||
):
|
||||
arr = fj.get(arr_key) or []
|
||||
if not isinstance(arr, list):
|
||||
arr = [arr]
|
||||
for item in arr:
|
||||
if item and str(item) not in kf and len(str(item)) <= 40:
|
||||
kf.append(str(item))
|
||||
prods_vis = fj.get("product_categories_visible") or []
|
||||
if isinstance(prods_vis, list):
|
||||
for c in prods_vis[:3]:
|
||||
if c and str(c) not in kf:
|
||||
kf.append(str(c))
|
||||
promo = fj.get("promotion_elements") or []
|
||||
if isinstance(promo, list) and promo:
|
||||
kf.append("促销活动:" + "、".join(str(p) for p in promo[:2]))
|
||||
if has_person:
|
||||
for pf in _extract_person_features(fj):
|
||||
if pf not in kf:
|
||||
kf.append(pf)
|
||||
if text_on_package:
|
||||
kf.append("文字:" + "/".join(text_on_package[:3]))
|
||||
kf = kf[:8] or ["门店场景"]
|
||||
atmosphere = fj.get("atmosphere") or fj.get("mood") or mood
|
||||
# portrait_prompt: 品牌+store_type+主色+atmosphere+scene+核心视觉元素
|
||||
pieces = []
|
||||
if brand != "无法判断":
|
||||
pieces.append(brand)
|
||||
pieces.append(store_type)
|
||||
if color_names:
|
||||
pieces.append("、".join(color_names[:3]) + "配色")
|
||||
if atmosphere:
|
||||
pieces.append(atmosphere)
|
||||
if scene and scene != "通用":
|
||||
pieces.append(scene)
|
||||
visual = fj.get("visual_elements") or []
|
||||
if isinstance(visual, list) and visual:
|
||||
pieces.append("、".join(str(v) for v in visual[:3] if v))
|
||||
pieces.append("门店实拍")
|
||||
portrait_prompt = ",".join(p for p in pieces if p)
|
||||
if len(portrait_prompt) > 200:
|
||||
portrait_prompt = portrait_prompt[:200].rstrip(",")
|
||||
summary = f"{brand} {store_type}" if brand != "无法判断" else f"{store_type}场景"
|
||||
return {
|
||||
"name": name[:30],
|
||||
"brand": str(brand)[:30],
|
||||
"category": category,
|
||||
"appearance": appearance[:300],
|
||||
"packaging": "门店场景无包装",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": atmosphere,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": summary[:40],
|
||||
"has_person": bool(fj.get("has_person", False)),
|
||||
"_source": "v2_fast_json_v6_store",
|
||||
}
|
||||
|
||||
# 场景类(纯场景图,无产品/人物)
|
||||
if vtype == "scene":
|
||||
main_subject = fj.get("main_subject") or ""
|
||||
scene_type = fj.get("scene_type") or "场景图"
|
||||
name = main_subject or scene_type
|
||||
brand = "无法判断"
|
||||
category = scene_type or "场景图"
|
||||
# appearance: 空间布局 + 关键元素
|
||||
app_parts = []
|
||||
if fj.get("spatial_layout"):
|
||||
app_parts.append(str(fj["spatial_layout"]))
|
||||
key_elements = fj.get("key_elements") or []
|
||||
if not isinstance(key_elements, list):
|
||||
key_elements = [key_elements]
|
||||
if key_elements:
|
||||
app_parts.append("关键元素:" + "、".join(str(k) for k in key_elements[:4] if k))
|
||||
if fj.get("lighting"):
|
||||
app_parts.append("光线:" + str(fj["lighting"]))
|
||||
appearance = ";".join(app_parts) if app_parts else "场景环境"
|
||||
# key_features: key_elements + environment_objects(去重,最多8)
|
||||
kf = []
|
||||
for k in key_elements:
|
||||
if k and str(k) not in kf and len(str(k)) <= 40:
|
||||
kf.append(str(k))
|
||||
env_obj = fj.get("environment_objects") or []
|
||||
if not isinstance(env_obj, list):
|
||||
env_obj = [env_obj]
|
||||
for o in env_obj:
|
||||
if o and str(o) not in kf and len(str(o)) <= 40:
|
||||
kf.append(str(o))
|
||||
if has_person:
|
||||
for pf in _extract_person_features(fj):
|
||||
if pf not in kf:
|
||||
kf.append(pf)
|
||||
if text_on_package:
|
||||
kf.append("文字:" + "/".join(text_on_package[:3]))
|
||||
kf = kf[:8] or ["场景元素"]
|
||||
scene_name = fj.get("scene") or scene_type
|
||||
scene_mood = fj.get("mood") or mood
|
||||
packaging = "场景无包装"
|
||||
# portrait_prompt: main_subject + key_elements + lighting + atmosphere + composition + style
|
||||
pieces = []
|
||||
if main_subject:
|
||||
pieces.append(main_subject)
|
||||
if key_elements:
|
||||
pieces.append("、".join(str(k) for k in key_elements[:3] if k))
|
||||
if fj.get("lighting"):
|
||||
pieces.append(str(fj["lighting"]) + "光线")
|
||||
if scene_mood:
|
||||
pieces.append(scene_mood + "氛围")
|
||||
if fj.get("composition"):
|
||||
pieces.append(str(fj["composition"]))
|
||||
if fj.get("style"):
|
||||
pieces.append(str(fj["style"]))
|
||||
pieces.append("场景实拍,摄影级画质")
|
||||
portrait_prompt = ",".join(p for p in pieces if p)
|
||||
if len(portrait_prompt) > 200:
|
||||
portrait_prompt = portrait_prompt[:200].rstrip(",")
|
||||
summary = main_subject or scene_type
|
||||
return {
|
||||
"name": str(name)[:30],
|
||||
"brand": brand,
|
||||
"category": str(category)[:20],
|
||||
"appearance": appearance[:300],
|
||||
"packaging": packaging,
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene_name,
|
||||
"mood": scene_mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": str(summary)[:40],
|
||||
"has_person": bool(fj.get("has_person", False)),
|
||||
"_source": "v2_fast_json_v6_scene",
|
||||
}
|
||||
|
||||
# other 兜底:当 description 为空时,用 environment_objects/scene 拼基本描述
|
||||
desc = fj.get("description") or ""
|
||||
if not desc.strip():
|
||||
env_obj = fj.get("environment_objects") or []
|
||||
if not isinstance(env_obj, list):
|
||||
env_obj = [env_obj]
|
||||
env_names = [str(e) for e in env_obj if e and len(str(e)) <= 25][:5]
|
||||
scene_text = fj.get("scene") or ""
|
||||
if env_names:
|
||||
desc = f"{scene_text}场景中" + "、".join(env_names) if scene_text else "、".join(env_names)
|
||||
elif scene_text and scene_text != "通用":
|
||||
desc = f"{scene_text}场景"
|
||||
else:
|
||||
desc = "未识别"
|
||||
kf_other = []
|
||||
if desc != "未识别":
|
||||
kf_other.append(desc[:30])
|
||||
env_obj = fj.get("environment_objects") or []
|
||||
if not isinstance(env_obj, list):
|
||||
env_obj = [env_obj]
|
||||
for o in env_obj:
|
||||
if o and str(o) not in kf_other and len(str(o)) <= 40:
|
||||
kf_other.append(str(o))
|
||||
kf_other = kf_other[:8] or ["无法判断"]
|
||||
return {
|
||||
"name": desc[:30],
|
||||
"brand": "无法判断",
|
||||
"category": "非产品图",
|
||||
"appearance": desc[:200],
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf_other,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": f"{scene},{mood}氛围,{desc}"[:200],
|
||||
"summary": desc[:40],
|
||||
"has_person": False,
|
||||
"_source": "v2_fast_json_v4_other",
|
||||
}
|
||||
|
||||
|
||||
def _assemble_old(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"""旧扁平schema(兼容存量prompt或pro兜底输出)"""
|
||||
portrait_prompt = _build_portrait_prompt_old(fj)
|
||||
name = _infer_name_old(fj, ocr_texts)
|
||||
brand = _infer_brand_old(fj, ocr_texts)
|
||||
category = _infer_category_old(fj)
|
||||
appearance = _build_appearance_old(fj)
|
||||
key_features = _build_key_features_old(fj, ocr_texts)
|
||||
scene = fj.get("scene") or "通用"
|
||||
mood = fj.get("mood") or ""
|
||||
packaging = "无法判断"
|
||||
text_on_package = ocr_texts[:8]
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or "穿搭"
|
||||
style = fj.get("style") or ""
|
||||
summary = f"{style}{up}" if style and style not in up else up
|
||||
elif brand != "无法判断" and name != brand:
|
||||
summary = f"{brand} {name}"
|
||||
else:
|
||||
summary = name
|
||||
return {
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
@@ -293,15 +920,5 @@ def assemble_result(
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary,
|
||||
"_source": "v2_fast_json",
|
||||
"has_person": bool(fj.get("has_person", False)),
|
||||
}
|
||||
|
||||
|
||||
def _build_summary(fj: dict, name: str, brand: str, category: str) -> str:
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or "穿搭"
|
||||
style = fj.get("style") or ""
|
||||
base = f"{style}{up}" if style and style not in up else up
|
||||
return base
|
||||
if brand != "无法判断" and name != brand:
|
||||
return f"{brand} {name}"
|
||||
return name
|
||||
|
||||
@@ -23,10 +23,10 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 超时(可通过环境变量覆盖)
|
||||
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "12"))
|
||||
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "12"))
|
||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_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"))
|
||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "25"))
|
||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
|
||||
|
||||
_FALLBACK_RESULT = {
|
||||
"name": "未识别",
|
||||
@@ -58,27 +58,29 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
||||
|
||||
fj_result: dict[str, Any] | None = None
|
||||
ocr_result: list[str] = []
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
||||
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||
try:
|
||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||
try:
|
||||
res = fut.result(timeout=1)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||
continue
|
||||
if fut is f_fj and isinstance(res, dict):
|
||||
fj_result = res
|
||||
elif fut is f_ocr and isinstance(res, list):
|
||||
ocr_result = res
|
||||
except TimeoutError:
|
||||
for f in (f_fj, f_ocr):
|
||||
if not f.done():
|
||||
f.cancel()
|
||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||
|
||||
fast_elapsed = time.time() - t0
|
||||
fast_elapsed = 0.0
|
||||
pool = ThreadPoolExecutor(max_workers=2)
|
||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
||||
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||
try:
|
||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||
try:
|
||||
res = fut.result(timeout=1)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||
continue
|
||||
if fut is f_fj and isinstance(res, dict):
|
||||
fj_result = res
|
||||
elif fut is f_ocr and isinstance(res, list):
|
||||
ocr_result = res
|
||||
except TimeoutError:
|
||||
for f in (f_fj, f_ocr):
|
||||
if not f.done():
|
||||
f.cancel()
|
||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||
finally:
|
||||
fast_elapsed = time.time() - t0
|
||||
pool.shutdown(wait=False) # 不等待未完成的线程,避免计时膨胀
|
||||
|
||||
if fj_result:
|
||||
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""VLM 返回文本的稳健 JSON 提取工具。
|
||||
|
||||
背景:复杂门店图 VLM 输出经常被 max_tokens 截断(finish_reason=length),
|
||||
json.loads 失败后整个结果被丢弃,导致"未识别"。本工具提供:
|
||||
1. markdown 代码块剥离(含只开不闭的截断场景)
|
||||
2. 最外层 { } 切片
|
||||
3. 非法控制字符清理
|
||||
4. 直接 json.loads
|
||||
5. 截断 JSON 括号/引号栈补全修复
|
||||
6. 尾部逐字符截断重试(去除最后一个不完整 token 后修复)
|
||||
|
||||
成功返回 dict;截断修复产物带 _partial=True 标记;彻底失败返回 None。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_CODE_FENCE_RE = re.compile(r"^```(?:json)?\s*\n?(.*?)\n?```\s*$", re.DOTALL)
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
m = _CODE_FENCE_RE.match(s)
|
||||
if m:
|
||||
return m.group(1).strip()
|
||||
# 兼容开头 ```json 但结尾无 ```(截断场景)
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
|
||||
def _repair_truncated_json(text: str) -> str:
|
||||
"""尝试补全被截断的JSON:维护 bracket/quote 栈,在末尾补闭合符。"""
|
||||
stack: list[str] = []
|
||||
in_string = False
|
||||
escape = False
|
||||
for ch in text:
|
||||
if escape:
|
||||
escape = False
|
||||
continue
|
||||
if ch == "\\" and in_string:
|
||||
escape = True
|
||||
continue
|
||||
if ch == '"':
|
||||
in_string = not in_string
|
||||
continue
|
||||
if in_string:
|
||||
continue
|
||||
if ch in "{[":
|
||||
stack.append(ch)
|
||||
elif ch == "}":
|
||||
if stack and stack[-1] == "{":
|
||||
stack.pop()
|
||||
elif ch == "]":
|
||||
if stack and stack[-1] == "[":
|
||||
stack.pop()
|
||||
repair = ""
|
||||
if in_string:
|
||||
repair += '"'
|
||||
for opener in reversed(stack):
|
||||
repair += "}" if opener == "{" else "]"
|
||||
if repair:
|
||||
logger.info(
|
||||
"[json_utils] 截断JSON修复: 补全%d个闭合符 in_string=%s",
|
||||
len(repair),
|
||||
in_string,
|
||||
)
|
||||
return text + repair
|
||||
|
||||
|
||||
def _clean_invalid_chars(text: str) -> str:
|
||||
"""清理JSON中非法的控制字符(tab/newline 之外的 0x00-0x1f 段)。"""
|
||||
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]", "", text)
|
||||
|
||||
|
||||
def extract_json_object(text: str) -> dict | None:
|
||||
"""从VLM返回文本中稳健提取JSON对象。
|
||||
|
||||
返回 dict 或 None。成功的 dict 可能带 _partial=True 标记,
|
||||
表示原始文本被截断、经括号补全后得到的产物。
|
||||
"""
|
||||
if not text or not isinstance(text, str):
|
||||
return None
|
||||
# 1. 剥离 markdown
|
||||
text = _strip_code_fence(text)
|
||||
# 2. 找最外层 { }
|
||||
lpos = text.find("{")
|
||||
if lpos < 0:
|
||||
return None
|
||||
rpos = text.rfind("}")
|
||||
if rpos > lpos:
|
||||
text = text[lpos : rpos + 1]
|
||||
else:
|
||||
# 截断场景:无任何闭合 },取到末尾交给修复器
|
||||
text = text[lpos:]
|
||||
# 3. 清理非法控制字符
|
||||
text = _clean_invalid_chars(text)
|
||||
# 4. 直接 loads
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
return obj if isinstance(obj, dict) else None
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 5. 尝试截断修复
|
||||
repaired = _repair_truncated_json(text)
|
||||
try:
|
||||
obj = json.loads(repaired)
|
||||
if isinstance(obj, dict):
|
||||
obj["_partial"] = True
|
||||
return obj
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 6. 尾部逐字符截断重试(去除最后一个不完整 token)
|
||||
for _ in range(50):
|
||||
last_comma = repaired.rfind(",")
|
||||
last_brace = max(repaired.rfind("}"), repaired.rfind("]"))
|
||||
cut = max(last_comma, last_brace)
|
||||
if cut < 10:
|
||||
break
|
||||
repaired = repaired[: cut + 1]
|
||||
repaired = _repair_truncated_json(repaired)
|
||||
try:
|
||||
obj = json.loads(repaired)
|
||||
if isinstance(obj, dict):
|
||||
obj["_partial"] = True
|
||||
return obj
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
return None
|
||||
@@ -1,114 +1,27 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 兜底路径:qwen3.7-plus(阿里云百炼/DashScope)单图调用。
|
||||
"""V2 兜底路径:image_analysis(默认 qwen-vl-plus 视觉模型,fallback qwen3.7-plus / DashScope)单图调用。
|
||||
|
||||
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
|
||||
设计要点:
|
||||
- 直接 httpx 直连 DashScope,不走 ai_client
|
||||
- enable_thinking=false + response_format=json_object
|
||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
||||
- enable_thinking=False + response_format=json_object
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- timeout=25s
|
||||
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
|
||||
- 返回 dict 字段与旧 _normalize() 兼容,下游零改动
|
||||
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
|
||||
- timeout=30s
|
||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import _prompt
|
||||
from . import _prompt, assembler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
_PRO_MODEL = "qwen3.7-plus"
|
||||
_DEFAULT_TIMEOUT = 25
|
||||
_DEFAULT_MAX_TOKENS = 800
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def _assemble_pp(obj: dict[str, Any]) -> str:
|
||||
"""从 JSON 字段组装 portrait_prompt(60-100字人物穿搭描述,给 Seedream t2i 用)。"""
|
||||
if not obj.get("has_person"):
|
||||
name = obj.get("product_name") or "商品"
|
||||
brand = obj.get("brand") or ""
|
||||
kf = obj.get("key_features") or []
|
||||
scene = obj.get("scene") or ""
|
||||
mood = obj.get("mood") or ""
|
||||
outfit = obj.get("outfit") or ""
|
||||
if outfit:
|
||||
return outfit
|
||||
pieces = []
|
||||
if brand:
|
||||
pieces.append(brand)
|
||||
pieces.append(str(name))
|
||||
if isinstance(kf, list):
|
||||
pieces.extend(str(x) for x in kf[:2] if x)
|
||||
if mood:
|
||||
pieces.append(str(mood) + "氛围")
|
||||
if scene:
|
||||
pieces.append(str(scene) + "场景")
|
||||
pieces.append("产品特写")
|
||||
p = ",".join(x for x in pieces if x)
|
||||
return p if len(p) >= 10 else "产品展示图,特写镜头"
|
||||
|
||||
parts: list[str] = []
|
||||
gender = obj.get("gender") or ""
|
||||
age = obj.get("age_range") or ""
|
||||
subj = ""
|
||||
if age == "儿童":
|
||||
subj = "小女孩" if gender == "女" else ("小男孩" if gender == "男" else "儿童")
|
||||
elif age == "青少年":
|
||||
subj = "少女" if gender == "女" else ("少年" if gender == "男" else "青少年")
|
||||
else:
|
||||
prefix_map = {"青年": "年轻", "中年": "中年", "老年": "老年"}
|
||||
gw = {"男": "男性", "女": "女性"}.get(gender, "")
|
||||
prefix = prefix_map.get(age, "")
|
||||
subj = (prefix + gw) if (prefix or gw) else "人物"
|
||||
parts.append(f"一位{subj}")
|
||||
|
||||
outfit = obj.get("outfit") or ""
|
||||
if outfit:
|
||||
parts.append(f"身着{outfit}")
|
||||
|
||||
hair = obj.get("hair") or ""
|
||||
if hair:
|
||||
parts.append(str(hair))
|
||||
|
||||
pose = obj.get("pose") or ""
|
||||
expr = obj.get("expression") or ""
|
||||
det = []
|
||||
if expr and expr not in ("自然", "平静"):
|
||||
det.append(f"神情{expr}")
|
||||
if pose and pose not in ("站立",):
|
||||
det.append(str(pose))
|
||||
if det:
|
||||
parts.append(",".join(det))
|
||||
|
||||
style_parts = []
|
||||
mood = obj.get("mood") or ""
|
||||
scene = obj.get("scene") or ""
|
||||
if mood:
|
||||
style_parts.append(str(mood))
|
||||
if scene and scene != "通用":
|
||||
style_parts.append(str(scene))
|
||||
if style_parts:
|
||||
parts.append("".join(style_parts) + "风格")
|
||||
else:
|
||||
parts.append("人像写真")
|
||||
|
||||
full = ",".join(p for p in parts if p)
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
_DEFAULT_TIMEOUT = 45
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
@@ -116,109 +29,87 @@ def call_pro_vlm(
|
||||
idx: int,
|
||||
*,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置。"""
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
|
||||
api_key = _api_key()
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="fallback")
|
||||
if not client or not client.is_available:
|
||||
logger.warning("[vision.v2] pro vision client 不可用,跳过")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": _PRO_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.3,
|
||||
"max_tokens": _DEFAULT_MAX_TOKENS,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
try:
|
||||
r = httpx.post(
|
||||
f"{_BASE_URL}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if r.status_code != 200:
|
||||
logger.warning("[vision.v2] pro HTTP %d elapsed=%.1fs body=%s", r.status_code, elapsed, r.text[:200])
|
||||
return None
|
||||
data = r.json()
|
||||
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] pro 返回空 elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
usage = data.get("usage") or {}
|
||||
reasoning_tokens = usage.get("reasoning_tokens", 0)
|
||||
ctd = usage.get("completion_tokens_details") or {}
|
||||
if not reasoning_tokens:
|
||||
reasoning_tokens = ctd.get("reasoning_tokens", 0)
|
||||
logger.info(
|
||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_PRO_MODEL,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
reasoning_tokens,
|
||||
)
|
||||
s = raw.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
l, rr = s.find("{"), s.rfind("}")
|
||||
if l >= 0 and rr > l:
|
||||
s = s[l : rr + 1]
|
||||
try:
|
||||
obj = json.loads(s)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("[vision.v2] pro JSON 解析失败 head=%s", raw[:200])
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
return None
|
||||
|
||||
pp = _assemble_pp(obj)
|
||||
kf = obj.get("key_features")
|
||||
if not isinstance(kf, list):
|
||||
kf = [str(kf)] if kf else ["无法判断"]
|
||||
else:
|
||||
kf = [str(x) for x in kf if x] or ["无法判断"]
|
||||
|
||||
name = obj.get("product_name") or "未识别"
|
||||
if obj.get("has_person") and (not name or name == "未识别"):
|
||||
name = obj.get("outfit") or "人物穿搭"
|
||||
brand = obj.get("brand") or "无法判断"
|
||||
category = obj.get("category") or ("服饰" if obj.get("has_person") else "非产品图")
|
||||
return {
|
||||
"name": str(name),
|
||||
"brand": str(brand),
|
||||
"category": str(category),
|
||||
"appearance": str(obj.get("outfit") or "无法判断"),
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": kf[:6],
|
||||
"scene": str(obj.get("scene") or "通用"),
|
||||
"mood": str(obj.get("mood") or ""),
|
||||
"portrait_prompt": pp,
|
||||
"summary": str(name),
|
||||
"_source": "vlm_pro",
|
||||
call_kwargs: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"images": None, # 图片已在 messages 中
|
||||
"temperature": 0.3,
|
||||
"timeout": timeout,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
# 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
|
||||
|
||||
raw = None
|
||||
obj = None
|
||||
for _outer in range(2):
|
||||
kw = dict(call_kwargs)
|
||||
if _outer == 1:
|
||||
kw.pop("response_format", None)
|
||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||
cont = [dict(c) for c in list(msgs2[1]["content"])]
|
||||
cont[-1] = {"type": "text", "text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。"}
|
||||
msgs2[1] = {"role": "user", "content": cont}
|
||||
kw["messages"] = msgs2
|
||||
raw = client.vision_completion(**kw)
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] pro 返回空 outer=%s", _outer)
|
||||
continue
|
||||
obj = extract_json_object(raw)
|
||||
if obj is not None:
|
||||
break
|
||||
logger.warning("[vision.v2] pro 非JSON(100字) outer=%s: %s", _outer, raw[:100])
|
||||
|
||||
elapsed = time.time() - t0
|
||||
if obj is None:
|
||||
logger.warning("[vision.v2] pro 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
if obj.get("_partial"):
|
||||
logger.warning("[vision.v2] pro 返回截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||
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, [])
|
||||
result["_source"] = "vlm_pro"
|
||||
result["_fallback_used"] = True
|
||||
return result
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
|
||||
@@ -1,22 +1,20 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 快速路径:qwen3.8-flash(阿里云百炼/DashScope)强约束 JSON-only 调用。
|
||||
"""V2 快速路径:image_analysis capability(默认 qwen-vl-plus 视觉模型 / DashScope)强约束 JSON-only 调用。
|
||||
|
||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
||||
设计要点:
|
||||
- 直接用 httpx 发最小 payload 到 DashScope OpenAI 兼容 endpoint,不走 ai_client 包装
|
||||
- enable_thinking=false 关闭推理链(reasoning 是延迟主因)
|
||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
||||
- enable_thinking=False 关闭推理链(reasoning 是延迟主因)
|
||||
- response_format=json_object 强约束JSON输出
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- max_tokens=350、temperature=0.1(稳定输出 JSON)
|
||||
- timeout=12s(失败由外层走 pro 兜底)
|
||||
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
|
||||
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
|
||||
- temperature=0.1(稳定输出 JSON)
|
||||
- timeout=15s(失败由外层走 pro 兜底)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
@@ -24,124 +22,99 @@ from . import _prompt
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# DashScope OpenAI 兼容 endpoint
|
||||
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
_FAST_MODEL = "qwen3.8-flash"
|
||||
_DEFAULT_TIMEOUT = 12
|
||||
_DEFAULT_MAX_TOKENS = 350
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
_DEFAULT_TIMEOUT = 20
|
||||
|
||||
|
||||
def call_fast_json(
|
||||
img_url: str,
|
||||
*,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int = _DEFAULT_MAX_TOKENS,
|
||||
max_tokens: int | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""调用 qwen3.8-flash 返回结构化 dict;失败/非 JSON 返回 None。"""
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
"""调用 vision client 返回结构化 dict;失败/非 JSON 返回 None。
|
||||
|
||||
api_key = _api_key()
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 fast_json")
|
||||
max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置;
|
||||
显式传入时作为覆盖。
|
||||
"""
|
||||
t0 = time.time()
|
||||
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="primary")
|
||||
if not client or not client.is_available:
|
||||
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
|
||||
url = f"{_BASE_URL}/chat/completions"
|
||||
payload: dict[str, Any] = {
|
||||
"model": _FAST_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.1,
|
||||
"max_tokens": max_tokens,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
try:
|
||||
resp = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code == 400 and "enable_thinking" in resp.text[:300].lower():
|
||||
logger.warning("[vision.v2] fast_json HTTP 400 thinking 参数不兼容,重试 elapsed=%.1fs", elapsed)
|
||||
payload.pop("enable_thinking", None)
|
||||
resp = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code != 200:
|
||||
call_kwargs: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"images": None, # 图片已在 messages 中
|
||||
"temperature": 0.1,
|
||||
"timeout": timeout,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
if max_tokens is not None:
|
||||
call_kwargs["max_tokens"] = max_tokens
|
||||
|
||||
# 双重防护:第1次正常调用;第2次去掉json_object强约束(部分模型在该约束下
|
||||
# 反而幻觉),并加严格指令。解析全部走 json_utils,截断partial产物可用。
|
||||
from .json_utils import extract_json_object
|
||||
|
||||
raw = None
|
||||
obj = None
|
||||
for _outer in range(2):
|
||||
kw = dict(call_kwargs)
|
||||
if _outer == 1:
|
||||
kw.pop("response_format", None)
|
||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||
cont = list(msgs2[1]["content"])
|
||||
cont = [dict(c) for c in cont]
|
||||
cont[-1] = {"type": "text", "text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。"}
|
||||
msgs2[1] = {"role": "user", "content": cont}
|
||||
kw["messages"] = msgs2
|
||||
raw = client.vision_completion(**kw)
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] fast_json 返回空 outer=%s", _outer)
|
||||
continue
|
||||
obj = extract_json_object(raw)
|
||||
if obj is not None:
|
||||
break
|
||||
logger.warning(
|
||||
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
|
||||
"[vision.v2] fast_json 非JSON(100字) outer=%s: %s",
|
||||
_outer,
|
||||
raw[:100],
|
||||
)
|
||||
|
||||
elapsed = time.time() - t0
|
||||
if obj is None:
|
||||
logger.warning("[vision.v2] fast_json 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
data = resp.json()
|
||||
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] fast_json 返回空 elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
usage = data.get("usage") or {}
|
||||
reasoning_tokens = usage.get("reasoning_tokens", 0)
|
||||
ctd = usage.get("completion_tokens_details") or {}
|
||||
if not reasoning_tokens:
|
||||
reasoning_tokens = ctd.get("reasoning_tokens", 0)
|
||||
if obj.get("_partial"):
|
||||
logger.warning("[vision.v2] fast_json 返回截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_FAST_MODEL,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
reasoning_tokens,
|
||||
)
|
||||
text = _strip_code_fence(raw)
|
||||
lpos, r = text.find("{"), text.rfind("}")
|
||||
if lpos >= 0 and r > lpos:
|
||||
text = text[lpos : r + 1]
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("[vision.v2] fast_json JSON 解析失败 elapsed=%.1fs head=%s", elapsed, raw[:200])
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
logger.warning("[vision.v2] fast_json 非 dict: %s", type(obj))
|
||||
return None
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 elapsed=%.1fs has_person=%s has_product=%s category=%s",
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs has_person=%s type=%s",
|
||||
client.model,
|
||||
elapsed,
|
||||
obj.get("has_person"),
|
||||
obj.get("has_product"),
|
||||
obj.get("category"),
|
||||
obj.get("type"),
|
||||
)
|
||||
return obj
|
||||
except Exception as e:
|
||||
|
||||
@@ -351,12 +351,25 @@ class CosyVoiceService:
|
||||
用于私有 bucket 下,将裸 URL 转为预签名 URL,
|
||||
确保 CosyVoice 服务器能下载参考音频.
|
||||
"""
|
||||
# 优先从 ai_router 获取 DB 配置
|
||||
_router_key, _router_url, _router_model = "", "", ""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
tts_client = ai_router.get_tts_client("tts")
|
||||
if tts_client and tts_client.is_available:
|
||||
_router_key = tts_client.api_key
|
||||
_router_url = tts_client.base_url
|
||||
_router_model = tts_client.model
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
settings = get_shared_settings()
|
||||
|
||||
self._api_key = api_key or settings.cosyvoice_api_key
|
||||
self._base_url = base_url or settings.cosyvoice_base_url
|
||||
self._model = model or settings.cosyvoice_model
|
||||
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "voice-enrollment")
|
||||
self._api_key = api_key or _router_key or settings.cosyvoice_api_key
|
||||
self._base_url = base_url or _router_url or settings.cosyvoice_base_url
|
||||
self._model = model or _router_model or settings.cosyvoice_model
|
||||
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "")
|
||||
self._audio_url_signer = audio_url_signer
|
||||
|
||||
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
|
||||
|
||||
@@ -110,10 +110,12 @@ _INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案
|
||||
|
||||
_INTENT_USER = """用户原始文案:{user_copy_text}
|
||||
所属行业:{industry}
|
||||
营销目的:{marketing_purpose}
|
||||
图片分析结果(供参考):
|
||||
{image_analysis}
|
||||
图片类型推断:{image_category_hint}
|
||||
|
||||
请理解用户意图,按标签格式输出。"""
|
||||
请理解用户意图,按标签格式输出。注意:theme和emotion_tone应与图片类型和营销目的匹配——门店类图片偏向"门店探店/到店体验",商品图偏向"好物分享/产品种草",人物图偏向"穿搭/人物故事"。"""
|
||||
|
||||
_INTENT_EXAMPLE = """<intent_summary>一款厨房去油污神器,喷一喷油污就掉</intent_summary>
|
||||
<core_messages>
|
||||
@@ -176,7 +178,7 @@ _FUSION_EXAMPLE = """<title>厨房重油污,别再用洗洁精硬擦了</title
|
||||
<segment duration_sec="4" image_index="0">39块钱625ml,厨房重油污的可以试一瓶</segment>
|
||||
</script_segments>
|
||||
<voiceover_script>这油污我真的忍很久了,用洗洁精擦半天都没用。后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净。39块钱625ml,厨房重油污的可以试一瓶。</voiceover_script>
|
||||
<overview_theme>厨房油污清洁好物分享</overview_theme>
|
||||
<overview_theme>厨房好物分享·产品种草</overview_theme>
|
||||
<scene_and_lighting>简洁明亮的厨房台面场景,自然光从窗户洒入,色调温暖柔和,突出产品白色瓶身与去油污对比效果。</scene_and_lighting>
|
||||
<word_count>58</word_count>
|
||||
<estimated_duration>13</estimated_duration>"""
|
||||
@@ -215,6 +217,8 @@ _STORYBOARD_USER = """目标时长:{duration}秒
|
||||
图片分析结果:
|
||||
{image_analysis}
|
||||
|
||||
重要:overview_theme 必须与图片实际内容和营销目的匹配。门店/餐饮/服务类图片用"门店探店·到店体验";商品图用"好物分享·产品种草";人物图用"穿搭分享·人物故事";场景图用"空间体验·场景氛围"。不要对所有图片都使用"好物分享"。
|
||||
|
||||
请按标签格式输出分镜。"""
|
||||
|
||||
_STORYBOARD_EXAMPLE = """<clips>
|
||||
|
||||
@@ -55,9 +55,14 @@ _LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
|
||||
class Reviewer:
|
||||
def __init__(self, client=None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = get_doubao_client()
|
||||
client = ai_router.get_llm_client("copy_review")
|
||||
except Exception:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
self.client = client
|
||||
|
||||
# ── 审核 ────────────────────────────────────────────────────────────
|
||||
@@ -96,6 +101,7 @@ class Reviewer:
|
||||
],
|
||||
temperature=0.2,
|
||||
max_tokens=1024,
|
||||
timeout=25,
|
||||
)
|
||||
if not raw:
|
||||
return None
|
||||
@@ -242,6 +248,7 @@ class Reviewer:
|
||||
],
|
||||
temperature=0.5,
|
||||
max_tokens=2048,
|
||||
timeout=25,
|
||||
)
|
||||
if not raw:
|
||||
return self._rule_fix(fusion, review)
|
||||
|
||||
+27
-35
@@ -80,54 +80,46 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
|
||||
cosyvoice_api_key: str = ""
|
||||
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
cosyvoice_model: str = "cosyvoice-v3-flash"
|
||||
cosyvoice_voice: str = "longxiaochun_v3" # 默认音色(v3 系列系统音色带 _v3 后缀)
|
||||
cosyvoice_base_url: str = ""
|
||||
cosyvoice_model: str = ""
|
||||
cosyvoice_voice: str = "longxiaochun_v3"
|
||||
cosyvoice_sample_rate: int = 22050
|
||||
cosyvoice_format: str = "mp3" # 输出格式:mp3/wav/pcm
|
||||
cosyvoice_format: str = "mp3"
|
||||
# 音色克隆模型名(固定为 voice-enrollment)
|
||||
cosyvoice_clone_model: str = "voice-enrollment"
|
||||
cosyvoice_clone_model: str = ""
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
# AI模型路由化:model/base_url 默认值清空,由 DB ai_models/ai_capability_configs 配置驱动。
|
||||
# 环境变量仍可覆盖(兼容旧部署);无任何配置时 ai_router fallback 提供最终默认值。
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # #2181: lite方舟侧100%超时,默认fast_model也走pro;方舟恢复lite后通过ENV DOUBAO_FAST_MODEL切回
|
||||
)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 45 # #2180: 方舟LLM高峰期响应6-8s,原30s太紧提到45s
|
||||
doubao_max_retries: int = 1 # #2180: timeout调大后一次调用就够,1次重试防偶发抖动;避免6次重试叠加到351s
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_use_lite: bool = True # #2188: lite恢复稳定,爆款视频默认lite-first提速(20-30s)
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
doubao_image_model: str = (
|
||||
"doubao-seedream-5-0-flash-260915" # #2173: 信任链 Seedream 改 flash 模型(实测 pro 46.5s→flash 13s;pro AI化图仍被Seedance拦截)
|
||||
)
|
||||
doubao_image_size: str = "1K" # #2173: 1K 已足够做 Seedance 参考图,2K 在 flash 下也 22s,1K 13s
|
||||
doubao_image_timeout: int = 60 # #2173: flash+1K 通常15s内,给60s余量
|
||||
doubao_trust_chain_enabled: bool = (
|
||||
True # #2173: 信任链总开关;若Seedream产物仍被Seedance拦截,可配 False 关闭直接t2v降级
|
||||
)
|
||||
doubao_model: str = ""
|
||||
doubao_fast_model: str = ""
|
||||
doubao_base_url: str = ""
|
||||
doubao_timeout: int = 45
|
||||
doubao_max_retries: int = 3
|
||||
doubao_vision_model: str = ""
|
||||
doubao_vision_lite_model: str = ""
|
||||
doubao_vision_use_lite: bool = True
|
||||
doubao_embedding_model: str = ""
|
||||
doubao_video_model: str = ""
|
||||
doubao_video_timeout: int = 600
|
||||
doubao_video_poll_interval: int = 10
|
||||
doubao_image_model: str = ""
|
||||
doubao_image_size: str = "1K"
|
||||
doubao_image_timeout: int = 60
|
||||
doubao_trust_chain_enabled: bool = True
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
|
||||
dashscope_base_url: str = ""
|
||||
dashscope_video_timeout: int = 900
|
||||
dashscope_video_poll_interval: int = 10
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
mediakit_base_url: str = ""
|
||||
mediakit_timeout: int = 60
|
||||
mediakit_cover_enabled: bool = False # 封面抽帧是否走MediaKit(默认false走本地ffmpeg+cv2,<2s完成)
|
||||
mediakit_cover_enabled: bool = False
|
||||
|
||||
# ── 积分/会员系统 (#1895) ────────────────────────────────────────────
|
||||
# 积分系统总开关(产品要求 #1895:暂停积分系统但保留全部代码/表/接口)。
|
||||
|
||||
@@ -191,13 +191,45 @@ class ViralVideoJob:
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_image_analyzed(self, **kwargs) -> None:
|
||||
if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING):
|
||||
"""阶段2入口:允许从 IMAGE_ANALYZED/PENDING 首次进入,也允许从 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。
|
||||
|
||||
重新生成时清空上一轮文案产物(copy_result/intent_result/storyboard/generated_copy_text),
|
||||
并重置 completed_at/result_video_url/error_msg,确保前端轮询能看到新的阶段2进度。
|
||||
"""
|
||||
_allowed = (
|
||||
ViralVideoStatus.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.PENDING,
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
)
|
||||
if self.status not in _allowed:
|
||||
raise ValueError(f"Cannot resume from {self.status} to copy-gen")
|
||||
_is_regen = self.status in (
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
)
|
||||
for k, v in kwargs.items():
|
||||
if hasattr(self, k) and v not in (None, "", []):
|
||||
setattr(self, k, v)
|
||||
if _is_regen:
|
||||
# 清空上一轮文案/视频产物,避免前端拿到旧数据
|
||||
self.intent_result = None
|
||||
self.copy_result = None
|
||||
self.storyboard = None
|
||||
self.generated_copy_text = ""
|
||||
self.result_video_url = ""
|
||||
self.current_stage = ""
|
||||
self.phase_message = ""
|
||||
self.error_msg = ""
|
||||
self.completed_at = None
|
||||
self.heartbeat_at = None
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
self.updated_at = _now
|
||||
|
||||
def resume_from_copy_generated(self, edited_copy: str | None = None) -> None:
|
||||
"""阶段2->阶段3:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
|
||||
@@ -206,14 +238,20 @@ class ViralVideoJob:
|
||||
if edited_copy and isinstance(self.copy_result, dict):
|
||||
self.copy_result = {**self.copy_result, "voiceover_script": edited_copy}
|
||||
self.generated_copy_text = edited_copy
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
self.updated_at = _now
|
||||
|
||||
def resume_from_confirm(self) -> None:
|
||||
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise ValueError(f"Cannot resume from {self.status}")
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
self.updated_at = _now
|
||||
|
||||
def mark_completed(self, video_url: str) -> None:
|
||||
self.status = ViralVideoStatus.COMPLETED
|
||||
|
||||
+107
-27
@@ -170,14 +170,35 @@ class DoubaoClient:
|
||||
未配置 API Key 时 is_available 为 False,调用方应降级处理。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str = "",
|
||||
base_url: str = "",
|
||||
model: str = "",
|
||||
timeout: int = 0,
|
||||
max_retries: int = 0,
|
||||
max_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
extra_params: dict | None = None,
|
||||
provider: str = "volcengine",
|
||||
) -> None:
|
||||
settings = get_shared_settings()
|
||||
self.api_key: str = settings.doubao_api_key
|
||||
self.model: str = settings.doubao_model
|
||||
self.base_url: str = settings.doubao_base_url.rstrip("/")
|
||||
self.timeout: int = settings.doubao_timeout
|
||||
self.max_retries: int = settings.doubao_max_retries
|
||||
self.provider: str = provider
|
||||
if provider == "dashscope":
|
||||
self.api_key: str = api_key or getattr(settings, "dashscope_api_key", "")
|
||||
self.model: str = model or getattr(settings, "dashscope_model", "")
|
||||
self.base_url: str = (base_url or getattr(settings, "dashscope_base_url", "")).rstrip("/")
|
||||
else: # volcengine (default)
|
||||
self.api_key = api_key or settings.doubao_api_key
|
||||
self.model = model or settings.doubao_model
|
||||
self.base_url = (base_url or settings.doubao_base_url).rstrip("/")
|
||||
self.timeout: int = timeout or settings.doubao_timeout
|
||||
self.max_retries: int = max_retries or settings.doubao_max_retries
|
||||
self.max_tokens: int | None = max_tokens
|
||||
self.temperature: float | None = temperature
|
||||
self.extra_params: dict = extra_params or {}
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
self.last_finish_reason: str = ""
|
||||
self.vision_lite_model: str = settings.doubao_vision_lite_model
|
||||
self.fast_model: str = settings.doubao_fast_model
|
||||
self.embedding_model: str = settings.doubao_embedding_model
|
||||
@@ -190,6 +211,13 @@ class DoubaoClient:
|
||||
# 最近一次图片生成的详细错误,供上层读取
|
||||
self.last_image_error: dict = {}
|
||||
|
||||
def _resolve_timeout(self, timeout) -> "httpx.Timeout":
|
||||
"""将整数超时转为 httpx.Timeout,区分 connect/read/write/pool,避免 read 卡到 TCP 120s 默认值."""
|
||||
if isinstance(timeout, httpx.Timeout):
|
||||
return timeout
|
||||
t = int(timeout) if timeout else 60
|
||||
return httpx.Timeout(connect=10, read=max(t, 10), write=10, pool=5)
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
if not self.is_available or not text or not text.strip():
|
||||
@@ -240,16 +268,17 @@ class DoubaoClient:
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int = 1024,
|
||||
max_tokens: int | None = None,
|
||||
model: str | None = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs,
|
||||
) -> Optional[str]:
|
||||
"""调用 Chat Completion 接口.
|
||||
|
||||
Args:
|
||||
messages: 对话消息列表,[{"role": "user"/"system"/"assistant", "content": "..."}]
|
||||
temperature: 采样温度,0-2,默认0.7
|
||||
max_tokens: 最大生成token数,默认1024
|
||||
max_tokens: 最大生成token数,默认 None(使用实例 self.max_tokens DB 配置,兜底 1024)
|
||||
|
||||
Returns:
|
||||
模型返回的文本内容,失败返回 None
|
||||
@@ -262,18 +291,24 @@ class DoubaoClient:
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
"max_tokens": effective_max_tokens,
|
||||
}
|
||||
# 合并实例级额外参数和调用方传入的额外参数
|
||||
if self.extra_params:
|
||||
payload.update(self.extra_params)
|
||||
if kwargs:
|
||||
payload.update(kwargs)
|
||||
|
||||
last_error: Optional[Exception] = None
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
_req_timeout = timeout if timeout is not None else self.timeout
|
||||
_req_timeout = self._resolve_timeout(timeout if timeout is not None else self.timeout)
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers=headers,
|
||||
@@ -282,7 +317,24 @@ class DoubaoClient:
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
|
||||
if finish_reason == "length" and attempt < self.max_retries:
|
||||
# 输出被 max_tokens 截断:2.0x 扩容后重试(计入 max_retries,不额外增加)
|
||||
old_max = int(payload["max_tokens"])
|
||||
new_max = int(old_max * 2)
|
||||
payload["max_tokens"] = new_max
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
|
||||
old_max,
|
||||
new_max,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
self.last_finish_reason = finish_reason
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%d",
|
||||
@@ -291,6 +343,7 @@ class DoubaoClient:
|
||||
data.get("usage", {}).get("completion_tokens", 0),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
_req_timeout,
|
||||
)
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
@@ -314,20 +367,21 @@ class DoubaoClient:
|
||||
self,
|
||||
messages: list[dict],
|
||||
images: list[str] | None = None,
|
||||
max_tokens: int = 2048,
|
||||
max_tokens: int | None = None,
|
||||
temperature: float = 0.3,
|
||||
timeout: int | None = None,
|
||||
model: str | None = None,
|
||||
**kwargs,
|
||||
) -> Optional[str]:
|
||||
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
|
||||
|
||||
将 images 附加到最后一条 user message 的 content 中,
|
||||
使用 vision_model(默认 doubao-1-5-vision-pro-250328)。
|
||||
使用构造函数传入的 self.model(DB capability 绑定的视觉模型,默认 qwen-vl-plus)。
|
||||
|
||||
Args:
|
||||
messages: 对话消息列表。最后一条 user message 会被注入图片内容。
|
||||
images: 图片列表,支持 base64 data URI 或 HTTP(S) URL。
|
||||
max_tokens: 最大生成 token 数,默认 2048。
|
||||
max_tokens: 最大生成 token 数,默认 None(使用实例 self.max_tokens DB 配置,兜底 2048)。
|
||||
temperature: 采样温度,默认 0.3(视觉任务偏低更稳定)。
|
||||
timeout: 单次请求超时秒数,不传则使用默认 self.timeout。
|
||||
|
||||
@@ -367,14 +421,19 @@ class DoubaoClient:
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 2048)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.vision_model,
|
||||
"model": model or self.model,
|
||||
"messages": vision_messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
"max_tokens": effective_max_tokens,
|
||||
}
|
||||
if self.extra_params:
|
||||
payload.update(self.extra_params)
|
||||
if kwargs:
|
||||
payload.update(kwargs)
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
req_timeout = self._resolve_timeout(timeout or self.timeout)
|
||||
last_error: Optional[Exception] = None
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
@@ -387,7 +446,24 @@ class DoubaoClient:
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
|
||||
if finish_reason == "length" and attempt < self.max_retries:
|
||||
# 视觉输出被 max_tokens 截断:2.0x 扩容后重试(计入 max_retries)
|
||||
old_max = int(payload["max_tokens"])
|
||||
new_max = int(old_max * 2)
|
||||
payload["max_tokens"] = new_max
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"视觉输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
|
||||
old_max,
|
||||
new_max,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
self.last_finish_reason = finish_reason
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] vision_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d",
|
||||
@@ -589,28 +665,32 @@ class DoubaoClient:
|
||||
# 信任链只作用于 doubao provider;DashScope(Wan) 保持原行为。
|
||||
trust_chain_applied = False
|
||||
if provider == "doubao" and getattr(self, "trust_chain_enabled", True) and pre_trusted_images:
|
||||
# #2220: 稀疏列表模式——pre_trusted_images 与 raw_portrait_urls 等长,
|
||||
# None 位保留原图,非 None 位用 AI 人像替换。
|
||||
raw_portrait_urls: list[str] = []
|
||||
if image_url:
|
||||
raw_portrait_urls.append(image_url)
|
||||
for u in ref_imgs:
|
||||
if u not in raw_portrait_urls:
|
||||
raw_portrait_urls.append(u)
|
||||
trusted_urls: list[str] = []
|
||||
if len(pre_trusted_images) >= 1:
|
||||
trusted_urls = list(pre_trusted_images)
|
||||
_n_trusted = sum(1 for _x in pre_trusted_images if _x)
|
||||
if _n_trusted >= 1 and len(pre_trusted_images) >= len(raw_portrait_urls):
|
||||
merged: list[str] = []
|
||||
for _i, _orig in enumerate(raw_portrait_urls):
|
||||
_ai = pre_trusted_images[_i] if _i < len(pre_trusted_images) else None
|
||||
merged.append(str(_ai) if _ai else _orig)
|
||||
trust_chain_applied = True
|
||||
logger.info(
|
||||
"[trust-chain] 使用预热t2i结果 %d 张,替换原参考图走 reference_image 模式(原n=%d)",
|
||||
len(trusted_urls),
|
||||
"[trust-chain] 稀疏替换 %d/%d 张为AI人像(场景/商品图保留原图),走reference_image模式",
|
||||
_n_trusted,
|
||||
len(raw_portrait_urls),
|
||||
)
|
||||
if trust_chain_applied and trusted_urls:
|
||||
# 替换:原 image_url 用第一张 AI 图,ref_imgs 用剩余
|
||||
if image_url and trusted_urls:
|
||||
image_url = trusted_urls[0]
|
||||
ref_imgs = trusted_urls[1:] if len(trusted_urls) > 1 else []
|
||||
# 替换:image_url 用第一张(可能是AI或原图),ref_imgs 用剩余
|
||||
if image_url and merged:
|
||||
image_url = merged[0]
|
||||
ref_imgs = merged[1:] if len(merged) > 1 else []
|
||||
else:
|
||||
ref_imgs = trusted_urls
|
||||
ref_imgs = merged
|
||||
# ─────────────────────────────────────────────────────────────────
|
||||
|
||||
# 判断任务模式:
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
"""AI 配置版本号管理 — Redis 通知机制.
|
||||
|
||||
admin 后台修改 ai_models / ai_capability_configs 后调用 bump_version(),
|
||||
SaaS 端 AIRouter 每次取配置前比对版本号,变了才重新查 DB。
|
||||
|
||||
Redis key: xiaoxia:ai_config:version = 时间戳字符串
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_REDIS_KEY = "xiaoxia:ai_config:version"
|
||||
|
||||
|
||||
def _get_redis_client():
|
||||
"""获取 Redis 客户端(复用 Celery broker 连接)."""
|
||||
try:
|
||||
import redis as _redis
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
settings = get_shared_settings()
|
||||
redis_url = getattr(settings, "redis_url", None) or getattr(
|
||||
settings, "celery_broker_url", "redis://localhost:6379/0"
|
||||
)
|
||||
return _redis.Redis.from_url(redis_url, decode_responses=True, socket_timeout=2)
|
||||
except Exception as e:
|
||||
logger.warning("AI config version: Redis 客户端初始化失败: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def bump_version() -> str:
|
||||
"""写入新版本号(当前时间戳),返回版本号字符串。失败返回空串。"""
|
||||
r = _get_redis_client()
|
||||
if r is None:
|
||||
logger.warning("AI config bump_version: Redis 不可用,跳过版本号更新")
|
||||
return ""
|
||||
try:
|
||||
ver = str(int(time.time() * 1000))
|
||||
r.set(_REDIS_KEY, ver)
|
||||
logger.info("AI config version bumped to %s", ver)
|
||||
return ver
|
||||
except Exception as e:
|
||||
logger.warning("AI config bump_version 失败: %s", e)
|
||||
return ""
|
||||
|
||||
|
||||
def get_version() -> Optional[str]:
|
||||
"""读取当前版本号。Redis 不可用或异常返回 None。"""
|
||||
r = _get_redis_client()
|
||||
if r is None:
|
||||
return None
|
||||
try:
|
||||
return r.get(_REDIS_KEY)
|
||||
except Exception as e:
|
||||
logger.warning("AI config get_version 失败: %s", e)
|
||||
return None
|
||||
@@ -0,0 +1,516 @@
|
||||
"""AI 模型路由层 — 统一模型配置读取与客户端构建.
|
||||
|
||||
业务代码通过 AIRouter 获取客户端,不再硬编码 model/api_key/base_url。
|
||||
配置来源:DB ai_capability_configs JOIN ai_models → Redis 版本号缓存 → SharedSettings fallback。
|
||||
|
||||
使用方式:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_llm_client("intent_parsing")
|
||||
result = client.chat_completion(messages=[...])
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 配置数据类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelConfig:
|
||||
"""单个 AI 模型配置(来自 ai_models 表)"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
provider: str
|
||||
model_key: str
|
||||
api_key: str
|
||||
api_base: str
|
||||
api_version: str | None
|
||||
status: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CapabilityConfig:
|
||||
"""业务能力配置(来自 ai_capability_configs JOIN ai_models)"""
|
||||
|
||||
capability_key: str
|
||||
capability_name: str
|
||||
primary_model: ModelConfig | None
|
||||
lite_model: ModelConfig | None
|
||||
fallback_model: ModelConfig | None
|
||||
timeout_seconds: int
|
||||
max_retries: int
|
||||
max_tokens: int | None
|
||||
temperature: float | None
|
||||
concurrency: int
|
||||
extra_params: dict
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
# ── 简单包装类(TTS / ImageGen / VideoGen)──────────────────────────────────
|
||||
|
||||
|
||||
class TTSClient:
|
||||
"""TTS 客户端(简单配置持有者,实际调用由 CosyVoiceService 完成)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 60,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class ImageGenClient:
|
||||
"""图片生成客户端(简单配置持有者)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 60,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class VideoGenClient:
|
||||
"""视频生成客户端(简单配置持有者)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 600,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
# ── DB Session 获取 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_session():
|
||||
"""获取 DB session,兼容 api / worker / 独立脚本场景"""
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is not None:
|
||||
return SessionLocal()
|
||||
|
||||
try:
|
||||
from worker_app.db import SessionLocal as WorkerSL
|
||||
|
||||
if WorkerSL is not None:
|
||||
return WorkerSL()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from app.db import SessionLocal as ApiSL
|
||||
|
||||
if ApiSL is not None:
|
||||
return ApiSL()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# ── 核心路由类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AIRouter:
|
||||
"""AI 模型路由器 — 统一配置读取与客户端构建.
|
||||
|
||||
缓存策略:
|
||||
1. 本地内存缓存 {capability_key: CapabilityConfig}
|
||||
2. 每次读取前比对 Redis 版本号,变了则清缓存重新查 DB
|
||||
3. DB 无配置 / Redis 不可用 → fallback 到 SharedSettings 环境变量
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._cache: dict[str, CapabilityConfig] = {}
|
||||
self._local_ver: str | None = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def _check_version(self) -> bool:
|
||||
"""检查 Redis 版本号,变了返回 True(需要刷新缓存)"""
|
||||
from packages.shared.ai_config_version import get_version
|
||||
|
||||
current_ver = get_version()
|
||||
if current_ver is None:
|
||||
return False
|
||||
if self._local_ver != current_ver:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _load_from_db(self, capability_key: str) -> CapabilityConfig | None:
|
||||
"""从 DB 加载配置(ai_capability_configs JOIN ai_models)"""
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
logger.warning("AI Router: 无法获取 DB session")
|
||||
return None
|
||||
try:
|
||||
from sqlalchemy import text
|
||||
|
||||
sql = text("""
|
||||
SELECT
|
||||
cc.capability_key, cc.capability_name, cc.timeout_seconds,
|
||||
cc.max_retries, cc.max_tokens, cc.temperature,
|
||||
cc.concurrency, cc.extra_params, cc.is_enabled,
|
||||
pm.id AS pm_id, pm.name AS pm_name, pm.provider AS pm_provider,
|
||||
pm.model_key AS pm_model_key, pm.api_key AS pm_api_key,
|
||||
pm.api_base AS pm_api_base, pm.api_version AS pm_api_version,
|
||||
pm.status AS pm_status,
|
||||
lm.id AS lm_id, lm.name AS lm_name, lm.provider AS lm_provider,
|
||||
lm.model_key AS lm_model_key, lm.api_key AS lm_api_key,
|
||||
lm.api_base AS lm_api_base, lm.api_version AS lm_api_version,
|
||||
lm.status AS lm_status,
|
||||
fm.id AS fm_id, fm.name AS fm_name, fm.provider AS fm_provider,
|
||||
fm.model_key AS fm_model_key, fm.api_key AS fm_api_key,
|
||||
fm.api_base AS fm_api_base, fm.api_version AS fm_api_version,
|
||||
fm.status AS fm_status
|
||||
FROM ai_capability_configs cc
|
||||
LEFT JOIN ai_models pm ON cc.primary_model_id = pm.id AND pm.deleted_at IS NULL
|
||||
LEFT JOIN ai_models lm ON cc.lite_model_id = lm.id AND lm.deleted_at IS NULL
|
||||
LEFT JOIN ai_models fm ON cc.fallback_model_id = fm.id AND fm.deleted_at IS NULL
|
||||
WHERE cc.capability_key = :key AND cc.is_enabled = true
|
||||
""")
|
||||
row = session.execute(sql, {"key": capability_key}).first()
|
||||
if not row:
|
||||
return None
|
||||
|
||||
def _to_model(prefix: str) -> ModelConfig | None:
|
||||
mid = getattr(row, f"{prefix}_id", None)
|
||||
if not mid:
|
||||
return None
|
||||
return ModelConfig(
|
||||
id=mid,
|
||||
name=getattr(row, f"{prefix}_name", "") or "",
|
||||
provider=getattr(row, f"{prefix}_provider", "") or "",
|
||||
model_key=getattr(row, f"{prefix}_model_key", "") or "",
|
||||
api_key=getattr(row, f"{prefix}_api_key", "") or "",
|
||||
api_base=getattr(row, f"{prefix}_api_base", "") or "",
|
||||
api_version=getattr(row, f"{prefix}_api_version", None),
|
||||
status=getattr(row, f"{prefix}_status", "active") or "active",
|
||||
)
|
||||
|
||||
return CapabilityConfig(
|
||||
capability_key=row.capability_key,
|
||||
capability_name=row.capability_name,
|
||||
primary_model=_to_model("pm"),
|
||||
lite_model=_to_model("lm"),
|
||||
fallback_model=_to_model("fm"),
|
||||
timeout_seconds=row.timeout_seconds or 30,
|
||||
max_retries=row.max_retries or 1,
|
||||
max_tokens=row.max_tokens,
|
||||
temperature=row.temperature,
|
||||
concurrency=row.concurrency or 2,
|
||||
extra_params=row.extra_params or {},
|
||||
is_enabled=row.is_enabled,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("AI Router: DB 查询失败 (key=%s): %s", capability_key, e)
|
||||
return None
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
def get_capability(self, key: str) -> CapabilityConfig | None:
|
||||
"""获取业务能力配置(带缓存)"""
|
||||
with self._lock:
|
||||
if self._check_version():
|
||||
self._cache.clear()
|
||||
from packages.shared.ai_config_version import get_version
|
||||
|
||||
self._local_ver = get_version()
|
||||
|
||||
if key in self._cache:
|
||||
return self._cache[key]
|
||||
|
||||
config = self._load_from_db(key)
|
||||
if config:
|
||||
self._cache[key] = config
|
||||
return config
|
||||
|
||||
def _get_model_or_fallback(self, cap: CapabilityConfig, variant: str = "primary") -> ModelConfig | None:
|
||||
"""按 variant 选择模型,不存在则降级。
|
||||
|
||||
- primary: primary → fallback
|
||||
- lite: lite → primary
|
||||
- fallback: fallback → primary(修复点:此前 fallback variant 被忽略,错误地使用了 primary 模型)
|
||||
"""
|
||||
if variant == "fallback":
|
||||
if cap.fallback_model:
|
||||
return cap.fallback_model
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
elif variant == "lite":
|
||||
if cap.lite_model:
|
||||
return cap.lite_model
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
else: # primary
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
if cap.fallback_model:
|
||||
return cap.fallback_model
|
||||
return None
|
||||
|
||||
# ── 构建客户端 ─────────────────────────────────────────────────────────
|
||||
|
||||
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig):
|
||||
"""构建 LLM 客户端 — 返回 DoubaoClient 实例"""
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
max_retries=cap.max_retries,
|
||||
max_tokens=cap.max_tokens,
|
||||
temperature=cap.temperature,
|
||||
extra_params=cap.extra_params,
|
||||
provider=model.provider,
|
||||
)
|
||||
|
||||
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig):
|
||||
"""构建 VLM 客户端 — 返回 DoubaoClient 实例(DoubaoClient 已支持 vision_completion)"""
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
max_retries=cap.max_retries,
|
||||
max_tokens=cap.max_tokens,
|
||||
temperature=cap.temperature,
|
||||
extra_params=cap.extra_params,
|
||||
provider=model.provider,
|
||||
)
|
||||
|
||||
def _build_tts_client(self, model: ModelConfig, cap: CapabilityConfig) -> TTSClient:
|
||||
return TTSClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
def _build_image_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> ImageGenClient:
|
||||
return ImageGenClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
def _build_video_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> VideoGenClient:
|
||||
return VideoGenClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
# ── 公开接口 ────────────────────────────────────────────────────────────
|
||||
|
||||
def get_llm_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 LLM 客户端(返回 DoubaoClient 实例)"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
if model and model.api_key:
|
||||
return self._build_llm_client(model, cap)
|
||||
|
||||
return self._fallback_llm_client(key)
|
||||
|
||||
def get_vision_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 VLM 客户端(返回 DoubaoClient 实例)"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
if model and model.api_key:
|
||||
return self._build_vision_client(model, cap)
|
||||
|
||||
return self._fallback_vision_client(key)
|
||||
|
||||
def get_tts_client(self, key: str = "tts") -> TTSClient | None:
|
||||
"""获取 TTS 客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_tts_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_tts_client()
|
||||
|
||||
def get_image_gen_client(self, key: str = "image_generation") -> ImageGenClient | None:
|
||||
"""获取图片生成客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_image_gen_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_image_gen_client()
|
||||
|
||||
def get_video_gen_client(self, key: str = "video_generation") -> VideoGenClient | None:
|
||||
"""获取视频生成客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_video_gen_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_video_gen_client()
|
||||
|
||||
# ── Fallback 方法(读 SharedSettings 环境变量)──────────────────────────
|
||||
|
||||
def _fallback_llm_client(self, key: str):
|
||||
"""Fallback LLM 客户端 — 从 settings 读取配置,不硬编码"""
|
||||
settings = get_shared_settings()
|
||||
model_map = {
|
||||
"intent_parsing": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"copy_fusion": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"storyboard": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"copy_review": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"asset_classify": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
}
|
||||
if key in model_map:
|
||||
model_id, base_url, api_key = model_map[key]
|
||||
else:
|
||||
model_id = settings.doubao_model
|
||||
base_url = settings.doubao_base_url
|
||||
api_key = settings.doubao_api_key
|
||||
|
||||
if not api_key:
|
||||
return None
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model_id,
|
||||
timeout=settings.doubao_timeout,
|
||||
max_retries=settings.doubao_max_retries,
|
||||
)
|
||||
|
||||
def _fallback_vision_client(self, key: str):
|
||||
"""Fallback VLM 客户端 — 从 settings 读取 dashscope 配置,不硬编码"""
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "dashscope_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "dashscope_base_url", "") or ""
|
||||
model = getattr(settings, "dashscope_model", "") or getattr(settings, "doubao_vision_model", "")
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
provider="dashscope",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
def _fallback_tts_client(self) -> TTSClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "cosyvoice_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "cosyvoice_base_url", "")
|
||||
model = getattr(settings, "cosyvoice_model", "")
|
||||
|
||||
return TTSClient(provider="dashscope", api_key=api_key, base_url=base_url, model=model)
|
||||
|
||||
def _fallback_image_gen_client(self) -> ImageGenClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_image_model", "")
|
||||
|
||||
return ImageGenClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=getattr(settings, "doubao_image_timeout", 60),
|
||||
)
|
||||
|
||||
def _fallback_video_gen_client(self) -> VideoGenClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_video_model", "")
|
||||
|
||||
return VideoGenClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=getattr(settings, "doubao_video_timeout", 600),
|
||||
)
|
||||
|
||||
def invalidate(self):
|
||||
"""清空本地缓存"""
|
||||
with self._lock:
|
||||
self._cache.clear()
|
||||
self._local_ver = None
|
||||
|
||||
|
||||
# ── 全局单例 ──────────────────────────────────────────────────────────────
|
||||
|
||||
ai_router = AIRouter()
|
||||
@@ -0,0 +1,395 @@
|
||||
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# ── Pre-mock heavy import chain to avoid pulling in full app ──
|
||||
_mock_config = MagicMock()
|
||||
_mock_settings = MagicMock()
|
||||
_mock_settings.doubao_model = "doubao-seed-2-1-pro-260915"
|
||||
_mock_settings.doubao_fast_model = "doubao-seed-2-1-pro-260915"
|
||||
_mock_settings.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
_mock_settings.doubao_api_key = "test-key"
|
||||
_mock_settings.doubao_timeout = 45
|
||||
_mock_settings.doubao_max_retries = 1
|
||||
_mock_settings.doubao_image_model = "doubao-seedream-5-0-flash-260915"
|
||||
_mock_settings.doubao_image_timeout = 60
|
||||
_mock_settings.doubao_video_model = "doubao-seedance-2-5-260628"
|
||||
_mock_settings.doubao_video_timeout = 600
|
||||
_mock_settings.dashscope_api_key = "ds-key"
|
||||
_mock_settings.cosyvoice_api_key = "cv-key"
|
||||
_mock_settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||
_mock_settings.cosyvoice_model = "cosyvoice-v3-flash"
|
||||
_mock_settings.redis_url = "redis://localhost:6379/0"
|
||||
_mock_settings.celery_broker_url = "redis://localhost:6379/0"
|
||||
_mock_config.get_shared_settings.return_value = _mock_settings
|
||||
|
||||
# Prevent the full packages.shared from loading
|
||||
for mod_name in list(sys.modules.keys()):
|
||||
if "packages.shared" in mod_name and "ai_router" not in mod_name and "ai_config_version" not in mod_name:
|
||||
pass # don't remove, just prevent new imports
|
||||
|
||||
# Direct import of our modules (bypassing __init__.py)
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
|
||||
def _load_module_from_file(name, path):
|
||||
spec = importlib.util.spec_from_file_location(name, path)
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
sys.modules[name] = mod
|
||||
spec.loader.exec_module(mod)
|
||||
return mod
|
||||
|
||||
|
||||
# Load ai_config_version
|
||||
_ai_config_version = _load_module_from_file(
|
||||
"packages.shared.ai_config_version",
|
||||
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_config_version.py"),
|
||||
)
|
||||
# Patch get_shared_settings in the loaded module
|
||||
_ai_config_version.get_shared_settings = lambda: _mock_settings
|
||||
|
||||
# Load ai_router - needs packages.shared.config to be available
|
||||
sys.modules["packages.shared.config"] = MagicMock()
|
||||
sys.modules["packages.shared.config"].get_shared_settings = lambda: _mock_settings
|
||||
|
||||
# Mock packages.shared.ai_client to avoid triggering packages.shared.__init__ chain
|
||||
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
|
||||
_mock_ai_client = MagicMock()
|
||||
|
||||
class _FakeDoubaoClient:
|
||||
"""Fake DoubaoClient for testing - mimics the real interface."""
|
||||
def __init__(self, api_key="", base_url="", model="", timeout=0, max_retries=0,
|
||||
max_tokens=None, temperature=None, extra_params=None, provider="volcengine"):
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.max_retries = max_retries
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.extra_params = extra_params or {}
|
||||
self.provider = provider
|
||||
self.vision_model = model
|
||||
|
||||
@property
|
||||
def is_available(self):
|
||||
return bool(self.api_key)
|
||||
|
||||
def chat_completion(self, messages, **kwargs):
|
||||
return None
|
||||
|
||||
def vision_completion(self, messages, **kwargs):
|
||||
return None
|
||||
|
||||
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
|
||||
sys.modules["packages.shared.ai_client"] = _mock_ai_client
|
||||
|
||||
_ai_router = _load_module_from_file(
|
||||
"packages.shared.ai_router",
|
||||
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_router.py"),
|
||||
)
|
||||
|
||||
|
||||
class TestAIConfigVersion(unittest.TestCase):
|
||||
"""Redis 版本号机制测试"""
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
def test_bump_version_success(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.set.return_value = True
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.bump_version()
|
||||
self.assertTrue(ver)
|
||||
self.assertTrue(ver.isdigit())
|
||||
mock_r.set.assert_called_once()
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
def test_bump_version_redis_unavailable(self, mock_redis_fn):
|
||||
mock_redis_fn.return_value = None
|
||||
ver = _ai_config_version.bump_version()
|
||||
self.assertEqual(ver, "")
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
def test_get_version_success(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.get.return_value = "1234567890"
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.get_version()
|
||||
self.assertEqual(ver, "1234567890")
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
def test_get_version_redis_down(self, mock_redis_fn):
|
||||
mock_redis_fn.return_value = None
|
||||
ver = _ai_config_version.get_version()
|
||||
self.assertIsNone(ver)
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
def test_get_version_exception(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.get.side_effect = Exception("connection refused")
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.get_version()
|
||||
self.assertIsNone(ver)
|
||||
|
||||
|
||||
class TestAIRouter(unittest.TestCase):
|
||||
"""AIRouter 路由/缓存/fallback 测试"""
|
||||
|
||||
def setUp(self):
|
||||
self.router = _ai_router.AIRouter()
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_capability_db_unavailable(self, mock_ver):
|
||||
with patch.object(_ai_router, "_get_session", return_value=None):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNone(cap)
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_capability_from_db(self, mock_ver):
|
||||
mock_session = MagicMock()
|
||||
mock_row = MagicMock()
|
||||
mock_row.capability_key = "intent_parsing"
|
||||
mock_row.capability_name = "文案意图解析"
|
||||
mock_row.timeout_seconds = 45
|
||||
mock_row.max_retries = 1
|
||||
mock_row.max_tokens = None
|
||||
mock_row.temperature = None
|
||||
mock_row.concurrency = 2
|
||||
mock_row.extra_params = {}
|
||||
mock_row.is_enabled = True
|
||||
mock_row.pm_id = "model-1"
|
||||
mock_row.pm_name = "豆包"
|
||||
mock_row.pm_provider = "volcengine"
|
||||
mock_row.pm_model_key = "doubao-seed-1-6-250615"
|
||||
mock_row.pm_api_key = "test-key"
|
||||
mock_row.pm_api_base = "https://ark.test.com"
|
||||
mock_row.pm_api_version = None
|
||||
mock_row.pm_status = "active"
|
||||
mock_row.lm_id = None
|
||||
mock_row.fm_id = None
|
||||
mock_session.execute.return_value.first.return_value = mock_row
|
||||
|
||||
with patch.object(_ai_router, "_get_session", return_value=mock_session):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNotNone(cap)
|
||||
self.assertEqual(cap.capability_key, "intent_parsing")
|
||||
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", side_effect=[None, "v2"])
|
||||
def test_cache_invalidation_on_version_change(self, mock_ver):
|
||||
with patch.object(self.router, "_load_from_db", return_value=None):
|
||||
self.router.get_capability("test_key")
|
||||
self.router._local_ver = "v1"
|
||||
self.assertTrue(self.router._check_version())
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value="same_ver")
|
||||
def test_cache_hit_same_version(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="test-model",
|
||||
api_key="key", api_base="https://test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test", primary_model=model,
|
||||
lite_model=None, fallback_model=None, timeout_seconds=30,
|
||||
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
|
||||
extra_params={}, is_enabled=True,
|
||||
)
|
||||
self.router._cache["test"] = cap
|
||||
self.router._local_ver = "same_ver"
|
||||
result = self.router.get_capability("test")
|
||||
self.assertEqual(result, cap)
|
||||
|
||||
def test_invalidate_clears_cache(self):
|
||||
self.router._cache["x"] = MagicMock()
|
||||
self.router._local_ver = "v1"
|
||||
self.router.invalidate()
|
||||
self.assertEqual(len(self.router._cache), 0)
|
||||
self.assertIsNone(self.router._local_ver)
|
||||
|
||||
@patch.object(_ai_router, "_get_session", return_value=None)
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_llm_client_fallback(self, mock_ver, mock_session):
|
||||
_ai_router.get_shared_settings = lambda: _mock_settings
|
||||
client = self.router.get_llm_client("intent_parsing")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||
self.assertEqual(client.api_key, "test-key")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_llm_client_from_db(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
||||
api_key="db-key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=350, temperature=0.1,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_llm_client("image_analysis")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "qwen3.8-flash")
|
||||
self.assertEqual(client.provider, "dashscope")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_vision_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_vision_client("image_analysis")
|
||||
self.assertIsNotNone(client)
|
||||
# #2220: vision client is now DoubaoClient with vision_completion
|
||||
self.assertTrue(hasattr(client, "vision_completion"))
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_tts_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="cosyvoice-v3-flash",
|
||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="tts", capability_name="语音合成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_tts_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "cosyvoice-v3-flash")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_image_gen_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="seedream-5.0-flash",
|
||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_generation", capability_name="图片生成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={"size": "1K"}, is_enabled=True,
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_image_gen_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedream-5.0-flash")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_video_gen_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="seedance-2.5",
|
||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="video_generation", capability_name="视频生成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=600, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=1, extra_params={}, is_enabled=True,
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_video_gen_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedance-2.5")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_lite_variant_preference(self, mock_ver):
|
||||
primary = _ai_router.ModelConfig(id="p1", name="pro", provider="volcengine", model_key="pro-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
lite = _ai_router.ModelConfig(id="l1", name="lite", provider="volcengine", model_key="lite-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=primary, lite_model=lite, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
)
|
||||
model = self.router._get_model_or_fallback(cap, "lite")
|
||||
self.assertEqual(model.model_key, "lite-model")
|
||||
model_primary = self.router._get_model_or_fallback(cap, "primary")
|
||||
self.assertEqual(model_primary.model_key, "pro-model")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_disabled_capability_returns_fallback(self, mock_ver):
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test",
|
||||
primary_model=None, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=False,
|
||||
)
|
||||
_ai_router.get_shared_settings = lambda: _mock_settings
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_llm_client("test")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_fallback_chain_primary_none(self, mock_ver):
|
||||
"""primary_model 为 None 时 fallback 到 fallback_model"""
|
||||
fb = _ai_router.ModelConfig(id="f1", name="fb", provider="volcengine", model_key="fb-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test",
|
||||
primary_model=None, lite_model=None, fallback_model=fb,
|
||||
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
)
|
||||
model = self.router._get_model_or_fallback(cap, "primary")
|
||||
self.assertEqual(model.model_key, "fb-model")
|
||||
|
||||
|
||||
class TestModelConfig(unittest.TestCase):
|
||||
"""数据类测试"""
|
||||
|
||||
def test_model_config_frozen(self):
|
||||
m = _ai_router.ModelConfig(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b", api_version=None, status="active")
|
||||
with self.assertRaises(AttributeError):
|
||||
m.model_key = "new"
|
||||
|
||||
def test_capability_config_frozen(self):
|
||||
c = _ai_router.CapabilityConfig(
|
||||
capability_key="k", capability_name="n", primary_model=None,
|
||||
lite_model=None, fallback_model=None, timeout_seconds=30,
|
||||
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
|
||||
extra_params={}, is_enabled=True,
|
||||
)
|
||||
with self.assertRaises(AttributeError):
|
||||
c.is_enabled = False
|
||||
|
||||
|
||||
class TestClientAvailability(unittest.TestCase):
|
||||
"""客户端可用性测试"""
|
||||
|
||||
def test_tts_client_available(self):
|
||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
def test_tts_client_unavailable_no_model(self):
|
||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="")
|
||||
self.assertFalse(c.is_available)
|
||||
|
||||
def test_image_gen_client_unavailable_no_url(self):
|
||||
c = _ai_router.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
|
||||
self.assertFalse(c.is_available)
|
||||
|
||||
def test_video_gen_client_available(self):
|
||||
c = _ai_router.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -73,15 +73,15 @@ class TestSharedSettingsDefaults:
|
||||
|
||||
def test_default_cosyvoice_settings(self):
|
||||
s = SharedSettings()
|
||||
assert s.cosyvoice_model == "cosyvoice-v3-flash"
|
||||
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
|
||||
assert s.cosyvoice_format == "mp3"
|
||||
assert s.cosyvoice_sample_rate == 22050
|
||||
|
||||
def test_default_doubao_settings(self):
|
||||
s = SharedSettings()
|
||||
assert "doubao" in s.doubao_model
|
||||
assert s.doubao_model == "" # 零硬编码:默认值已清空
|
||||
assert s.doubao_timeout == 45 # #2180 默认提到45s
|
||||
assert s.doubao_max_retries == 1
|
||||
assert s.doubao_max_retries == 3
|
||||
|
||||
|
||||
class TestAPISettingsDefaults:
|
||||
@@ -321,7 +321,7 @@ class TestWorkerSettingsDefaults:
|
||||
assert s.database_url # 继承自SharedSettings
|
||||
assert s.redis_url
|
||||
assert s.oss_endpoint
|
||||
assert s.cosyvoice_model == "cosyvoice-v3-flash"
|
||||
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
|
||||
|
||||
|
||||
class TestGetWorkerSettings:
|
||||
|
||||
@@ -102,17 +102,17 @@ class TestSharedSettingsDefaults:
|
||||
def test_default_cosyvoice_config(self):
|
||||
"""CosyVoice 默认配置"""
|
||||
s = self._make_settings()
|
||||
assert s.cosyvoice_model == "cosyvoice-v3-flash"
|
||||
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
|
||||
assert s.cosyvoice_sample_rate == 22050
|
||||
assert s.cosyvoice_format == "mp3"
|
||||
assert s.cosyvoice_clone_model == "voice-enrollment"
|
||||
assert s.cosyvoice_clone_model == "" # 零硬编码:默认值已清空
|
||||
|
||||
def test_default_doubao_config(self):
|
||||
"""豆包默认配置"""
|
||||
s = self._make_settings()
|
||||
assert s.doubao_timeout == 45 # #2180 默认提到45s
|
||||
assert s.doubao_max_retries == 1
|
||||
assert "volces.com" in s.doubao_base_url
|
||||
assert s.doubao_max_retries == 3
|
||||
assert s.doubao_base_url == "" # 零硬编码:默认值已清空
|
||||
|
||||
def test_default_empty_api_keys(self):
|
||||
"""API Key 默认空字符串"""
|
||||
|
||||
@@ -27,7 +27,10 @@ def mock_client() -> MagicMock:
|
||||
@pytest.fixture
|
||||
def service(mock_client: MagicMock) -> CosyVoiceService:
|
||||
"""Create CosyVoiceService with mocked HTTP client and config."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
settings.cosyvoice_api_key = "sk-test-12345678"
|
||||
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||
@@ -37,6 +40,7 @@ def service(mock_client: MagicMock) -> CosyVoiceService:
|
||||
settings.cosyvoice_voice = "longxiaochun_v3"
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None # ai_router returns None in tests
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
svc.CLONE_POLL_INTERVAL = 0.001 # 加速测试
|
||||
svc.RETRY_BACKOFF = 0.001
|
||||
@@ -48,7 +52,10 @@ class TestInitConfig:
|
||||
|
||||
def test_base_url_with_old_text2audio_path_gets_normalized(self, mock_client: MagicMock) -> None:
|
||||
"""旧版 base_url 带 text2audio 路径应自动修正."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
settings.cosyvoice_api_key = "sk-test"
|
||||
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio"
|
||||
@@ -58,12 +65,16 @@ class TestInitConfig:
|
||||
settings.cosyvoice_voice = "test"
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
assert svc._base_url == "https://dashscope.aliyuncs.com/api/v1"
|
||||
|
||||
def test_custom_params_override_config(self, mock_client: MagicMock) -> None:
|
||||
"""显式传入参数覆盖配置."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
settings.cosyvoice_api_key = "sk-config"
|
||||
settings.cosyvoice_base_url = "https://config.example.com"
|
||||
@@ -73,6 +84,7 @@ class TestInitConfig:
|
||||
settings.cosyvoice_voice = "test"
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(
|
||||
api_key="sk-custom",
|
||||
base_url="https://custom.example.com/api/v1",
|
||||
@@ -87,7 +99,10 @@ class TestInitConfig:
|
||||
|
||||
def test_context_manager(self, mock_client: MagicMock) -> None:
|
||||
"""上下文管理器正常工作."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = "sk-test"
|
||||
@@ -105,6 +120,7 @@ class TestInitConfig:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
with svc as s:
|
||||
assert s is svc
|
||||
@@ -113,7 +129,10 @@ class TestInitConfig:
|
||||
|
||||
def test_owns_client_gets_closed(self) -> None:
|
||||
"""自有client在close时被关闭."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = "sk-test"
|
||||
@@ -131,6 +150,7 @@ class TestInitConfig:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
with patch("packages.application.cosyvoice_service.httpx.Client") as mock_cls:
|
||||
mock_instance = MagicMock()
|
||||
mock_cls.return_value = mock_instance
|
||||
@@ -173,7 +193,10 @@ class TestSubmitCloneTask:
|
||||
|
||||
def test_no_api_key_raises_auth_error(self, mock_client: MagicMock) -> None:
|
||||
"""无API Key抛认证错误."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = ""
|
||||
@@ -191,6 +214,7 @@ class TestSubmitCloneTask:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
with pytest.raises(CosyVoiceAuthError, match="API Key 未配置"):
|
||||
svc.submit_clone_task(audio_url="https://example.com/audio.mp3")
|
||||
@@ -258,7 +282,10 @@ class TestSubmitCloneTask:
|
||||
|
||||
def test_audio_url_signer_is_called(self, mock_client: MagicMock) -> None:
|
||||
"""配置了audio_url_signer时会被调用预签名."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = "sk-test"
|
||||
@@ -276,6 +303,7 @@ class TestSubmitCloneTask:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
signer = MagicMock(return_value="https://signed.example.com/audio.mp3?token=xxx")
|
||||
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
|
||||
|
||||
@@ -297,7 +325,10 @@ class TestSubmitCloneTask:
|
||||
|
||||
def test_signer_failure_falls_back_to_original_url(self, mock_client: MagicMock) -> None:
|
||||
"""预签名失败时回退到原始URL."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = "sk-test"
|
||||
@@ -315,6 +346,7 @@ class TestSubmitCloneTask:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
signer = MagicMock(side_effect=RuntimeError("sign failed"))
|
||||
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
|
||||
|
||||
@@ -375,7 +407,10 @@ class TestQueryVoiceStatus:
|
||||
|
||||
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
|
||||
"""无API Key抛认证错误."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = ""
|
||||
@@ -393,6 +428,7 @@ class TestQueryVoiceStatus:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
with pytest.raises(CosyVoiceAuthError):
|
||||
svc.query_voice_status("v1")
|
||||
@@ -589,7 +625,10 @@ class TestSubmitSynthesizeTask:
|
||||
|
||||
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
|
||||
"""无API Key抛认证错误."""
|
||||
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
|
||||
with (
|
||||
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
|
||||
patch("packages.shared.ai_router.ai_router") as mock_router,
|
||||
):
|
||||
settings = MagicMock()
|
||||
|
||||
settings.cosyvoice_api_key = ""
|
||||
@@ -607,6 +646,7 @@ class TestSubmitSynthesizeTask:
|
||||
settings.cosyvoice_clone_model = "voice-enrollment"
|
||||
|
||||
mock_settings.return_value = settings
|
||||
mock_router.get_tts_client.return_value = None
|
||||
svc = CosyVoiceService(http_client=mock_client)
|
||||
with pytest.raises(CosyVoiceAuthError):
|
||||
svc.submit_synthesize_task(text="你好", voice_id="v1")
|
||||
|
||||
@@ -521,3 +521,81 @@ class TestIngestJob:
|
||||
storage_key="k",
|
||||
)
|
||||
assert job.error_message == ""
|
||||
|
||||
|
||||
class TestViralVideoResumeForRegenerate:
|
||||
"""#2222: resume_from_image_analyzed 应支持 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。"""
|
||||
|
||||
def test_regen_from_copy_generated_clears_old_copy(self):
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
job = ViralVideoJob(user_id="u1", images=["img1"])
|
||||
# 模拟已经生成过文案和视频
|
||||
job.status = ViralVideoStatus.COPY_GENERATED
|
||||
job.copy_result = {"shots": [{"x": 1}], "voiceover_script": "旧文案"}
|
||||
job.intent_result = {"intent": "旧意图"}
|
||||
job.storyboard = [{"x": 1}]
|
||||
job.generated_copy_text = "旧文案"
|
||||
job.result_video_url = "http://old.mp4"
|
||||
job.completed_at = datetime(2026, 10, 6, tzinfo=timezone.utc)
|
||||
job.error_msg = ""
|
||||
job.current_stage = "tts_generation"
|
||||
job.phase_message = "TTS完成"
|
||||
|
||||
# 重新生成
|
||||
job.resume_from_image_analyzed()
|
||||
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
assert job.copy_result is None
|
||||
assert job.intent_result is None
|
||||
assert job.storyboard is None
|
||||
assert job.generated_copy_text == ""
|
||||
assert job.result_video_url == ""
|
||||
assert job.completed_at is None
|
||||
assert job.error_msg == ""
|
||||
assert job.current_stage == ""
|
||||
assert job.phase_message == ""
|
||||
|
||||
def test_regen_from_completed_clears_old_copy(self):
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
job = ViralVideoJob(user_id="u1", images=["img1"])
|
||||
job.status = ViralVideoStatus.COMPLETED
|
||||
job.copy_result = {"shots": [], "voiceover_script": "xx"}
|
||||
job.intent_result = {"intent": "x"}
|
||||
job.result_video_url = "http://v.mp4"
|
||||
|
||||
job.resume_from_image_analyzed()
|
||||
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
assert job.copy_result is None
|
||||
assert job.intent_result is None
|
||||
assert job.result_video_url == ""
|
||||
|
||||
def test_first_call_from_image_analyzed_keeps_fields(self):
|
||||
"""首次进入(IMAGE_ANALYZED)不应清空任何已有的字段。"""
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
job = ViralVideoJob(user_id="u1", images=["img1"])
|
||||
job.status = ViralVideoStatus.IMAGE_ANALYZED
|
||||
job.image_analysis = {"products": []}
|
||||
job.industry = "美妆"
|
||||
|
||||
job.resume_from_image_analyzed()
|
||||
|
||||
assert job.status == ViralVideoStatus.RUNNING
|
||||
assert job.image_analysis == {"products": []}
|
||||
assert job.industry == "美妆"
|
||||
|
||||
def test_wait_user_confirm_rejected(self):
|
||||
"""wait_user_confirm 中间状态应被拒绝(前端正在编辑/确认文案)。"""
|
||||
import pytest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
job = ViralVideoJob(user_id="u1", images=["img1"])
|
||||
job.status = ViralVideoStatus.WAIT_USER_CONFIRM
|
||||
with pytest.raises(ValueError, match="Cannot resume"):
|
||||
job.resume_from_image_analyzed()
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -108,8 +108,14 @@ class TestGetTtsService:
|
||||
|
||||
def test_empty_env_falls_back_to_auto_detect(self):
|
||||
"""环境变量为空时自动检测."""
|
||||
with patch.dict(os.environ, {"TTS_PROVIDER": ""}):
|
||||
with (
|
||||
patch.dict(os.environ, {"TTS_PROVIDER": ""}),
|
||||
patch("packages.shared.config.get_shared_settings") as mock_settings,
|
||||
):
|
||||
# 没有 cosyvoice_api_key 时应该用 mock
|
||||
settings = MagicMock()
|
||||
settings.cosyvoice_api_key = ""
|
||||
mock_settings.return_value = settings
|
||||
service = get_tts_service(None)
|
||||
assert service.provider_name == "mock"
|
||||
|
||||
|
||||
@@ -431,7 +431,7 @@ class TestGenerateCopy:
|
||||
assert resp.id == "job-gc"
|
||||
|
||||
def test_generate_copy_rejects_wrong_status(self):
|
||||
"""任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。"""
|
||||
"""wait_user_confirm 等中间状态不允许调用 generate-copy(状态保护)。"""
|
||||
import pytest
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import GenerateCopyRequest
|
||||
@@ -441,7 +441,8 @@ class TestGenerateCopy:
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
|
||||
# wait_user_confirm 属于前端在编辑/确认文案的中间状态,应拒绝重新触发生成
|
||||
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
@@ -450,6 +451,31 @@ class TestGenerateCopy:
|
||||
vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session)
|
||||
assert exc.value.status_code == 409
|
||||
|
||||
def test_generate_copy_allows_regenerate_from_copy_generated(self):
|
||||
"""#2222: COPY_GENERATED/COMPLETED 状态下点「重新生成文案」应放行入队,不返回 409。"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import GenerateCopyRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
for regen_status in (ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
|
||||
job = _make_job(job_id=f"job-regen-{regen_status}", user_id="u1", status=regen_status)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.generate_copy(f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session)
|
||||
mock_send.assert_called_once()
|
||||
job.resume_from_image_analyzed.assert_called()
|
||||
assert job.retry_count >= 1
|
||||
assert resp.id == f"job-regen-{regen_status}"
|
||||
|
||||
def test_generate_copy_persists_voice_and_ratio(self):
|
||||
"""generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。"""
|
||||
from unittest.mock import patch
|
||||
|
||||
Executable
+291
@@ -0,0 +1,291 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""vision v4 prompt / assembler 单元测试:
|
||||
|
||||
- assembler 正确识别 v4 嵌套 schema 与旧扁平 schema
|
||||
- v4 product/person/store/other 四类输出组装出下游必出字段
|
||||
- 旧扁平 schema 行为不变
|
||||
- _prompt._resolve:DB 有 active prompt 时原样使用(不追加硬编码 schema);
|
||||
DB 无记录时回落到硬编码 JSON schema
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from worker_app.tasks.vision import _prompt, assembler
|
||||
|
||||
REQUIRED_KEYS = {
|
||||
"name",
|
||||
"brand",
|
||||
"category",
|
||||
"appearance",
|
||||
"packaging",
|
||||
"text_on_package",
|
||||
"key_features",
|
||||
"scene",
|
||||
"mood",
|
||||
"portrait_prompt",
|
||||
"summary",
|
||||
"_source",
|
||||
}
|
||||
|
||||
|
||||
# ---------- schema 识别 ----------
|
||||
|
||||
|
||||
def test_is_v4_schema_products_list() -> None:
|
||||
assert assembler._is_v4_schema({"type": "product", "products": []})
|
||||
|
||||
|
||||
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 = {
|
||||
"type": "product",
|
||||
"products": [
|
||||
{"product_name": "次要商品", "brand": "B"},
|
||||
{"product_name": "主商品", "brand": "A", "position": "main"},
|
||||
],
|
||||
}
|
||||
r = assembler.assemble_result(1, fj, [])
|
||||
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, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert r["name"] == "社区便利店"
|
||||
assert r["brand"] == "全家FamilyMart"
|
||||
assert r["category"] == "门店场景"
|
||||
assert any("饮料" in str(f) for f in r["key_features"])
|
||||
assert "门店实拍" in r["portrait_prompt"]
|
||||
|
||||
|
||||
# ---------- v4 other ----------
|
||||
|
||||
|
||||
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 = {
|
||||
"has_person": True,
|
||||
"gender": "男",
|
||||
"age_range": "中年",
|
||||
"upper_wear": "西装",
|
||||
"upper_color": "深灰色",
|
||||
"lower_wear": "西裤",
|
||||
"lower_color": "黑色",
|
||||
"accessories": ["手表"],
|
||||
"hairstyle": "短发",
|
||||
"expression": "严肃",
|
||||
"scene": "办公室",
|
||||
"style": "商务",
|
||||
"mood": "专业",
|
||||
}
|
||||
r = assembler.assemble_result(0, fj, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert "中年男性" in r["portrait_prompt"]
|
||||
assert r["_source"] == "v2_fast_json"
|
||||
|
||||
|
||||
def test_assemble_old_flat_product() -> None:
|
||||
fj = {
|
||||
"has_person": False,
|
||||
"product_name": "口红",
|
||||
"brand": "Dior",
|
||||
"category": "美妆",
|
||||
"colors": ["红色"],
|
||||
"scene": "通用",
|
||||
"style": "商业",
|
||||
"mood": "高级",
|
||||
}
|
||||
r = assembler.assemble_result(0, fj, ["Dior"])
|
||||
assert r["name"] == "口红"
|
||||
assert r["brand"] == "Dior"
|
||||
assert r["text_on_package"] == ["Dior"]
|
||||
|
||||
|
||||
def test_assemble_none_input() -> None:
|
||||
r = assembler.assemble_result(0, None, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
|
||||
|
||||
# ---------- _prompt 解析 ----------
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_prompt_cache() -> Any:
|
||||
_prompt.invalidate_cache()
|
||||
yield
|
||||
_prompt.invalidate_cache()
|
||||
|
||||
|
||||
def _fake_tpl(system_prompt: str = "v4 system prompt 只返回JSON") -> Any:
|
||||
return types.SimpleNamespace(
|
||||
system_prompt=system_prompt,
|
||||
user_prompt_template="分析 {image_count} 张图",
|
||||
version=4,
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PROMPT_XYZ"))
|
||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
assert sys_prompt == "DB_V4_PROMPT_XYZ"
|
||||
assert "DB_V4_PROMPT_XYZ" not in _prompt._FAST_JSON_APPEND # sanity: 旧append是另一段文本
|
||||
assert "分析 1 张图" in user_prompt
|
||||
|
||||
|
||||
def test_resolve_pro_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PRO_PROMPT"))
|
||||
sys_prompt, _ = _prompt.resolve_pro_prompt()
|
||||
assert sys_prompt == "DB_V4_PRO_PROMPT"
|
||||
assert "【输出格式要求】" not in sys_prompt
|
||||
|
||||
|
||||
def test_resolve_falls_back_when_no_db(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
|
||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
assert sys_prompt == _prompt._FAST_JSON_SCHEMA
|
||||
assert user_prompt == _prompt.DEFAULT_FAST_USER
|
||||
|
||||
|
||||
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
calls = {"n": 0}
|
||||
|
||||
def _load() -> Any:
|
||||
calls["n"] += 1
|
||||
return _fake_tpl("CACHED_PROMPT")
|
||||
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", _load)
|
||||
s1, _ = _prompt.resolve_fast_prompt()
|
||||
s2, _ = _prompt.resolve_fast_prompt()
|
||||
assert s1 == s2 == "CACHED_PROMPT"
|
||||
assert calls["n"] == 1
|
||||
Reference in New Issue
Block a user