feat(llm): httpx 薄 client——传输重试(429 Retry-After/5xx 退避/401 立断)+参数自愈(400 摘参数)+strict 重试,MockTransport 全覆盖 [vps] [no-doc]
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user