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
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:
@@ -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("}")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user