"""LLM 客户端抽象层 提供统一的 `generate(prompt, json_mode)` 接口,底层实现可切换: - OllamaLLMClient:走 Ollama HTTP /api/generate(本地或远程 Ollama 服务) - OpenAICompatibleLLMClient:走 OpenAI 兼容 /v1/chat/completions(OpenAI / DeepSeek / 智谱 / Qwen API 等) 工厂 `create_llm_client(purpose)` 根据 runtime_settings 选择实现与参数。 所有客户端共享同一接口,业务侧无需感知底层协议差异。 """ from __future__ import annotations from typing import Literal, Protocol, runtime_checkable import httpx import structlog from app.config import settings from app.core.runtime_settings import LlmProviderConfig, get_runtime_settings from app.services.ollama import OllamaClient logger = structlog.get_logger() LlmPurpose = Literal["summarize", "query", "classify"] @runtime_checkable class LLMClient(Protocol): """LLM 客户端统一接口""" async def generate(self, prompt: str, json_mode: bool = False) -> str: """生成文本 Args: prompt: 输入提示词 json_mode: True 时约束输出为 JSON(不支持时降级为普通生成) Returns: 生成的文本内容 """ ... async def is_available(self) -> bool: """检查服务是否可用""" ... class OllamaLLMClient: """Ollama HTTP API 客户端(包装现有 OllamaClient,便于统一接口)""" def __init__(self, base_url: str, model: str, timeout: float = 120.0) -> None: self._inner = OllamaClient(base_url=base_url, model=model, timeout=timeout) async def generate(self, prompt: str, json_mode: bool = False) -> str: return await self._inner.generate(prompt, json_mode=json_mode) async def is_available(self) -> bool: return await self._inner.is_available() class OpenAICompatibleLLMClient: """OpenAI 兼容 chat/completions 客户端 适用于 OpenAI 官方 API、DeepSeek、智谱 ChatGLM、Qwen DashScope、Moonshot 等 所有兼容 OpenAI /v1/chat/completions 协议的服务。 """ def __init__(self, base_url: str, api_key: str, model: str, timeout: float = 120.0, temperature: float = 0.3) -> None: self.base_url = base_url.rstrip("/") self.api_key = api_key self.model = model self.timeout = timeout self.temperature = temperature async def generate(self, prompt: str, json_mode: bool = False) -> str: url = f"{self.base_url}/chat/completions" payload: dict = { "model": self.model, "messages": [{"role": "user", "content": prompt}], "temperature": self.temperature, "stream": False, } if json_mode: payload["response_format"] = {"type": "json_object"} headers = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } async with httpx.AsyncClient(timeout=self.timeout) as client: resp = await client.post(url, json=payload, headers=headers) resp.raise_for_status() data = resp.json() # OpenAI 标准响应结构:choices[0].message.content choices = data.get("choices") or [] if not choices: logger.warning("OpenAI 兼容响应无 choices", model=self.model, raw_keys=list(data.keys())) return "" message = choices[0].get("message") or {} content = message.get("content") or "" logger.debug("OpenAI 兼容生成完成", model=self.model, output_length=len(content)) return content async def is_available(self) -> bool: """简单探测:调 /models 列表接口(多数 OpenAI 兼容服务支持)""" try: async with httpx.AsyncClient(timeout=5.0) as client: resp = await client.get( f"{self.base_url}/models", headers={"Authorization": f"Bearer {self.api_key}"}, ) return resp.status_code == 200 except httpx.HTTPError: return False # ---------------------------------------------------------------------------- # # 工厂 # ---------------------------------------------------------------------------- # # 进程级客户端缓存(避免每请求新建 httpx client) _client_cache: dict[str, LLMClient] = {} def _resolve_provider_config(purpose: LlmPurpose) -> LlmProviderConfig: """从 runtime_settings 取指定用途的 LLM 配置""" rt = get_runtime_settings() return getattr(rt.models, purpose) def _build_client(purpose: LlmPurpose) -> LLMClient: """按 runtime_settings 构造 LLM 客户端""" cfg = _resolve_provider_config(purpose) if cfg.provider == "ollama": base_url = cfg.base_url or settings.ollama_base_url model = cfg.model or settings.ollama_model return OllamaLLMClient(base_url=base_url, model=model, timeout=cfg.timeout) if cfg.provider == "openai_compatible": base_url = cfg.base_url or settings.openai_base_url api_key = cfg.api_key or settings.openai_api_key # openai_compatible 默认模型:若 cfg.model 为空,退化到 OpenAI 通用 chat 模型 model = cfg.model or "gpt-4o-mini" if not api_key: logger.warning("OpenAI 兼容 provider 缺少 api_key,调用大概率会失败", purpose=purpose) return OpenAICompatibleLLMClient( base_url=base_url, api_key=api_key, model=model, timeout=cfg.timeout, temperature=cfg.temperature ) raise ValueError(f"未知 LLM provider: {cfg.provider}") def create_llm_client(purpose: LlmPurpose, *, use_cache: bool = True) -> LLMClient: """创建指定用途的 LLM 客户端 Args: purpose: 用途(summarize / query / classify) use_cache: True 时复用进程级客户端单例(默认);False 每次新建(测试用) Returns: LLMClient 实例 """ if not use_cache: return _build_client(purpose) cache_key = f"{purpose}" if cache_key not in _client_cache: _client_cache[cache_key] = _build_client(purpose) return _client_cache[cache_key] def invalidate_llm_client_cache(purpose: LlmPurpose | None = None) -> None: """清除客户端缓存(runtime_settings 更新后调用,确保后续读取新配置)""" if purpose is None: _client_cache.clear() else: _client_cache.pop(purpose, None)