refactor: PR#2220 review fixes — DoubaoClient accepts DB config, remove redundant client wrappers
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 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 3s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web 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 API 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 / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
PR Automation / Auto Approve on CI Green (pull_request) Has been cancelled
Preview Deploy / Deploy Preview Environment (pull_request) Has been cancelled

1. DoubaoClient.__init__() now accepts explicit params (api_key, base_url, model, etc.)
   so AIRouter can inject DB-sourced config instead of DoubaoClient reading settings itself.
   Added provider field, **kwargs on chat_completion/vision_completion for extra params.

2. ai_router.py: Removed LLMClient/VisionClient classes. _build_llm_client() and
   _build_vision_client() now return DoubaoClient instances directly. Kept simple
   TTSClient/ImageGenClient/VideoGenClient wrappers. Fixed all fallback methods to
   read from settings instead of hardcoding URLs/model names.

3. viral_video.py: Intent parsing and storyboard sections now use ai_router.get_llm_client()
   to get DoubaoClient instances directly, instead of manually extracting model_key and
   passing to a shared _llm_client.

4. vision/vlm_fast_json.py & vlm_fallback.py: Now use ai_router.get_vision_client()
   to get DoubaoClient and call vision_completion() with **kwargs (enable_thinking,
   response_format), instead of building their own httpx requests.

5. tests: Updated test_ai_router.py to mock DoubaoClient import (avoiding Python 3.10
   compatibility chain), removed LLMClient/VisionClient references, added vision_completion check.
This commit is contained in:
Xiaoxia Agent
2026-10-06 15:30:27 +08:00
parent ed118d4444
commit 8f949ae5d8
6 changed files with 208 additions and 292 deletions
+26 -42
View File
@@ -499,30 +499,22 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
"suggested_title": "",
}
_s = get_shared_settings()
try:
from packages.shared.ai_router import ai_router
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
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")]:
_client_fast = ai_router.get_llm_client("intent_parsing", variant="primary")
_client_pro = ai_router.get_llm_client("intent_parsing", variant="lite")
for _client, _lbl in [(_client_fast, "fast"), (_client_pro, "pro-fallback")]:
if not _client or not _client.is_available:
continue
try:
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
raw = _llm_client.chat_completion(
logger.info("[爆款视频] 意图解析 model=%s label=%s", _client.model, _lbl)
raw = _client.chat_completion(
[{"role": "system", "content": system}, {"role": "user", "content": user}],
temperature=0.4,
max_tokens=1024,
model=_m,
timeout=60,
) # #2180/#2215: 直接用 client.chat_completion 传 messages list,不再走 call_llm 字符串包装
)
if not raw:
continue
parsed = _parse(raw)
@@ -903,13 +895,14 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
image_analysis=products_summary,
)
def _try_gen(model: str, temp: float, max_tok: int, label: str, tmo: int = 25):
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", model, label, tmo)
raw = _llm_client2.chat_completion(
def _try_gen(client, temp: float, max_tok: int, label: str, tmo: int = 25):
if not client or not client.is_available:
return None
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", client.model, label, tmo)
raw = client.chat_completion(
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
temperature=temp,
max_tokens=max_tok,
model=model,
timeout=tmo,
)
if not raw:
@@ -940,34 +933,22 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
)
return None if is_fallback else normalized
_s = get_shared_settings()
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
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
_client_fast = ai_router.get_llm_client("storyboard", variant="primary")
_client_pro = ai_router.get_llm_client("storyboard", variant="lite")
_script_fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "150"))
_script_pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "150"))
try:
# #2217: doubao-seed-2-1-pro生成长编导脚本高峰期>90s,上调到150s,支持ENV覆盖
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
normalized = _try_gen(_client_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
if normalized is not None:
return normalized
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
normalized = _try_gen(_client_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
if normalized is not None:
return normalized
# 第三次:用主力模型兜底
if _pro and _pro != _fast:
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
# 第三次:用 lite/pro 模型兜底
if _client_pro and _client_pro.is_available:
normalized = _try_gen(_client_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
if normalized is not None:
return normalized
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
@@ -975,6 +956,9 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
except Exception as e:
logger.warning("[爆款视频] 编导脚本生成异常: %s,使用兜底脚本", e, exc_info=True)
return _fallback_script(job)
except Exception as e:
logger.warning("[爆款视频] 编导脚本生成异常: %s,使用兜底脚本", e, exc_info=True)
return _fallback_script(job)
def _step_review(job: ViralVideoJob, copy_result: dict) -> dict:
@@ -3,8 +3,8 @@
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
设计要点:
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
- enable_thinking=false + response_format=json_object
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
- enable_thinking=False + response_format=json_object
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
- timeout=25s
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
@@ -25,25 +25,6 @@ _DEFAULT_TIMEOUT = 30
_DEFAULT_MAX_TOKENS = 800
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(
img_url: str,
idx: int,
@@ -51,70 +32,52 @@ def call_pro_vlm(
timeout: int = _DEFAULT_TIMEOUT,
) -> dict[str, Any] | None:
t0 = time.time()
import httpx
api_key, base_url, model = _get_vision_config("primary")
if not api_key:
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
try:
from packages.shared.ai_router import ai_router
client = ai_router.get_vision_client("image_analysis", variant="primary")
if not client or not client.is_available:
logger.warning("[vision.v2] pro vision client 不可用,跳过")
return None
except Exception as e:
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
return None
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
payload: dict[str, Any] = {
"model": model,
"messages": [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
],
"temperature": 0.3,
"max_tokens": _DEFAULT_MAX_TOKENS,
"stream": False,
"enable_thinking": False,
"response_format": {"type": "json_object"},
}
messages = [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
]
try:
r = httpx.post(
f"{base_url.rstrip('/')}/chat/completions",
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
raw = client.vision_completion(
messages=messages,
images=None, # 图片已在 messages 中
temperature=0.3,
max_tokens=_DEFAULT_MAX_TOKENS,
timeout=timeout,
enable_thinking=False,
response_format={"type": "json_object"},
)
elapsed = time.time() - t0
if r.status_code != 200:
logger.warning("[vision.v2] pro HTTP %d elapsed=%.1fs body=%s", r.status_code, elapsed, r.text[:200])
return None
data = r.json()
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
if not raw:
logger.warning("[vision.v2] pro 返回空 elapsed=%.1fs", elapsed)
return None
usage = data.get("usage") or {}
reasoning_tokens = usage.get("reasoning_tokens", 0)
ctd = usage.get("completion_tokens_details") or {}
if not reasoning_tokens:
reasoning_tokens = ctd.get("reasoning_tokens", 0)
logger.info(
"[vision.v2] pro 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
model,
"[vision.v2] pro 完成 model=%s elapsed=%.1fs",
client.model,
elapsed,
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
reasoning_tokens,
)
s = raw.strip()
if s.startswith("```"):
lines = s.split("\n")
if lines and lines[0].startswith("```"):
lines = lines[1:]
if lines and lines[-1].strip().startswith("```"):
lines = lines[:-1]
s = "\n".join(lines).strip()
s = _strip_code_fence(raw)
lpos, rr = s.find("{"), s.rfind("}")
if lpos >= 0 and rr > lpos:
s = s[lpos : rr + 1]
@@ -135,3 +98,15 @@ def call_pro_vlm(
elapsed = time.time() - t0
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
return None
def _strip_code_fence(s: str) -> str:
s = s.strip()
if s.startswith("```"):
lines = s.split("\n")
if lines and lines[0].startswith("```"):
lines = lines[1:]
if lines and lines[-1].strip().startswith("```"):
lines = lines[:-1]
s = "\n".join(lines).strip()
return s
@@ -3,8 +3,8 @@
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
设计要点:
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
- enable_thinking=false 关闭推理链(reasoning 是延迟主因)
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
- enable_thinking=False 关闭推理链(reasoning 是延迟主因)
- response_format=json_object 强约束JSON输出
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
- max_tokens=350、temperature=0.1(稳定输出 JSON)
@@ -26,24 +26,6 @@ _DEFAULT_TIMEOUT = 15
_DEFAULT_MAX_TOKENS = 350
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:
s = s.strip()
if s.startswith("```"):
@@ -62,76 +44,52 @@ def call_fast_json(
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int = _DEFAULT_MAX_TOKENS,
) -> dict[str, Any] | None:
"""调用 qwen3.8-flash 返回结构化 dict;失败/非 JSON 返回 None。"""
"""调用 vision client 返回结构化 dict;失败/非 JSON 返回 None。"""
t0 = time.time()
import httpx
api_key, base_url, model = _get_vision_config()
if not api_key:
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 fast_json")
try:
from packages.shared.ai_router import ai_router
client = ai_router.get_vision_client("image_analysis", variant="primary")
if not client or not client.is_available:
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
return None
except Exception as e:
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
return None
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
url = f"{base_url.rstrip('/')}/chat/completions"
payload: dict[str, Any] = {
"model": model,
"messages": [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
],
"temperature": 0.1,
"max_tokens": max_tokens,
"stream": False,
"enable_thinking": False,
"response_format": {"type": "json_object"},
}
messages = [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
]
try:
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
raw = client.vision_completion(
messages=messages,
images=None, # 图片已在 messages 中
temperature=0.1,
max_tokens=max_tokens,
timeout=timeout,
enable_thinking=False,
response_format={"type": "json_object"},
)
elapsed = time.time() - t0
if resp.status_code == 400 and "enable_thinking" in resp.text[:300].lower():
logger.warning("[vision.v2] fast_json HTTP 400 thinking 参数不兼容,重试 elapsed=%.1fs", elapsed)
payload.pop("enable_thinking", None)
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if resp.status_code != 200:
logger.warning(
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
)
return None
data = resp.json()
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
if not raw:
logger.warning("[vision.v2] fast_json 返回空 elapsed=%.1fs", elapsed)
return None
usage = data.get("usage") or {}
reasoning_tokens = usage.get("reasoning_tokens", 0)
ctd = usage.get("completion_tokens_details") or {}
if not reasoning_tokens:
reasoning_tokens = ctd.get("reasoning_tokens", 0)
logger.info(
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
model,
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs",
client.model,
elapsed,
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
reasoning_tokens,
)
text = _strip_code_fence(raw)
lpos, r = text.find("{"), text.rfind("}")
+32 -6
View File
@@ -170,13 +170,28 @@ class DoubaoClient:
未配置 API Key 时 is_available 为 False,调用方应降级处理。
"""
def __init__(self) -> None:
def __init__(
self,
api_key: str = "",
base_url: str = "",
model: str = "",
timeout: int = 0,
max_retries: int = 0,
max_tokens: int | None = None,
temperature: float | None = None,
extra_params: dict | None = None,
provider: str = "volcengine",
) -> None:
settings = get_shared_settings()
self.api_key: str = settings.doubao_api_key
self.model: str = settings.doubao_model
self.base_url: str = settings.doubao_base_url.rstrip("/")
self.timeout: int = settings.doubao_timeout
self.max_retries: int = settings.doubao_max_retries
self.api_key: str = api_key or settings.doubao_api_key
self.model: str = model or settings.doubao_model
self.base_url: str = (base_url or settings.doubao_base_url).rstrip("/")
self.timeout: int = timeout or settings.doubao_timeout
self.max_retries: int = max_retries or settings.doubao_max_retries
self.max_tokens: int | None = max_tokens
self.temperature: float | None = temperature
self.extra_params: dict = extra_params or {}
self.provider: str = provider
self.vision_model: str = settings.doubao_vision_model
self.vision_lite_model: str = settings.doubao_vision_lite_model
self.fast_model: str = settings.doubao_fast_model
@@ -243,6 +258,7 @@ class DoubaoClient:
max_tokens: int = 1024,
model: str | None = None,
timeout: int | None = None,
**kwargs,
) -> Optional[str]:
"""调用 Chat Completion 接口.
@@ -268,6 +284,11 @@ class DoubaoClient:
"temperature": temperature,
"max_tokens": max_tokens,
}
# 合并实例级额外参数和调用方传入的额外参数
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
last_error: Optional[Exception] = None
_t0 = time.time()
@@ -319,6 +340,7 @@ class DoubaoClient:
temperature: float = 0.3,
timeout: int | None = None,
model: str | None = None,
**kwargs,
) -> Optional[str]:
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
@@ -374,6 +396,10 @@ class DoubaoClient:
"temperature": temperature,
"max_tokens": max_tokens,
}
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
req_timeout = timeout or self.timeout
last_error: Optional[Exception] = None
+36 -92
View File
@@ -56,80 +56,11 @@ class CapabilityConfig:
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)
# ── 简单包装类(TTS / ImageGen / VideoGen)──────────────────────────────────
class TTSClient:
"""TTS 客户端"""
"""TTS 客户端(简单配置持有者,实际调用由 CosyVoiceService 完成)"""
def __init__(self, provider: str, api_key: str, base_url: str, model: str, timeout: int = 60, extra_params: dict | None = None):
self.provider = provider
@@ -145,7 +76,7 @@ class TTSClient:
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
@@ -161,7 +92,7 @@ class ImageGenClient:
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
@@ -181,13 +112,11 @@ class VideoGenClient:
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
@@ -196,7 +125,6 @@ def _get_session():
except ImportError:
pass
# 方式3:尝试 api 的 db 模块
try:
from app.db import SessionLocal as ApiSL
@@ -334,9 +262,13 @@ class AIRouter:
return cap.fallback_model
return None
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig) -> LLMClient:
return LLMClient(
provider=model.provider,
# ── 构建客户端 ─────────────────────────────────────────────────────────
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig):
"""构建 LLM 客户端 — 返回 DoubaoClient 实例"""
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
@@ -345,11 +277,14 @@ class AIRouter:
max_tokens=cap.max_tokens,
temperature=cap.temperature,
extra_params=cap.extra_params,
provider=model.provider,
)
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig) -> VisionClient:
return VisionClient(
provider=model.provider,
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig):
"""构建 VLM 客户端 — 返回 DoubaoClient 实例(DoubaoClient 已支持 vision_completion)"""
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
api_key=model.api_key,
base_url=model.api_base,
model=model.model_key,
@@ -358,6 +293,7 @@ class AIRouter:
max_tokens=cap.max_tokens,
temperature=cap.temperature,
extra_params=cap.extra_params,
provider=model.provider,
)
def _build_tts_client(self, model: ModelConfig, cap: CapabilityConfig) -> TTSClient:
@@ -390,8 +326,10 @@ class AIRouter:
extra_params=cap.extra_params,
)
def get_llm_client(self, key: str, variant: str = "primary") -> LLMClient | None:
"""获取 LLM 客户端"""
# ── 公开接口 ────────────────────────────────────────────────────────────
def get_llm_client(self, key: str, variant: str = "primary"):
"""获取 LLM 客户端(返回 DoubaoClient 实例)"""
cap = self.get_capability(key)
if cap and cap.is_enabled:
model = self._get_model_or_fallback(cap, variant)
@@ -400,8 +338,8 @@ class AIRouter:
return self._fallback_llm_client(key)
def get_vision_client(self, key: str, variant: str = "primary") -> VisionClient | None:
"""获取 VLM 客户端"""
def get_vision_client(self, key: str, variant: str = "primary"):
"""获取 VLM 客户端(返回 DoubaoClient 实例)"""
cap = self.get_capability(key)
if cap and cap.is_enabled:
model = self._get_model_or_fallback(cap, variant)
@@ -436,7 +374,8 @@ class AIRouter:
# ── Fallback 方法(读 SharedSettings 环境变量)──────────────────────────
def _fallback_llm_client(self, key: str) -> LLMClient | None:
def _fallback_llm_client(self, key: str):
"""Fallback LLM 客户端 — 从 settings 读取配置,不硬编码"""
settings = get_shared_settings()
model_map = {
"intent_parsing": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
@@ -455,7 +394,9 @@ class AIRouter:
if not api_key:
return None
return LLMClient(
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
provider="volcengine",
api_key=api_key,
base_url=base_url,
@@ -464,15 +405,18 @@ class AIRouter:
max_retries=settings.doubao_max_retries,
)
def _fallback_vision_client(self, key: str) -> VisionClient | None:
def _fallback_vision_client(self, key: str):
"""Fallback VLM 客户端 — 从 settings 读取 dashscope 配置,不硬编码"""
settings = get_shared_settings()
api_key = getattr(settings, "dashscope_api_key", "")
if not api_key:
return None
base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1"
model = "qwen3.8-flash"
base_url = getattr(settings, "dashscope_base_url", "") or ""
model = getattr(settings, "dashscope_model", "") or getattr(settings, "doubao_vision_model", "")
return VisionClient(
from packages.shared.ai_client import DoubaoClient
return DoubaoClient(
provider="dashscope",
api_key=api_key,
base_url=base_url,
+36 -7
View File
@@ -59,6 +59,38 @@ _ai_config_version.get_shared_settings = lambda: _mock_settings
sys.modules["packages.shared.config"] = MagicMock()
sys.modules["packages.shared.config"].get_shared_settings = lambda: _mock_settings
# Mock packages.shared.ai_client to avoid triggering packages.shared.__init__ chain
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
_mock_ai_client = MagicMock()
class _FakeDoubaoClient:
"""Fake DoubaoClient for testing - mimics the real interface."""
def __init__(self, api_key="", base_url="", model="", timeout=0, max_retries=0,
max_tokens=None, temperature=None, extra_params=None, provider="volcengine"):
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.max_retries = max_retries
self.max_tokens = max_tokens
self.temperature = temperature
self.extra_params = extra_params or {}
self.provider = provider
self.vision_model = model
@property
def is_available(self):
return bool(self.api_key)
def chat_completion(self, messages, **kwargs):
return None
def vision_completion(self, messages, **kwargs):
return None
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
sys.modules["packages.shared.ai_client"] = _mock_ai_client
_ai_router = _load_module_from_file(
"packages.shared.ai_router",
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_router.py"),
@@ -223,7 +255,8 @@ class TestAIRouter(unittest.TestCase):
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"))
# #2220: vision client is now DoubaoClient with vision_completion
self.assertTrue(hasattr(client, "vision_completion"))
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_tts_client(self, mock_ver):
@@ -341,14 +374,10 @@ class TestModelConfig(unittest.TestCase):
class TestClientAvailability(unittest.TestCase):
"""客户端可用性测试"""
def test_llm_client_available(self):
c = _ai_router.LLMClient(provider="p", api_key="k", base_url="u", model="m")
def test_tts_client_available(self):
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
self.assertTrue(c.is_available)
def test_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)