Add direct protocol MVP core
This commit is contained in:
@@ -0,0 +1,95 @@
|
||||
"""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,
|
||||
) -> OllamaGeneration:
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"think": False,
|
||||
"stream": False,
|
||||
"options": {
|
||||
"temperature": 0.0,
|
||||
"num_ctx": num_ctx,
|
||||
"num_predict": num_predict,
|
||||
},
|
||||
}
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user