feat(llm): httpx 薄 client——传输重试(429 Retry-After/5xx 退避/401 立断)+参数自愈(400 摘参数)+strict 重试,MockTransport 全覆盖 [vps] [no-doc]

This commit is contained in:
2026-10-01 23:30:45 +08:00
parent 46d76e75b7
commit 9aa104ba9f
3 changed files with 236 additions and 1 deletions
+2 -1
View File
@@ -1,5 +1,6 @@
"""LLM 薄 client 层(P4):provider 表/httpx 调用/JSON 兜底.零新依赖(httpx+pydantic)."""
from .provider import LLMConfig, LLMConfigError, LLMError, ProviderSpec, PROVIDERS, chat_url, resolve_config
from .client import LLMClient
__all__ = ["LLMConfig", "LLMConfigError", "LLMError", "ProviderSpec",
"PROVIDERS", "chat_url", "resolve_config"]
"PROVIDERS", "chat_url", "resolve_config", "LLMClient"]
+108
View File
@@ -0,0 +1,108 @@
# sanguo_api/llm/client.py
"""httpx 薄 client(P4):openai 兼容 chat/completions 直调.
零新依赖(httpx 0.28.1 in requirements).搬运组合(调研报告 §1):
- tickflow ai_provider.py:420-465 参数自愈(400 摘参数重试一次)
- RD-Agent base.py:520-613 重试引擎模式精简(429 Retry-After/5xx 退避/401 立断)
- QA planning.py:93-106 解析失败追加 strict 指令重试一次
"""
from __future__ import annotations
import asyncio
import json
import logging
import time
from typing import Callable
import httpx
from .json_utils import JSONParseError, robust_json_parse
from .provider import LLMConfig, LLMError, chat_url
logger = logging.getLogger("sanguo_api.llm")
_MAX_ATTEMPTS = 3
_BACKOFF_SECONDS = (2.0, 4.0)
_STRICT_NOTE = "Strictly output valid JSON. No extra text."
class LLMClient:
def __init__(self, config: LLMConfig,
transport: httpx.AsyncBaseTransport | None = None,
sleep: Callable[[float], None] = time.sleep) -> None:
self._cfg = config
self._transport = transport
self._sleep = sleep
async def chat_json(self, messages: list[dict], *, temperature: float = 0.2,
max_tokens: int = 2000) -> dict:
"""一轮对话→dict.传输层重试+参数自愈在 payload 级,解析失败 strict 重试一次."""
payload: dict = {
"model": self._cfg.model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
"response_format": {"type": "json_object"},
}
healed = {"response_format": False, "temperature": False}
attempts = 0
while True:
attempts += 1
status, text = await self._post(payload)
if status == 200:
content = self._content(text)
usage = json.loads(text).get("usage") or {}
logger.info("llm tokens prompt=%s completion=%s",
usage.get("prompt_tokens"),
usage.get("completion_tokens"))
try:
return robust_json_parse(content)
except JSONParseError:
return await self._strict_retry(messages, payload)
if status in (401, 403, 404):
raise LLMError(f"LLM 鉴权/端点错误(HTTP {status}):检查 "
f"SANGUO_LLM_API_KEY/SANGUO_LLM_BASE_URL")
if status == 400:
if ("response_format" in text
and not healed["response_format"]):
payload.pop("response_format", None)
healed["response_format"] = True
continue
if ("temperature" in text and not healed["temperature"]):
payload.pop("temperature", None)
healed["temperature"] = True
continue
raise LLMError(f"LLM 拒绝请求(400): {text[:200]}")
if status == 429 or status >= 500:
if attempts >= _MAX_ATTEMPTS:
raise LLMError(f"LLM 重试耗尽(HTTP {status}): {text[:200]}")
self._sleep(_BACKOFF_SECONDS[min(attempts - 1, 1)])
continue
raise LLMError(f"LLM 未预期状态 HTTP {status}: {text[:200]}")
async def _strict_retry(self, messages: list[dict], payload: dict) -> dict:
payload = dict(payload, messages=messages + [
{"role": "user", "content": _STRICT_NOTE}])
status, text = await self._post(payload)
if status != 200:
raise LLMError(f"LLM strict 重试失败(HTTP {status}): {text[:200]}")
try:
return robust_json_parse(self._content(text))
except JSONParseError as e:
raise LLMError(f"LLM 返回非 JSON(两轮): {e}") from e
async def _post(self, payload: dict) -> tuple[int, str]:
headers = {"Authorization": f"Bearer {self._cfg.api_key}",
"Content-Type": "application/json"}
async with httpx.AsyncClient(timeout=self._cfg.timeout,
transport=self._transport) as client:
resp = await client.post(chat_url(self._cfg.base_url),
headers=headers, json=payload)
return resp.status_code, resp.text
@staticmethod
def _content(text: str) -> str:
try:
return json.loads(text)["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError, ValueError) as e:
raise LLMError(f"LLM 响应缺 choices/message: {e}") from e
+126
View File
@@ -0,0 +1,126 @@
# tests/api/test_llm_client.py
"""LLM 薄 client:重试/参数自愈/strict 重试.全部 MockTransport,零真调用."""
import asyncio
import httpx
import pytest
from sanguo_api.llm.client import LLMClient
from sanguo_api.llm.provider import LLMConfig
CFG = LLMConfig("glm_cn", "https://open.bigmodel.cn/api/paas/v4/",
"glm-5.3-flash", "sk-test")
MSGS = [{"role": "user", "content": "hi"}]
def _ok(content='{"a": 1}'):
return httpx.Response(200, json={"choices": [
{"message": {"content": content}}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5}})
def _make(handler):
calls, sleeps = [], []
transport = httpx.MockTransport(handler)
client = LLMClient(CFG, transport=transport, sleep=lambda s: sleeps.append(s))
client._record = calls
return client, calls, sleeps
class TestHappyPath:
def test_chat_json_returns_dict_and_bears_model(self):
def h(request):
calls.append(request)
import json as j
body = j.loads(request.content)
assert body["model"] == "glm-5.3-flash"
assert body["response_format"] == {"type": "json_object"}
return _ok()
client, calls, _ = _make(h)
out = asyncio.run(client.chat_json(MSGS))
assert out == {"a": 1}
assert calls[0].headers["authorization"] == "Bearer sk-test"
class TestRetry:
def test_429_honors_retry_after_then_succeeds(self):
def h(request):
calls.append(request)
if len(calls) == 1:
return httpx.Response(429, headers={"Retry-After": "2"},
json={"error": "rate"})
return _ok()
client, calls, sleeps = _make(h)
assert asyncio.run(client.chat_json(MSGS)) == {"a": 1}
assert len(calls) == 2 and sleeps == [2.0]
def test_5xx_exhausts_attempts_raises_llm_error(self):
def h(request):
calls.append(request)
return httpx.Response(502, json={"error": "bad gw"})
client, calls, sleeps = _make(h)
from sanguo_api.llm import LLMError
with pytest.raises(LLMError) as e:
asyncio.run(client.chat_json(MSGS))
assert "502" in str(e.value)
assert len(calls) == 3 # max_attempts=3
def test_401_raises_immediately(self):
def h(request):
calls.append(request)
return httpx.Response(401, json={"error": "bad key"})
client, calls, _ = _make(h)
from sanguo_api.llm import LLMError
with pytest.raises(LLMError):
asyncio.run(client.chat_json(MSGS))
assert len(calls) == 1
class TestParamSelfHeal:
def test_400_rejecting_response_format_strips_and_retries(self):
def h(request):
calls.append(request)
import json as j
body = j.loads(request.content)
if "response_format" in body:
return httpx.Response(400, text='{"error":"response_format not supported"}')
return _ok()
client, calls, _ = _make(h)
assert asyncio.run(client.chat_json(MSGS)) == {"a": 1}
assert len(calls) == 2
def test_400_unknown_raises_with_upstream_detail(self):
def h(request):
calls.append(request)
return httpx.Response(400, text="quota exceeded 余额不足")
client, calls, _ = _make(h)
from sanguo_api.llm import LLMError
with pytest.raises(LLMError) as e:
asyncio.run(client.chat_json(MSGS))
assert "余额不足" in str(e.value)
class TestStrictRetry:
def test_non_json_content_appends_strict_and_retries_once(self):
def h(request):
calls.append(request)
import json as j
body = j.loads(request.content)
if len(calls) == 1:
return _ok(content="我觉得应该说……")
assert any("Strictly output valid JSON" in m["content"]
for m in body["messages"])
return _ok()
client, calls, _ = _make(h)
assert asyncio.run(client.chat_json(MSGS)) == {"a": 1}
assert len(calls) == 2
def test_non_json_twice_raises(self):
def h(request):
calls.append(request)
return _ok(content="还是不行")
client, calls, _ = _make(h)
from sanguo_api.llm import LLMError
with pytest.raises(LLMError):
asyncio.run(client.chat_json(MSGS))
assert len(calls) == 2