Compare commits

...

1 Commits

Author SHA1 Message Date
Xiaoxia Agent 2d811fccf1 feat: AI模型路由层 — 统一模型配置读取与客户端构建
CI/CD Pipeline / Build Production API Image (pull_request) Blocked by required conditions
CI/CD Pipeline / Build Production Web Image (pull_request) Blocked by required conditions
CI/CD Pipeline / Build Production Worker Image (pull_request) Blocked by required conditions
CI/CD Pipeline / Deploy Production (pull_request) Blocked by required conditions
CI/CD Pipeline / Production Browser E2E (pull_request) Blocked by required conditions
CI/CD Pipeline / Canary Release to Production (pull_request) Blocked by required conditions
CI/CD Pipeline / CI Gate (pull_request) Blocked by required conditions
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 0s
AI Code Review / AI Code Review (pull_request) Has started running
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Validate - Style (pull_request) Has started running
CI/CD Pipeline / Validate - Security (pull_request) Has started running
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has started running
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 53s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 55s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m30s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m54s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m7s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has started running
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 3m22s
核心改动:
- 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个单测全绿
2026-10-06 13:14:05 +08:00
12 changed files with 1300 additions and 79 deletions
@@ -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"
)
)
+20 -8
View File
@@ -421,11 +421,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": ""}
@@ -478,8 +478,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)
@@ -829,11 +835,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)
@@ -908,8 +914,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),
+15 -3
View File
@@ -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
+6 -3
View File
@@ -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
View File
@@ -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:暂停积分系统但保留全部代码/表/接口)。
+60
View File
@@ -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
+535
View File
@@ -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()
+366
View File
@@ -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()
+3 -3
View File
@@ -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:
+2 -2
View File
@@ -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 默认空字符串"""