101 lines
3.2 KiB
Python
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)
|