Compare commits
27 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 |
@@ -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
|
||||
@@ -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;
|
||||
|
||||
@@ -443,11 +443,11 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
from packages.shared.ai_router import ai_router
|
||||
except ImportError:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
_llm_client = get_doubao_client()
|
||||
_llm_client = ai_router.get_llm_client("intent_parsing")
|
||||
if not _llm_client.is_available:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
@@ -499,19 +499,22 @@ 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")
|
||||
for _client, _lbl in [(_client_fast, "fast"), (_client_pro, "pro-fallback")]:
|
||||
if not _client or not _client.is_available:
|
||||
continue
|
||||
try:
|
||||
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
|
||||
raw = _llm_client.chat_completion(
|
||||
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/#2215: 直接用 client.chat_completion 传 messages list,不再走 call_llm 字符串包装
|
||||
)
|
||||
if not raw:
|
||||
continue
|
||||
parsed = _parse(raw)
|
||||
@@ -851,11 +854,11 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
GLOBAL_CONSTRAINTS,
|
||||
NEGATIVE_RULES,
|
||||
)
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
from packages.shared.ai_router import ai_router
|
||||
except ImportError:
|
||||
return _fallback_script(job)
|
||||
|
||||
_llm_client2 = get_doubao_client()
|
||||
_llm_client2 = ai_router.get_llm_client("storyboard")
|
||||
if not _llm_client2.is_available:
|
||||
return _fallback_script(job)
|
||||
|
||||
@@ -892,13 +895,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 = _llm_client2.chat_completion(
|
||||
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:
|
||||
@@ -929,22 +933,22 @@ 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", "150"))
|
||||
_script_pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "150"))
|
||||
try:
|
||||
# #2217: doubao-seed-2-1-pro生成长编导脚本高峰期>90s,上调到150s,支持ENV覆盖
|
||||
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
normalized = _try_gen(_client_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
normalized = _try_gen(_client_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
# 第三次:用主力模型兜底
|
||||
if _pro and _pro != _fast:
|
||||
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
|
||||
# 第三次:用 lite/pro 模型兜底
|
||||
if _client_pro and _client_pro.is_available:
|
||||
normalized = _try_gen(_client_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
|
||||
@@ -1702,7 +1706,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, "意图解析完成")
|
||||
|
||||
# 阶段:编导脚本生成(核心耗时环节,已用快模型)
|
||||
@@ -1760,7 +1764,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:
|
||||
@@ -1919,16 +1924,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, "正在进行出片前合规审核...")
|
||||
@@ -1941,12 +1956,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, "合规审核完成")
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
# -*- 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 读取
|
||||
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
|
||||
- timeout=30s
|
||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
||||
"""
|
||||
|
||||
@@ -15,7 +15,6 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
@@ -23,14 +22,7 @@ from . import _prompt, assembler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
_PRO_MODEL = "qwen3.7-plus"
|
||||
_DEFAULT_TIMEOUT = 30
|
||||
_DEFAULT_MAX_TOKENS = 800
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
@@ -38,72 +30,58 @@ 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="primary")
|
||||
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,
|
||||
)
|
||||
call_kwargs: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"images": None, # 图片已在 messages 中
|
||||
"temperature": 0.3,
|
||||
"timeout": timeout,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
if max_tokens is not None:
|
||||
call_kwargs["max_tokens"] = max_tokens
|
||||
raw = client.vision_completion(**call_kwargs)
|
||||
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,
|
||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs",
|
||||
client.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()
|
||||
s = _strip_code_fence(raw)
|
||||
lpos, rr = s.find("{"), s.rfind("}")
|
||||
if lpos >= 0 and rr > lpos:
|
||||
s = s[lpos : rr + 1]
|
||||
@@ -124,3 +102,15 @@ def call_pro_vlm(
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
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
|
||||
|
||||
@@ -1,22 +1,21 @@
|
||||
# -*- 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,15 +23,7 @@ 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 = 15
|
||||
_DEFAULT_MAX_TOKENS = 350
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
@@ -51,78 +42,60 @@ 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,
|
||||
)
|
||||
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
|
||||
raw = client.vision_completion(**call_kwargs)
|
||||
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:
|
||||
logger.warning(
|
||||
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
|
||||
)
|
||||
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)
|
||||
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_FAST_MODEL,
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs",
|
||||
client.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("}")
|
||||
|
||||
@@ -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 规范化:去掉末尾的路径残留(兼容旧版配置)
|
||||
|
||||
@@ -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
|
||||
|
||||
# ── 审核 ────────────────────────────────────────────────────────────
|
||||
|
||||
+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 = 1
|
||||
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,11 +191,40 @@ 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
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
|
||||
@@ -170,13 +170,33 @@ 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.vision_lite_model: str = settings.doubao_vision_lite_model
|
||||
self.fast_model: str = settings.doubao_fast_model
|
||||
@@ -240,16 +260,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,12 +283,18 @@ class DoubaoClient:
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": 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()
|
||||
@@ -282,6 +309,22 @@ 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 截断:1.5x 扩容后重试(计入 max_retries,不额外增加)
|
||||
old_max = int(payload["max_tokens"])
|
||||
new_max = int(old_max * 1.5)
|
||||
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"]
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
@@ -315,20 +358,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。
|
||||
|
||||
@@ -368,12 +412,17 @@ 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
|
||||
last_error: Optional[Exception] = None
|
||||
@@ -388,6 +437,22 @@ 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 截断:1.5x 扩容后重试(计入 max_retries)
|
||||
old_max = int(payload["max_tokens"])
|
||||
new_max = int(old_max * 1.5)
|
||||
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"]
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
|
||||
@@ -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,478 @@
|
||||
"""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 选择模型,不存在则 fallback"""
|
||||
if variant == "lite" and cap.lite_model:
|
||||
return cap.lite_model
|
||||
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,13 +73,13 @@ 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
|
||||
|
||||
@@ -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_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
|
||||
|
||||
Reference in New Issue
Block a user