"""统一 LLM 客户端:本地 Ollama / OpenAI 兼容 API(api_key)双模式。 设计:两个模式的底层都是 OpenAI 兼容 /v1 端点,切换只改配置(LLM_MODE), 业务代码一律调用 `chat()` / `embed()`,不感知底层。失败自动指数退避重试, 主模型失败切换备用模型(fallback_chat_model),再失败抛 LLMFailError(§1001), 调用方兜底话术见 FALLBACK_REPLY(降级链与原需求文档一致)。 """ from __future__ import annotations import asyncio import logging import httpx from config.settings import LLMCfg, settings from utils.exceptions import LLMFailError logger = logging.getLogger("tool.llm") # 兜底话术:主模型 + 备用模型均失败时返回 FALLBACK_REPLY = "抱歉,服务暂时不可用,请稍后再试,或拨打客服热线 400-XXX-XXXX。" class LLMClient: def __init__(self, cfg: LLMCfg | None = None): self.cfg = cfg or settings.llm # ---- 模式解析:端点与模型由 LLM_MODE 决定 --------------------------- @property def is_ollama(self) -> bool: return self.cfg.mode.lower() == "ollama" @property def base_url(self) -> str: return (self.cfg.ollama_base + "/v1") if self.is_ollama else self.cfg.api_base @property def chat_model(self) -> str: return self.cfg.ollama_chat_model if self.is_ollama else self.cfg.api_chat_model @property def embed_model(self) -> str: return self.cfg.ollama_embed_model if self.is_ollama else self.cfg.api_embed_model def _headers(self) -> dict: h = {"Content-Type": "application/json"} if not self.is_ollama: h["Authorization"] = f"Bearer {self.cfg.api_key}" return h def describe(self) -> str: """健康检查/日志用:当前模式 + 端点 + 模型。""" return f"{self.cfg.mode} [{self.base_url}] chat={self.chat_model} embed={self.embed_model}" # ---- 对话 ----------------------------------------------------------- async def chat( self, messages: list[dict], *, temperature: float | None = None, max_tokens: int | None = None, model: str | None = None, ) -> str: """chat/completions,返回 assistant 文本。 模型降级链:model > chat_model > fallback_chat_model;每模型内部退避重试。 """ models = [model or self.chat_model] if self.cfg.fallback_chat_model and self.cfg.fallback_chat_model not in models: models.append(self.cfg.fallback_chat_model) last_err: Exception | None = None for m in models: try: return await self._chat_once(m, messages, temperature, max_tokens) except Exception as e: # noqa: BLE001 重试/切换模型,异常向上收敛 logger.warning("chat model=%r failed: %s: %s", m, type(e).__name__, e) last_err = e raise LLMFailError() from last_err async def _chat_once( self, model: str, messages: list[dict], temperature: float | None, max_tokens: int | None, ) -> str: payload = { "model": model, "messages": messages, "temperature": self.cfg.temperature if temperature is None else temperature, "max_tokens": self.cfg.max_tokens if max_tokens is None else max_tokens, "stream": False, } url = f"{self.base_url}/chat/completions" # 指数退避:retry_backoff_sec → 翻倍 → …,最多 max_retries 次 for attempt in range(self.cfg.max_retries): try: async with httpx.AsyncClient(timeout=self.cfg.timeout) as client: r = await client.post(url, headers=self._headers(), json=payload) r.raise_for_status() return r.json()["choices"][0]["message"]["content"] except Exception: if attempt == self.cfg.max_retries - 1: raise await asyncio.sleep(self.cfg.retry_backoff_sec * (2**attempt)) raise LLMFailError() # 理论不可达,防御 # ---- Embedding(RAG 入口的统一向量化) ------------------------------- async def embed(self, texts: list[str]) -> list[list[float]]: """批量向量化,返回与 input 一一对应的向量列表(维度由模型决定)。""" url = f"{self.base_url}/embeddings" payload = {"model": self.embed_model, "input": texts} async with httpx.AsyncClient(timeout=self.cfg.timeout) as client: r = await client.post(url, headers=self._headers(), json=payload) r.raise_for_status() data = r.json() return [item["embedding"] for item in data["data"]] async def embed_one(self, text: str) -> list[float]: return (await self.embed([text]))[0] # ---- 就绪探测(启动/健康检查可选接入) -------------------------------- async def check_health(self) -> None: """验证当前模式端点可达。ollama 走 /api/tags,API 走 /models(需鉴权,401 也算可达)。""" base = self.cfg.ollama_base if self.is_ollama else self.cfg.api_base path = "/api/tags" if self.is_ollama else "/models" async with httpx.AsyncClient(timeout=min(self.cfg.timeout, 10)) as client: r = await client.get(f"{base}{path}", headers=self._headers()) if r.status_code >= 500: r.raise_for_status() # 全局单例:Agent 统一引入 llm = LLMClient()