feat: AI模型路由层 — 统一模型配置读取与客户端构建
核心改动: - packages/shared/ai_config_version.py: Redis版本号通知机制 - packages/shared/ai_router.py: AIRouter统一路由层(DB→Redis缓存→SharedSettings fallback) - alembic/versions/099: 补齐缺失模型seed和capability配置 - apps/worker/worker_app/tasks/vision/*: VLM硬编码替换为ai_router动态配置 - packages/application/viral_video/reviewer.py: 审核走copy_review capability - packages/application/cosyvoice_service.py: TTS走tts capability - apps/worker/worker_app/tasks/viral_video.py: 意图解析/分镜/文案走DB配置 - packages/config/base.py: 移除doubao/dashscope/mediakit/cosyvoice硬编码默认值 - tests/unit/test_ai_router.py: 26个单元测试覆盖路由/缓存/fallback Admin侧: - app/utils/ai_config_notify.py: bump_ai_config_version()工具函数 - routers/ai_models.py + ai_capability_configs.py: CRUD后调用bump_version() 验收标准: - admin改模型后30s内生效(Redis版本号通知) - Redis故障fallback环境变量 - base.py零硬编码(api_key保留空串) - 26个单测全绿
This commit is contained in:
@@ -0,0 +1,221 @@
|
||||
# -*- 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"
|
||||
)
|
||||
)
|
||||
@@ -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": ""}
|
||||
|
||||
@@ -500,8 +500,14 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
}
|
||||
|
||||
_s = get_shared_settings()
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = _s.doubao_model
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
_cap = ai_router.get_capability('intent_parsing')
|
||||
_fast = (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or _s.doubao_fast_model
|
||||
_pro = (_cap.lite_model.model_key if _cap and _cap.lite_model else None) or (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or _s.doubao_model
|
||||
except Exception:
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = _s.doubao_model
|
||||
for _m, _lbl in [(_fast, "fast"), (_pro, "pro-fallback")]:
|
||||
try:
|
||||
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
|
||||
@@ -851,11 +857,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)
|
||||
|
||||
@@ -930,8 +936,14 @@ 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
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
_cap = ai_router.get_capability('storyboard')
|
||||
_fast = (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or _s.doubao_fast_model
|
||||
_pro = (_cap.lite_model.model_key if _cap and _cap.lite_model else None) or (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or getattr(_s, "doubao_model", None) or _fast
|
||||
except Exception:
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = getattr(_s, "doubao_model", None) or _fast
|
||||
_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:
|
||||
|
||||
@@ -1,13 +1,12 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 兜底路径:qwen3.7-plus(阿里云百炼/DashScope)单图调用。
|
||||
"""V2 兜底路径:image_analysis lite/fallback(默认 qwen3.7-plus / DashScope)单图调用。
|
||||
|
||||
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
|
||||
设计要点:
|
||||
- 直接 httpx 直连 DashScope,不走 ai_client
|
||||
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
|
||||
- enable_thinking=false + response_format=json_object
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- timeout=25s
|
||||
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
|
||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
||||
"""
|
||||
|
||||
@@ -15,7 +14,6 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
@@ -23,14 +21,27 @@ 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 _get_vision_config(variant: str = "primary") -> tuple[str, str, str]:
|
||||
"""从 ai_router 获取 image_analysis 配置,返回 (api_key, base_url, model)。"""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
# 先尝试 lite,再 fallback
|
||||
client = ai_router.get_vision_client("image_analysis", variant=variant)
|
||||
if client and client.is_available:
|
||||
return client.api_key, client.base_url, client.model
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败 (%s),fallback 环境变量: %s", variant, e)
|
||||
|
||||
# Fallback: 环境变量
|
||||
import os
|
||||
|
||||
api_key = os.environ.get("DASHSCOPE_API_KEY", "")
|
||||
return api_key, "https://dashscope.aliyuncs.com/compatible-mode/v1", "qwen3.7-plus"
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
@@ -42,7 +53,7 @@ def call_pro_vlm(
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
|
||||
api_key = _api_key()
|
||||
api_key, base_url, model = _get_vision_config("primary")
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
|
||||
return None
|
||||
@@ -50,7 +61,7 @@ def call_pro_vlm(
|
||||
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": _PRO_MODEL,
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
@@ -69,7 +80,7 @@ def call_pro_vlm(
|
||||
}
|
||||
try:
|
||||
r = httpx.post(
|
||||
f"{_BASE_URL}/chat/completions",
|
||||
f"{base_url.rstrip('/')}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
@@ -90,7 +101,7 @@ def call_pro_vlm(
|
||||
reasoning_tokens = ctd.get("reasoning_tokens", 0)
|
||||
logger.info(
|
||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_PRO_MODEL,
|
||||
model,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
|
||||
@@ -1,22 +1,20 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 快速路径:qwen3.8-flash(阿里云百炼/DashScope)强约束 JSON-only 调用。
|
||||
"""V2 快速路径:image_analysis capability(默认 qwen3.8-flash / DashScope)强约束 JSON-only 调用。
|
||||
|
||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
||||
设计要点:
|
||||
- 直接用 httpx 发最小 payload 到 DashScope OpenAI 兼容 endpoint,不走 ai_client 包装
|
||||
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
|
||||
- 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 读取
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
@@ -24,15 +22,26 @@ 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 _get_vision_config() -> tuple[str, str, str]:
|
||||
"""从 ai_router 获取 image_analysis 配置,返回 (api_key, base_url, model)。"""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="primary")
|
||||
if client and client.is_available:
|
||||
return client.api_key, client.base_url, client.model
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败,fallback 环境变量: %s", e)
|
||||
|
||||
# Fallback: 环境变量
|
||||
import os
|
||||
|
||||
api_key = os.environ.get("DASHSCOPE_API_KEY", "")
|
||||
return api_key, "https://dashscope.aliyuncs.com/compatible-mode/v1", "qwen3.8-flash"
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
@@ -57,16 +66,16 @@ def call_fast_json(
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
|
||||
api_key = _api_key()
|
||||
api_key, base_url, model = _get_vision_config()
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 fast_json")
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
|
||||
url = f"{_BASE_URL}/chat/completions"
|
||||
url = f"{base_url.rstrip('/')}/chat/completions"
|
||||
payload: dict[str, Any] = {
|
||||
"model": _FAST_MODEL,
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
@@ -118,7 +127,7 @@ def call_fast_json(
|
||||
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,
|
||||
model,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
|
||||
@@ -351,11 +351,23 @@ 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._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", "voice-enrollment")
|
||||
self._audio_url_signer = audio_url_signer
|
||||
|
||||
|
||||
@@ -55,9 +55,12 @@ _LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
|
||||
class Reviewer:
|
||||
def __init__(self, client=None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
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:暂停积分系统但保留全部代码/表/接口)。
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
"""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,535 @@
|
||||
"""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 typing import Optional
|
||||
|
||||
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
|
||||
|
||||
|
||||
# ── 客户端包装 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""统一 LLM 客户端接口"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 45,
|
||||
max_retries: int = 1,
|
||||
max_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
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.max_retries = max_retries
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
def chat_completion(self, messages: list[dict], **kwargs) -> dict:
|
||||
"""调用 LLM chat completion API"""
|
||||
import httpx
|
||||
|
||||
url = f"{self.base_url.rstrip('/')}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
}
|
||||
if self.max_tokens is not None:
|
||||
payload["max_tokens"] = self.max_tokens
|
||||
if self.temperature is not None:
|
||||
payload["temperature"] = self.temperature
|
||||
payload.update(self.extra_params)
|
||||
payload.update(kwargs)
|
||||
|
||||
resp = httpx.post(url, json=payload, headers=headers, timeout=self.timeout)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class VisionClient(LLMClient):
|
||||
"""VLM 多模态客户端(继承 LLM,增加图片支持)"""
|
||||
|
||||
def call_with_images(self, image_urls: list[str], system_prompt: str, user_prompt: str, **kwargs) -> dict:
|
||||
"""VLM 多图片调用"""
|
||||
content: list[dict] = [{"type": "text", "text": user_prompt}]
|
||||
for url in image_urls:
|
||||
content.append({"type": "image_url", "image_url": {"url": url}})
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": content},
|
||||
]
|
||||
return self.chat_completion(messages, **kwargs)
|
||||
|
||||
|
||||
class TTSClient:
|
||||
"""TTS 客户端"""
|
||||
|
||||
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 / 独立脚本场景"""
|
||||
# 方式1:全局 SessionLocal(worker/api 启动时通过 build_session_factory 设置)
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is not None:
|
||||
return SessionLocal()
|
||||
|
||||
# 方式2:尝试 worker_app.db
|
||||
try:
|
||||
from worker_app.db import SessionLocal as WorkerSL
|
||||
|
||||
if WorkerSL is not None:
|
||||
return WorkerSL()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# 方式3:尝试 api 的 db 模块
|
||||
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) -> LLMClient:
|
||||
return LLMClient(
|
||||
provider=model.provider,
|
||||
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,
|
||||
)
|
||||
|
||||
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig) -> VisionClient:
|
||||
return VisionClient(
|
||||
provider=model.provider,
|
||||
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,
|
||||
)
|
||||
|
||||
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") -> LLMClient | None:
|
||||
"""获取 LLM 客户端"""
|
||||
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") -> VisionClient | None:
|
||||
"""获取 VLM 客户端"""
|
||||
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) -> LLMClient | None:
|
||||
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
|
||||
|
||||
return LLMClient(
|
||||
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) -> VisionClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "dashscope_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
model = "qwen3.8-flash"
|
||||
|
||||
return VisionClient(
|
||||
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", "https://dashscope.aliyuncs.com/api/v1")
|
||||
model = getattr(settings, "cosyvoice_model", "cosyvoice-v3-flash")
|
||||
|
||||
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", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
model = getattr(settings, "doubao_image_model", "doubao-seedream-5-0-flash-260915")
|
||||
|
||||
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", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
model = getattr(settings, "doubao_video_model", "doubao-seedance-2-5-260628")
|
||||
|
||||
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,366 @@
|
||||
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
# ── 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
|
||||
|
||||
_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)
|
||||
self.assertTrue(hasattr(client, "call_with_images"))
|
||||
|
||||
@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_llm_client_available(self):
|
||||
c = _ai_router.LLMClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
def test_llm_client_unavailable_no_key(self):
|
||||
c = _ai_router.LLMClient(provider="p", api_key="", base_url="u", model="m")
|
||||
self.assertFalse(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,7 +102,7 @@ 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"
|
||||
@@ -112,7 +112,7 @@ class TestSharedSettingsDefaults:
|
||||
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 默认空字符串"""
|
||||
|
||||
Reference in New Issue
Block a user