Files
meeting-lab/src/meeting_lab/llm/ollama.py
T

101 lines
3.2 KiB
Python

"""Minimal Ollama client behavior used by the direct protocol MVP."""
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Any
import requests
DEFAULT_ENDPOINT = "http://127.0.0.1:11434"
class OllamaError(RuntimeError):
"""Raised when Ollama cannot safely complete the requested operation."""
@dataclass(frozen=True)
class OllamaGeneration:
raw_response: dict[str, Any]
text: str
client_wall_time_seconds: float
def ollama_base_url(endpoint: str) -> str:
endpoint = endpoint.rstrip("/")
return endpoint.rsplit("/api/", 1)[0] if "/api/" in endpoint else endpoint
def generate_url(endpoint: str) -> str:
return f"{ollama_base_url(endpoint)}/api/generate"
def require_model(endpoint: str, model: str, timeout: int = 10) -> dict[str, Any]:
base_url = ollama_base_url(endpoint)
try:
response = requests.get(f"{base_url}/api/tags", timeout=timeout)
response.raise_for_status()
data = response.json()
except (requests.RequestException, ValueError) as exc:
raise OllamaError(f"Ollama endpoint is not reachable at {base_url}: {exc}") from exc
models = data.get("models") if isinstance(data, dict) else None
if not isinstance(models, list):
raise OllamaError("Ollama /api/tags returned a malformed response.")
installed = {
item.get("name")
for item in models
if isinstance(item, dict) and isinstance(item.get("name"), str)
}
if model not in installed:
raise OllamaError(f"Requested model is not installed in Ollama: {model}")
return {"base_url": base_url, "model": model, "installed": True}
def generate_once(
endpoint: str,
model: str,
prompt: str,
*,
timeout: int,
num_ctx: int,
num_predict: int,
num_thread: int | None = None,
) -> OllamaGeneration:
payload = {
"model": model,
"prompt": prompt,
"think": False,
"stream": False,
"options": {
"temperature": 0.0,
"num_ctx": num_ctx,
"num_predict": num_predict,
},
}
if num_thread is not None:
if type(num_thread) is not int or num_thread <= 0:
raise ValueError("Protocol thread count must be a positive integer.")
payload["options"]["num_thread"] = num_thread
started = time.perf_counter()
try:
response = requests.post(generate_url(endpoint), json=payload, timeout=timeout)
response.raise_for_status()
data = response.json()
except requests.RequestException as exc:
raise OllamaError(f"Ollama generation request failed: {exc}") from exc
except ValueError as exc:
raise OllamaError("Ollama generation response is not valid JSON.") from exc
wall_time = time.perf_counter() - started
if not isinstance(data, dict):
raise OllamaError("Ollama generation response must be a JSON object.")
text = data.get("response")
if not isinstance(text, str):
raise OllamaError("Ollama generation response has no string 'response' field.")
if not text.strip():
raise OllamaError("Ollama returned an empty protocol.")
return OllamaGeneration(data, text, wall_time)