155 lines
5.7 KiB
Python
155 lines
5.7 KiB
Python
"""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
|