SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
5.72 kB
from __future__ import annotations
import json
import urllib.error
import urllib.request
from typing import Any, Callable
class OllamaError(RuntimeError):
pass
class OllamaClient:
def __init__(
self,
base_url: str,
model: str,
timeout: float = 2.5,
chat_max_tokens: int | None = None,
) -> None:
self.base_url = base_url.rstrip("/")
self.model = model
self.timeout = timeout
self.chat_max_tokens = chat_max_tokens
def _chat_system(self, system: str) -> str:
"""Return the application system prompt unchanged."""
return system
def _num_predict(self, default: int, *, chat: bool = True) -> int:
"""Qwen3's reasoning commonly needs more than a short-chat token budget."""
if chat and self.chat_max_tokens is not None:
return self.chat_max_tokens
return 1024 if self.model.casefold().startswith("qwen3") else default
def is_available(self, timeout: float = 0.35) -> bool:
request = urllib.request.Request(f"{self.base_url}/api/tags", method="GET")
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
return response.status == 200
except (OSError, urllib.error.URLError, TimeoutError):
return False
def list_models(self, timeout: float = 2.0) -> list[str]:
request = urllib.request.Request(f"{self.base_url}/api/tags", method="GET")
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
payload = json.loads(response.read().decode("utf-8"))
except (OSError, urllib.error.URLError, TimeoutError, json.JSONDecodeError):
return []
return [
str(model.get("name"))
for model in payload.get("models", [])
if isinstance(model, dict) and model.get("name")
]
def generate_json(self, system: str, prompt: str) -> dict[str, Any]:
raw_response = self._generate(system, prompt, json_format=True)
try:
parsed = json.loads(raw_response)
except (TypeError, json.JSONDecodeError) as exc:
raise OllamaError("Ollama did not return a valid JSON plan.") from exc
if not isinstance(parsed, dict):
raise OllamaError("Ollama returned an unsupported plan shape.")
return parsed
def generate_text(self, system: str, prompt: str) -> str:
response = self._generate(system, prompt, json_format=False).strip()
if not response:
raise OllamaError("Ollama returned an empty response.")
return response
def generate_text_stream(
self,
system: str,
prompt: str,
on_chunk: Callable[[str], None],
) -> str:
payload: dict[str, Any] = {
"model": self.model,
"system": self._chat_system(system),
"prompt": prompt,
"stream": True,
"options": {"temperature": 0.35, "num_predict": self._num_predict(180)},
}
request = urllib.request.Request(
f"{self.base_url}/api/generate",
data=json.dumps(payload).encode("utf-8"),
headers={"Content-Type": "application/json"},
method="POST",
)
pieces: list[str] = []
try:
with urllib.request.urlopen(request, timeout=self.timeout) as response:
for raw_line in response:
if not raw_line.strip():
continue
event = json.loads(raw_line.decode("utf-8"))
chunk = str(event.get("response", ""))
if chunk:
pieces.append(chunk)
on_chunk(chunk)
if event.get("done"):
break
except (OSError, urllib.error.URLError, json.JSONDecodeError) as exc:
raise OllamaError(f"Ollama streaming request failed: {exc}") from exc
result = "".join(pieces).strip()
if not result:
raise OllamaError("Ollama returned an empty response.")
return result
def _generate(self, system: str, prompt: str, *, json_format: bool) -> str:
payload: dict[str, Any] = {
"model": self.model,
"system": self._chat_system(system),
"prompt": prompt,
"stream": False,
"options": {
"temperature": 0.1 if json_format else 0.35,
"num_predict": self._num_predict(300 if json_format else 180, chat=not json_format),
},
}
if json_format:
payload["format"] = "json"
body = json.dumps(
payload
).encode("utf-8")
request = urllib.request.Request(
f"{self.base_url}/api/generate",
data=body,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(request, timeout=self.timeout) as response:
payload = json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
detail = ""
try:
error_payload = json.loads(exc.read().decode("utf-8"))
detail = str(error_payload.get("error", ""))
except Exception:
pass
raise OllamaError(detail or f"Ollama returned HTTP {exc.code}.") from exc
except (OSError, urllib.error.URLError, json.JSONDecodeError) as exc:
raise OllamaError(f"Ollama request failed: {exc}") from exc
return str(payload.get("response", ""))