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
|
||||
Reference in New Issue
Block a user