"""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)