"""Ollama API 客户端 通过 HTTP 流式调用本地 Ollama 服务,并测量真实推理指标: - TTFT: 首个 token 到达时间(从发起请求开始计时) - Prefill: prompt 预填充耗时(取自 Ollama 返回的 prompt_eval_duration) - Decode Speed: 解码速度 tokens/s(eval_count / eval_duration) - E2E: 端到端总耗时 依赖:requests(Ollama 服务需已启动,默认 http://localhost:11434) """ import json import logging import time import requests from . import config from .client_base import InferenceResult, ClientError logger = logging.getLogger(__name__) class OllamaError(ClientError): """Ollama 调用相关错误""" class OllamaClient: """封装对本地 Ollama 服务的调用""" def __init__(self, base_url: str = None, model: str = None, timeout: int = 180): self.base_url = (base_url or config.OLLAMA_BASE_URL).rstrip("/") self.model = model or config.DEFAULT_MODEL self.timeout = timeout self._session = requests.Session() # ---------- 服务探测 ---------- def ping(self) -> bool: """探测 Ollama 服务是否可用""" try: return self._session.get(f"{self.base_url}/api/tags", timeout=5).ok except Exception: return False def list_models(self) -> list: """列出本地已拉取的模型""" try: resp = self._session.get(f"{self.base_url}/api/tags", timeout=5) resp.raise_for_status() return [m["name"] for m in resp.json().get("models", [])] except requests.exceptions.ConnectionError as e: raise OllamaError( f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve'" ) from e except Exception as e: raise OllamaError(f"获取模型列表失败: {e}") from e # ---------- 文本生成(/api/generate)---------- def generate(self, prompt: str, system: str = None, options: dict = None) -> InferenceResult: """单轮文本生成""" payload = { "model": self.model, "prompt": prompt, "stream": True, "options": options or {}, } if system: payload["system"] = system return self._request(payload) # ---------- 多轮对话(/api/chat)---------- def chat(self, messages: list, options: dict = None) -> InferenceResult: """多轮对话,messages 为 [{"role": "user"/"assistant", "content": ...}]""" payload = { "model": self.model, "messages": messages, "stream": True, "options": options or {}, } return self._request(payload, chat=True) # ---------- 核心请求逻辑 ---------- def _request(self, payload: dict, chat: bool = False) -> InferenceResult: result = InferenceResult() chunks = [] final = {} ttft_start = None start = time.perf_counter() endpoint = "chat" if chat else "generate" try: resp = self._session.post( f"{self.base_url}/api/{endpoint}", json=payload, stream=True, timeout=self.timeout, ) resp.raise_for_status() for line in resp.iter_lines(decode_unicode=True): if not line: continue data = json.loads(line) if chat: piece = (data.get("message") or {}).get("content", "") else: piece = data.get("response", "") done = data.get("done", False) if not done: # 首个非空片段出现时刻记为 TTFT if ttft_start is None and piece: ttft_start = time.perf_counter() - start if piece: chunks.append(piece) else: final = data except requests.exceptions.ConnectionError as e: raise OllamaError( f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve' 且模型 {self.model} 已拉取" ) from e except requests.exceptions.Timeout as e: raise OllamaError(f"Ollama 请求超时({self.timeout}s): {e}") from e except (json.JSONDecodeError, KeyError) as e: raise OllamaError(f"Ollama 返回数据格式错误: {e}") from e result.response = "".join(chunks) result.e2e_ms = (time.perf_counter() - start) * 1000 result.ttft_ms = (ttft_start * 1000) if ttft_start else 0.0 # Ollama 结束块中的指标(单位为纳秒) prompt_eval_count = final.get("prompt_eval_count", 0) or 0 prompt_eval_dur_ns = final.get("prompt_eval_duration", 0) or 0 eval_count = final.get("eval_count", 0) or 0 eval_dur_ns = final.get("eval_duration", 0) or 0 result.prompt_tokens = int(prompt_eval_count) result.completion_tokens = int(eval_count) result.total_tokens = result.prompt_tokens + result.completion_tokens result.prefill_ms = prompt_eval_dur_ns / 1_000_000 if eval_dur_ns: result.decode_speed_tok_s = result.completion_tokens / (eval_dur_ns / 1e9) # TTFT 未测到时(如空回复),回退到 prefill 时间 if result.ttft_ms <= 0 and result.prefill_ms > 0: result.ttft_ms = result.prefill_ms logger.debug( "推理完成: e2e=%.1fms ttft=%.1fms decode=%.1f tok/s tokens=%d/%d", result.e2e_ms, result.ttft_ms, result.decode_speed_tok_s, result.prompt_tokens, result.completion_tokens, ) return result