Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 45d964df43 |
+158
-63
@@ -180,8 +180,11 @@ class DoubaoClient:
|
||||
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
|
||||
self.embedding_model: str = settings.doubao_embedding_model
|
||||
# 最近一次视频生成的详细错误(error_code + user_message + raw detail),供上层读取后展示给用户
|
||||
self.last_video_error: dict = {}
|
||||
# chat fallback 信号位:_do_chat_request 遇 404/model_not_exist 时置 True
|
||||
self._last_chat_model_not_found: bool = False
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
@@ -229,6 +232,88 @@ class DoubaoClient:
|
||||
"""是否可用(配置了 API Key)."""
|
||||
return bool(self.api_key)
|
||||
|
||||
@staticmethod
|
||||
def _is_model_not_found_err(e: Exception) -> bool:
|
||||
"""判断 HTTP 异常是否是模型/endpoint 不存在(404 或 body 含 not found 关键词)。"""
|
||||
try:
|
||||
resp = getattr(e, "response", None)
|
||||
if resp is not None:
|
||||
sc = getattr(resp, "status_code", 0)
|
||||
if sc == 404:
|
||||
return True
|
||||
try:
|
||||
body = resp.text or ""
|
||||
except Exception:
|
||||
body = ""
|
||||
bl = body.lower()
|
||||
if sc == 400 and any(
|
||||
k in bl
|
||||
for k in (
|
||||
"modelnotfound",
|
||||
"model_not_exist",
|
||||
"endpointnotfound",
|
||||
"model not found",
|
||||
"endpoint not found",
|
||||
"不存在",
|
||||
)
|
||||
):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
def _do_chat_request(self, url: str, headers: dict, payload: dict, model_label: str) -> Optional[str]:
|
||||
"""单次 chat 请求,带 5xx/网络重试;4xx 直接返回 None,同时置位 self._last_chat_model_not_found。"""
|
||||
self._last_chat_model_not_found = False
|
||||
last_err: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(url, headers=headers, json=payload, timeout=self.timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return (data["choices"][0]["message"]["content"] or "").strip()
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
sc = getattr(getattr(e, "response", None), "status_code", 0) or 0
|
||||
if 400 <= sc < 500:
|
||||
# 判断是否模型/endpoint 不存在
|
||||
is_mnf = sc == 404
|
||||
if not is_mnf:
|
||||
try:
|
||||
body = getattr(e.response, "text", "") or ""
|
||||
except Exception:
|
||||
body = ""
|
||||
bl = body.lower()
|
||||
is_mnf = any(
|
||||
k in bl
|
||||
for k in (
|
||||
"modelnotfound",
|
||||
"model_not_exist",
|
||||
"endpointnotfound",
|
||||
"model not found",
|
||||
"endpoint not found",
|
||||
)
|
||||
)
|
||||
if is_mnf:
|
||||
self._last_chat_model_not_found = True
|
||||
logger.warning(
|
||||
"[豆包chat] HTTP %d model=%s mnf=%s err=%s 不重试", sc, model_label, is_mnf, str(e)[:200]
|
||||
)
|
||||
return None
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"[豆包chat] 调用失败 %.1fs后重试(%d/%d) model=%s err=%s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
model_label,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
logger.error("[豆包chat] 调用最终失败 model=%s err=%s", model_label, last_err)
|
||||
return None
|
||||
|
||||
def chat_completion(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
@@ -238,13 +323,8 @@ class DoubaoClient:
|
||||
) -> Optional[str]:
|
||||
"""调用 Chat Completion 接口.
|
||||
|
||||
Args:
|
||||
messages: 对话消息列表,[{"role": "user"/"system"/"assistant", "content": "..."}]
|
||||
temperature: 采样温度,0-2,默认0.7
|
||||
max_tokens: 最大生成token数,默认1024
|
||||
|
||||
Returns:
|
||||
模型返回的文本内容,失败返回 None
|
||||
- 若指定 model 返回 404/model not found,自动 fallback 到默认推理模型 self.model
|
||||
- 5xx/网络错误按 max_retries 重试;4xx(鉴权/配额/参数/模型不存在)不重试
|
||||
"""
|
||||
if not self.is_available:
|
||||
return None
|
||||
@@ -254,40 +334,24 @@ class DoubaoClient:
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
req_model = model or self.model
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.model,
|
||||
"model": req_model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
}
|
||||
content = self._do_chat_request(url, headers, payload, req_model)
|
||||
if content is not None:
|
||||
return content
|
||||
|
||||
last_error: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
|
||||
logger.error("豆包API调用最终失败: %s", last_error)
|
||||
# fallback:仅当请求的模型是"模型不存在/endpoint 不存在"时才降级到默认模型
|
||||
# (鉴权/配额/限流等全局错误 fallback 也会失败,别浪费请求)
|
||||
if req_model != self.model and getattr(self, "_last_chat_model_not_found", False):
|
||||
logger.warning("[豆包chat] 模型 %s 不存在,fallback 到默认模型 %s", req_model, self.model)
|
||||
payload["model"] = self.model
|
||||
self._last_chat_model_not_found = False
|
||||
return self._do_chat_request(url, headers, payload, self.model)
|
||||
return None
|
||||
|
||||
def vision_completion(
|
||||
@@ -347,41 +411,72 @@ class DoubaoClient:
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
req_model = model or self.vision_model
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.vision_model,
|
||||
"model": req_model,
|
||||
"messages": vision_messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
}
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
last_error: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=req_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包视觉API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
vision_mnf = {"mnf": False}
|
||||
|
||||
logger.error("豆包视觉API调用最终失败: %s", last_error)
|
||||
def _do_vision(curr_model: str) -> Optional[str]:
|
||||
payload["model"] = curr_model
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(url, headers=headers, json=payload, timeout=req_timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return (data["choices"][0]["message"]["content"] or "").strip()
|
||||
except Exception as e:
|
||||
sc = getattr(getattr(e, "response", None), "status_code", 0) or 0
|
||||
if 400 <= sc < 500:
|
||||
is_mnf = sc == 404
|
||||
if not is_mnf:
|
||||
try:
|
||||
body = getattr(e.response, "text", "") or ""
|
||||
except Exception:
|
||||
body = ""
|
||||
bl = body.lower()
|
||||
is_mnf = any(
|
||||
k in bl
|
||||
for k in (
|
||||
"modelnotfound",
|
||||
"model_not_exist",
|
||||
"endpointnotfound",
|
||||
"model not found",
|
||||
"endpoint not found",
|
||||
)
|
||||
)
|
||||
if is_mnf:
|
||||
vision_mnf["mnf"] = True
|
||||
logger.warning(
|
||||
"[豆包vision] HTTP %d model=%s mnf=%s err=%s 不重试", sc, curr_model, is_mnf, str(e)[:200]
|
||||
)
|
||||
return None
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包视觉API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
return None
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
content = _do_vision(req_model)
|
||||
if content is not None:
|
||||
return content
|
||||
# 仅 404/模型不存在时才 fallback 到默认 chat 模型
|
||||
if req_model != self.model and vision_mnf["mnf"]:
|
||||
logger.warning("[豆包vision] 模型 %s 不存在,fallback 到默认模型 %s", req_model, self.model)
|
||||
content = _do_vision(self.model)
|
||||
if content is not None:
|
||||
return content
|
||||
return None
|
||||
|
||||
# ── 视频生成(Seedance 2.5,异步任务)────────────────────────────
|
||||
|
||||
@@ -656,3 +656,76 @@ class TestAiServiceLastVideoError:
|
||||
err = ai_service.get_last_video_error()
|
||||
assert err["error_code"] == "unknown"
|
||||
assert "user_message" in err
|
||||
|
||||
|
||||
class TestChatCompletion404Fallback:
|
||||
"""chat_completion / vision_completion 404 自动 fallback 到默认模型。"""
|
||||
|
||||
def test_chat_404_fallback_to_default_model(self, monkeypatch):
|
||||
"""指定模型 404 时自动用 self.model 重试一次。"""
|
||||
from packages.shared import ai_client as ac_mod
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeResp:
|
||||
def __init__(self, sc, body):
|
||||
self.status_code = sc
|
||||
self._body = body
|
||||
|
||||
def raise_for_status(self):
|
||||
import httpx
|
||||
|
||||
if self.status_code >= 400:
|
||||
raise httpx.HTTPStatusError("x", request=object(), response=self)
|
||||
|
||||
def json(self):
|
||||
return self._body
|
||||
|
||||
@property
|
||||
def text(self):
|
||||
import json as _json
|
||||
|
||||
return _json.dumps(self._body)
|
||||
|
||||
def fake_post(url, headers=None, json=None, timeout=None):
|
||||
calls.append(json.get("model"))
|
||||
if json.get("model") == "doubao-1-5-pro-32k-250115":
|
||||
return FakeResp(404, {"error": {"message": "model not found"}})
|
||||
return FakeResp(200, {"choices": [{"message": {"content": "ok-from-default"}}]})
|
||||
|
||||
monkeypatch.setattr(ac_mod.httpx, "post", fake_post)
|
||||
c = ac_mod.DoubaoClient()
|
||||
c.api_key = "test"
|
||||
c.model = "doubao-seed-1-6-250615"
|
||||
c.max_retries = 0
|
||||
out = c.chat_completion([{"role": "user", "content": "hi"}], model="doubao-1-5-pro-32k-250115")
|
||||
assert out == "ok-from-default"
|
||||
assert calls == ["doubao-1-5-pro-32k-250115", "doubao-seed-1-6-250615"]
|
||||
|
||||
def test_chat_401_does_not_retry(self, monkeypatch):
|
||||
"""401(鉴权失败)不重试也不 fallback,直接返回 None。"""
|
||||
from packages.shared import ai_client as ac_mod
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeResp:
|
||||
def __init__(self, sc):
|
||||
self.status_code = sc
|
||||
|
||||
def raise_for_status(self):
|
||||
import httpx
|
||||
|
||||
raise httpx.HTTPStatusError("x", request=object(), response=self)
|
||||
|
||||
def fake_post(url, headers=None, json=None, timeout=None):
|
||||
calls.append(json.get("model"))
|
||||
return FakeResp(401)
|
||||
|
||||
monkeypatch.setattr(ac_mod.httpx, "post", fake_post)
|
||||
c = ac_mod.DoubaoClient()
|
||||
c.api_key = "test"
|
||||
c.model = "doubao-seed-1-6-250615"
|
||||
c.max_retries = 0
|
||||
out = c.chat_completion([{"role": "user", "content": "hi"}], model="doubao-1-5-pro-32k-250115")
|
||||
assert out is None
|
||||
assert calls == ["doubao-1-5-pro-32k-250115"] # 不重试不 fallback(401 是全局鉴权错,fallback 也没用)
|
||||
|
||||
Reference in New Issue
Block a user