import asyncio import json import logging import urllib.error import urllib.request from typing import Any from .ai import AIClientError logger = logging.getLogger(__name__) class OllamaError(AIClientError): pass class OllamaClient: provider = "local" display_name = "Ollama" def __init__(self, base_url: str) -> None: self.base_url = base_url.rstrip("/") def normalize_model(self, model: str) -> str: return model.strip() async def list_models(self) -> list[str]: return await asyncio.to_thread(self._list_models) async def chat( self, model: str, messages: list[dict[str, str]], json_mode: bool = False, ) -> str: return await asyncio.to_thread(self._chat, model, messages, json_mode) def _list_models(self) -> list[str]: data = self._request_json("GET", "/api/tags", None, timeout=15) models = data.get("models", []) return sorted(model["name"] for model in models if model.get("name")) def _chat( self, model: str, messages: list[dict[str, str]], json_mode: bool, ) -> str: payload = { "model": model, "messages": messages, "stream": False, } if json_mode: payload["format"] = "json" data = self._request_json("POST", "/api/chat", payload, timeout=180) prompt_tokens = data.get("prompt_eval_count") completion_tokens = data.get("eval_count") logger.info( "AI usage provider=ollama model=%s input_tokens=%s " "output_tokens=%s total_tokens=%s prompt_eval_ms=%.1f " "generation_ms=%.1f total_ms=%.1f load_ms=%.1f", model, prompt_tokens, completion_tokens, ( prompt_tokens + completion_tokens if isinstance(prompt_tokens, int) and isinstance(completion_tokens, int) else None ), self._duration_ms(data.get("prompt_eval_duration")), self._duration_ms(data.get("eval_duration")), self._duration_ms(data.get("total_duration")), self._duration_ms(data.get("load_duration")), ) content = data.get("message", {}).get("content") if isinstance(content, str) and content.strip(): return content.strip() fallback = data.get("response") if isinstance(fallback, str) and fallback.strip(): return fallback.strip() raise OllamaError("Ollama returned an empty response.") @staticmethod def _duration_ms(value: Any) -> float: return value / 1_000_000 if isinstance(value, int | float) else 0.0 def _request_json( self, method: str, path: str, payload: dict[str, Any] | None, timeout: int, ) -> dict[str, Any]: url = f"{self.base_url}{path}" data = json.dumps(payload).encode("utf-8") if payload is not None else None request = urllib.request.Request( url, data=data, method=method, headers={"Content-Type": "application/json"}, ) try: with urllib.request.urlopen(request, timeout=timeout) as response: raw = response.read().decode("utf-8") except urllib.error.HTTPError as exc: body = exc.read().decode("utf-8", errors="replace")[:400] raise OllamaError(f"Ollama HTTP {exc.code}: {body}") from exc except urllib.error.URLError as exc: raise OllamaError(f"Ollama is unavailable at {self.base_url}: {exc.reason}") from exc except TimeoutError as exc: raise OllamaError("Ollama request timed out.") from exc try: parsed = json.loads(raw) except json.JSONDecodeError as exc: raise OllamaError("Ollama returned invalid JSON.") from exc if not isinstance(parsed, dict): raise OllamaError("Ollama returned an unexpected response.") return parsed