1504 lines
65 KiB
Python
1504 lines
65 KiB
Python
"""Native Hermes dashboard API for managing a local Ollama instance."""
|
||
from __future__ import annotations
|
||
|
||
import base64
|
||
import binascii
|
||
import io
|
||
import ipaddress
|
||
import json
|
||
import mimetypes
|
||
import os
|
||
import re
|
||
import socket
|
||
import sqlite3
|
||
import subprocess
|
||
import threading
|
||
import time
|
||
import uuid
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
from datetime import datetime, timedelta, timezone
|
||
from html import unescape
|
||
from html.parser import HTMLParser
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from urllib.error import HTTPError, URLError
|
||
from urllib.parse import unquote, urlencode, urlparse
|
||
from urllib.request import HTTPRedirectHandler, Request, build_opener, urlopen
|
||
from zoneinfo import ZoneInfo
|
||
|
||
from fastapi import APIRouter, HTTPException
|
||
from pydantic import BaseModel, Field
|
||
from hermes_constants import get_hermes_home
|
||
|
||
router = APIRouter()
|
||
|
||
|
||
def _running_in_container() -> bool:
|
||
return Path("/.dockerenv").exists() or Path("/run/.containerenv").exists()
|
||
|
||
|
||
def _ollama_base_url() -> str:
|
||
configured = os.environ.get("OLLAMA_HOST", "").strip()
|
||
if configured:
|
||
return configured.rstrip("/")
|
||
return "http://ollama:11434" if _running_in_container() else "http://localhost:11434"
|
||
|
||
|
||
def _discover_ollama_endpoint() -> None:
|
||
global LOCAL_OLLAMA
|
||
if os.environ.get("OLLAMA_HOST", "").strip() or not _running_in_container():
|
||
return
|
||
candidates = ("http://ollama:11434", "http://host.docker.internal:11434")
|
||
for candidate in candidates:
|
||
try:
|
||
_json_request(candidate + "/api/version", timeout=2)
|
||
except (HTTPError, URLError, OSError, ValueError):
|
||
continue
|
||
LOCAL_OLLAMA = candidate
|
||
return
|
||
|
||
|
||
LOCAL_OLLAMA = _ollama_base_url()
|
||
REMOTE_OLLAMA = "https://ollama.com"
|
||
CATALOG_FILE = "catalog.json"
|
||
MODEL_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/-]{0,190}$")
|
||
MELBOURNE = ZoneInfo("Australia/Melbourne")
|
||
MAX_ATTACHMENT_BYTES = 20 * 1024 * 1024
|
||
MAX_ATTACHMENT_TEXT = 80_000
|
||
MAX_URL_BYTES = 15 * 1024 * 1024
|
||
CHAT_KEEP_ALIVE = -1
|
||
|
||
_jobs: dict[str, dict[str, Any]] = {}
|
||
_jobs_lock = threading.Lock()
|
||
_chat_requests: dict[str, dict[str, Any]] = {}
|
||
_chat_requests_lock = threading.Lock()
|
||
_catalog_lock = threading.Lock()
|
||
_chat_db_init_lock = threading.Lock()
|
||
_chat_db_ready = False
|
||
|
||
|
||
def _chat_db_path() -> Path:
|
||
return _home() / "chat.sqlite3"
|
||
|
||
|
||
def _chat_db() -> sqlite3.Connection:
|
||
global _chat_db_ready
|
||
path = _chat_db_path()
|
||
path.parent.mkdir(parents=True, exist_ok=True)
|
||
with _chat_db_init_lock:
|
||
connection = sqlite3.connect(path, timeout=30)
|
||
connection.row_factory = sqlite3.Row
|
||
connection.execute("PRAGMA journal_mode=WAL")
|
||
connection.execute("PRAGMA foreign_keys=ON")
|
||
if not _chat_db_ready:
|
||
connection.executescript(
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS conversations (
|
||
id TEXT PRIMARY KEY,
|
||
title TEXT NOT NULL DEFAULT 'New conversation',
|
||
model TEXT NOT NULL DEFAULT '',
|
||
models_json TEXT NOT NULL DEFAULT '[]',
|
||
created_at REAL NOT NULL,
|
||
updated_at REAL NOT NULL
|
||
);
|
||
CREATE TABLE IF NOT EXISTS messages (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
||
request_id TEXT NOT NULL,
|
||
role TEXT NOT NULL,
|
||
content TEXT NOT NULL DEFAULT '',
|
||
model TEXT NOT NULL DEFAULT '',
|
||
attachments_json TEXT NOT NULL DEFAULT '[]',
|
||
created_at REAL NOT NULL,
|
||
UNIQUE(request_id, role, model)
|
||
);
|
||
CREATE INDEX IF NOT EXISTS idx_messages_conversation ON messages(conversation_id, id);
|
||
CREATE TABLE IF NOT EXISTS chat_metrics (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
||
request_id TEXT NOT NULL,
|
||
model TEXT NOT NULL,
|
||
status TEXT NOT NULL,
|
||
started_at REAL,
|
||
first_token_at REAL,
|
||
finished_at REAL,
|
||
prompt_eval_count INTEGER,
|
||
eval_count INTEGER,
|
||
total_duration_ns INTEGER,
|
||
load_duration_ns INTEGER,
|
||
prompt_eval_duration_ns INTEGER,
|
||
eval_duration_ns INTEGER,
|
||
error TEXT NOT NULL DEFAULT '',
|
||
created_at REAL NOT NULL,
|
||
UNIQUE(request_id, model)
|
||
);
|
||
CREATE INDEX IF NOT EXISTS idx_metrics_conversation ON chat_metrics(conversation_id, id);
|
||
"""
|
||
)
|
||
connection.commit()
|
||
try:
|
||
path.chmod(0o600)
|
||
except OSError:
|
||
pass
|
||
_chat_db_ready = True
|
||
return connection
|
||
|
||
|
||
def _conversation_id(value: str | None = None) -> str:
|
||
value = str(value or uuid.uuid4().hex).strip()
|
||
if not re.fullmatch(r"[A-Za-z0-9._-]{1,80}", value):
|
||
raise HTTPException(400, "Invalid conversation id")
|
||
return value
|
||
|
||
|
||
def _ensure_conversation(conversation_id: str, model: str, models: list[str], title: str = "") -> None:
|
||
now = time.time()
|
||
title = re.sub(r"\\s+", " ", title.strip())[:100] or "New conversation"
|
||
db = _chat_db()
|
||
try:
|
||
db.execute(
|
||
"INSERT INTO conversations(id,title,model,models_json,created_at,updated_at) VALUES(?,?,?,?,?,?) "
|
||
"ON CONFLICT(id) DO UPDATE SET model=excluded.model, models_json=excluded.models_json, updated_at=excluded.updated_at",
|
||
(conversation_id, title, model, json.dumps(models[:12]), now, now),
|
||
)
|
||
db.commit()
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
def _persist_message(conversation_id: str, request_id: str, role: str, content: str, model: str, attachments: list[dict[str, Any]] | None = None) -> None:
|
||
db = _chat_db()
|
||
try:
|
||
db.execute(
|
||
"INSERT OR IGNORE INTO messages(conversation_id,request_id,role,content,model,attachments_json,created_at) VALUES(?,?,?,?,?,?,?)",
|
||
(conversation_id, request_id, role, content, model, json.dumps(attachments or [])[:10000], time.time()),
|
||
)
|
||
db.execute("UPDATE conversations SET updated_at=? WHERE id=?", (time.time(), conversation_id))
|
||
db.commit()
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
def _metric_values(state: dict[str, Any], status: str | None = None, error: str = "") -> dict[str, Any]:
|
||
started = float(state.get("started_at") or time.time())
|
||
finished = float(state.get("finished_at") or time.time())
|
||
first = state.get("first_token_at")
|
||
prompt_count = state.get("prompt_eval_count")
|
||
eval_count = state.get("eval_count")
|
||
prompt_duration = state.get("prompt_eval_duration")
|
||
eval_duration = state.get("eval_duration")
|
||
total_duration = state.get("total_duration")
|
||
load_duration = state.get("load_duration")
|
||
return {
|
||
"status": status or str(state.get("state") or "unknown"),
|
||
"started_at": started,
|
||
"first_token_at": first,
|
||
"finished_at": finished,
|
||
"time_to_first_token_ms": round((float(first) - started) * 1000, 2) if first else None,
|
||
"total_latency_ms": round((finished - started) * 1000, 2),
|
||
"prompt_eval_count": prompt_count,
|
||
"eval_count": eval_count,
|
||
"total_tokens": (int(prompt_count) + int(eval_count)) if prompt_count is not None and eval_count is not None else None,
|
||
"prompt_eval_duration_ns": prompt_duration,
|
||
"eval_duration_ns": eval_duration,
|
||
"total_duration_ns": total_duration,
|
||
"load_duration_ns": load_duration,
|
||
"prompt_tokens_per_second": round(int(prompt_count) / (int(prompt_duration) / 1e9), 2) if prompt_count and prompt_duration else None,
|
||
"eval_tokens_per_second": round(int(eval_count) / (int(eval_duration) / 1e9), 2) if eval_count and eval_duration else None,
|
||
"error": error,
|
||
}
|
||
|
||
|
||
def _persist_metric(conversation_id: str, request_id: str, model: str, state: dict[str, Any], status: str | None = None, error: str = "") -> dict[str, Any]:
|
||
values = _metric_values(state, status=status, error=error)
|
||
db = _chat_db()
|
||
try:
|
||
db.execute(
|
||
"INSERT INTO chat_metrics(conversation_id,request_id,model,status,started_at,first_token_at,finished_at,prompt_eval_count,eval_count,total_duration_ns,load_duration_ns,prompt_eval_duration_ns,eval_duration_ns,error,created_at) VALUES(?,?,?,?,?,?,?,?,?,?,?,?,?,?,?) "
|
||
"ON CONFLICT(request_id,model) DO UPDATE SET status=excluded.status, started_at=excluded.started_at, first_token_at=excluded.first_token_at, finished_at=excluded.finished_at, prompt_eval_count=excluded.prompt_eval_count, eval_count=excluded.eval_count, total_duration_ns=excluded.total_duration_ns, load_duration_ns=excluded.load_duration_ns, prompt_eval_duration_ns=excluded.prompt_eval_duration_ns, eval_duration_ns=excluded.eval_duration_ns, error=excluded.error",
|
||
(conversation_id, request_id, model, values["status"], values["started_at"], values["first_token_at"], values["finished_at"], values["prompt_eval_count"], values["eval_count"], values["total_duration_ns"], values["load_duration_ns"], values["prompt_eval_duration_ns"], values["eval_duration_ns"], values["error"], time.time()),
|
||
)
|
||
db.commit()
|
||
finally:
|
||
db.close()
|
||
return values
|
||
|
||
|
||
def _row_metric(row: sqlite3.Row) -> dict[str, Any]:
|
||
value = dict(row)
|
||
started = value.get("started_at")
|
||
first = value.get("first_token_at")
|
||
finished = value.get("finished_at")
|
||
value["time_to_first_token_ms"] = round((first - started) * 1000, 2) if first and started else None
|
||
value["total_latency_ms"] = round((finished - started) * 1000, 2) if finished and started else None
|
||
prompt = value.get("prompt_eval_count")
|
||
output = value.get("eval_count")
|
||
value["total_tokens"] = (prompt + output) if prompt is not None and output is not None else None
|
||
value["prompt_tokens_per_second"] = round(prompt / (value["prompt_eval_duration_ns"] / 1e9), 2) if prompt and value.get("prompt_eval_duration_ns") else None
|
||
value["eval_tokens_per_second"] = round(output / (value["eval_duration_ns"] / 1e9), 2) if output and value.get("eval_duration_ns") else None
|
||
return value
|
||
|
||
CAPABILITY_INFO = {
|
||
"completion": "Text generation and chat completion.",
|
||
"tools": "Tool/function calling for agent workflows.",
|
||
"thinking": "Explicit reasoning/thinking output support.",
|
||
"vision": "Image input and visual understanding.",
|
||
"audio": "Audio input or audio-aware inference.",
|
||
"video": "Video input or video-aware inference.",
|
||
}
|
||
|
||
FAMILY_STRENGTHS = {
|
||
"qwen": ["general reasoning", "coding", "multilingual work", "tool use"],
|
||
"qwen35": ["general reasoning", "coding", "long-context work", "tool use"],
|
||
"gemma": ["general assistance", "reasoning", "tool use", "efficient local inference"],
|
||
"nemotron": ["reasoning", "agent workflows", "long-context work", "technical tasks"],
|
||
"deepseek": ["coding", "mathematical reasoning", "technical analysis"],
|
||
"gpt-oss": ["general reasoning", "coding", "agent workflows"],
|
||
"mistral": ["general assistance", "multilingual work", "coding"],
|
||
"kimi": ["long-context work", "reasoning", "coding"],
|
||
"minimax": ["agent workflows", "reasoning", "long-context work"],
|
||
"glm": ["reasoning", "coding", "multilingual work"],
|
||
"lfm2": ["fast local assistants", "low-resource inference", "general chat"],
|
||
}
|
||
|
||
|
||
def _home() -> Path:
|
||
path = get_hermes_home() / "ollama-manager"
|
||
path.mkdir(parents=True, exist_ok=True)
|
||
return path
|
||
|
||
|
||
def _json_request(url: str, method: str = "GET", payload: Any = None, timeout: int = 30) -> dict[str, Any]:
|
||
data = None if payload is None else json.dumps(payload).encode("utf-8")
|
||
headers = {"Accept": "application/json"}
|
||
if data is not None:
|
||
headers["Content-Type"] = "application/json"
|
||
request = Request(url, data=data, headers=headers, method=method)
|
||
with urlopen(request, timeout=timeout) as response:
|
||
raw = response.read()
|
||
value = json.loads(raw.decode("utf-8")) if raw else {}
|
||
return value if isinstance(value, dict) else {}
|
||
|
||
|
||
def _valid_name(name: str) -> str:
|
||
name = str(name or "").strip()
|
||
if not MODEL_RE.fullmatch(name):
|
||
raise HTTPException(400, "Invalid Ollama model name")
|
||
return name
|
||
|
||
|
||
def _local_tags() -> list[dict[str, Any]]:
|
||
_discover_ollama_endpoint()
|
||
try:
|
||
payload = _json_request(LOCAL_OLLAMA + "/api/tags", timeout=15)
|
||
models = payload.get("models", [])
|
||
return [item for item in models if isinstance(item, dict)]
|
||
except (HTTPError, URLError, OSError, ValueError):
|
||
return []
|
||
|
||
|
||
def _local_ps() -> list[dict[str, Any]]:
|
||
try:
|
||
payload = _json_request(LOCAL_OLLAMA + "/api/ps", timeout=10)
|
||
models = payload.get("models", [])
|
||
return [item for item in models if isinstance(item, dict)]
|
||
except (HTTPError, URLError, OSError, ValueError):
|
||
return []
|
||
|
||
|
||
def _local_ps() -> list[dict[str, Any]]:
|
||
try:
|
||
payload = _json_request(LOCAL_OLLAMA + "/api/ps", timeout=10)
|
||
models = payload.get("models", [])
|
||
return [item for item in models if isinstance(item, dict)]
|
||
except (HTTPError, URLError, OSError, ValueError):
|
||
return []
|
||
|
||
|
||
def _read_meminfo() -> dict[str, int]:
|
||
values: dict[str, int] = {}
|
||
try:
|
||
for line in Path("/proc/meminfo").read_text(encoding="utf-8").splitlines():
|
||
key, _, raw = line.partition(":")
|
||
match = re.search(r"([0-9]+)", raw)
|
||
if match:
|
||
values[key] = int(match.group(1)) * 1024
|
||
except OSError:
|
||
return {}
|
||
return values
|
||
|
||
|
||
def _host_ram_gib() -> float | None:
|
||
"""Return installed system RAM as GiB, detected from the running host."""
|
||
total = _read_meminfo().get("MemTotal", 0)
|
||
return round(total / (1024 ** 3), 1) if total else None
|
||
|
||
|
||
def _gpu_snapshot() -> dict[str, Any]:
|
||
"""Return NVIDIA GPU telemetry when available, without requiring CUDA."""
|
||
query = "name,memory.total,memory.used,memory.free"
|
||
try:
|
||
result = subprocess.run(
|
||
["nvidia-smi", f"--query-gpu={query}", "--format=csv,noheader,nounits"],
|
||
capture_output=True,
|
||
text=True,
|
||
timeout=4,
|
||
check=False,
|
||
)
|
||
except (OSError, subprocess.SubprocessError):
|
||
result = None
|
||
if result and result.returncode == 0:
|
||
gpus = []
|
||
for line in result.stdout.splitlines():
|
||
parts = [part.strip() for part in line.split(",")]
|
||
if len(parts) != 4:
|
||
continue
|
||
try:
|
||
total, used, free = (int(float(value)) * 1024 * 1024 for value in parts[1:])
|
||
except ValueError:
|
||
continue
|
||
gpus.append({"name": parts[0], "total_bytes": total, "used_bytes": used, "free_bytes": free})
|
||
if gpus:
|
||
return {"detected": True, "telemetry_available": True, "gpus": gpus}
|
||
nvidia_present = False
|
||
for vendor in Path("/sys/class/drm").glob("card*/device/vendor"):
|
||
try:
|
||
nvidia_present = nvidia_present or vendor.read_text().strip().lower() == "0x10de"
|
||
except OSError:
|
||
continue
|
||
return {"detected": nvidia_present, "telemetry_available": False, "gpus": []}
|
||
|
||
|
||
def _runtime_snapshot() -> dict[str, Any]:
|
||
mem = _read_meminfo()
|
||
total = mem.get("MemTotal", 0)
|
||
available = mem.get("MemAvailable", mem.get("MemFree", 0))
|
||
swap_total = mem.get("SwapTotal", 0)
|
||
swap_free = mem.get("SwapFree", 0)
|
||
ps_rows = _local_ps()
|
||
tag_rows = {str(row.get("name") or row.get("model")): row for row in _local_tags()}
|
||
model_memory = []
|
||
for row in ps_rows:
|
||
name = str(row.get("name") or row.get("model") or "")
|
||
total_bytes = int(row.get("size") or 0)
|
||
gpu_bytes = int(row.get("size_vram") or 0)
|
||
capability_view = _model_view(tag_rows.get(name, {"name": name}), row)
|
||
model_memory.append({
|
||
"name": name,
|
||
"total_bytes": total_bytes,
|
||
"gpu_bytes": gpu_bytes,
|
||
"ram_bytes": max(0, total_bytes - gpu_bytes),
|
||
"gpu_offload_percent": round(gpu_bytes * 100 / total_bytes, 1) if total_bytes else 0,
|
||
"capabilities": capability_view["capabilities"],
|
||
"capability_breakdown": capability_view["capability_breakdown"],
|
||
"input_modalities": capability_view["input_modalities"],
|
||
"family": capability_view["family"],
|
||
"context_length": capability_view["context_length"],
|
||
"parameter_size": capability_view["parameter_size"],
|
||
"quantization": capability_view["quantization"],
|
||
"permanent": True,
|
||
})
|
||
return {
|
||
"captured_at": time.time(),
|
||
"memory_total_bytes": total,
|
||
"memory_used_bytes": max(0, total - available),
|
||
"memory_available_bytes": available,
|
||
"swap_total_bytes": swap_total,
|
||
"swap_used_bytes": max(0, swap_total - swap_free),
|
||
"model_memory": model_memory,
|
||
"gpu": _gpu_snapshot(),
|
||
}
|
||
|
||
|
||
def _validate_public_url(value: str) -> str:
|
||
parsed = urlparse(value.strip())
|
||
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||
raise HTTPException(400, "URL attachments must use http:// or https://")
|
||
host = parsed.hostname
|
||
try:
|
||
addresses = {info[4][0] for info in socket.getaddrinfo(host, parsed.port or 443, type=socket.SOCK_STREAM)}
|
||
except (OSError, ValueError) as exc:
|
||
raise HTTPException(400, f"Could not resolve URL host: {exc}") from exc
|
||
for address in addresses:
|
||
ip = ipaddress.ip_address(address)
|
||
if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved or ip.is_multicast or ip.is_unspecified:
|
||
raise HTTPException(400, "Private or local URL targets are not allowed")
|
||
return value.strip()
|
||
|
||
|
||
class _SafeRedirectHandler(HTTPRedirectHandler):
|
||
def redirect_request(self, req, fp, code, msg, headers, newurl):
|
||
_validate_public_url(newurl)
|
||
return super().redirect_request(req, fp, code, msg, headers, newurl)
|
||
|
||
|
||
_SAFE_URL_OPENER = build_opener(_SafeRedirectHandler)
|
||
|
||
|
||
def _fetch_attachment_url(value: str) -> tuple[bytes, str, str]:
|
||
value = _validate_public_url(value)
|
||
request = Request(value, headers={"Accept": "text/html, text/plain, application/pdf, image/*", "User-Agent": "Hermes-Ollama-Manager/1.3"})
|
||
try:
|
||
with _SAFE_URL_OPENER.open(request, timeout=30) as response:
|
||
final_url = _validate_public_url(response.geturl())
|
||
content_type = response.headers.get_content_type() if response.headers else "application/octet-stream"
|
||
data = response.read(MAX_URL_BYTES + 1)
|
||
except (HTTPError, URLError, OSError, ValueError) as exc:
|
||
raise HTTPException(400, f"Could not fetch URL: {exc}") from exc
|
||
if len(data) > MAX_URL_BYTES:
|
||
raise HTTPException(413, "URL attachment is larger than 15 MiB")
|
||
return data, content_type, final_url
|
||
|
||
|
||
def _extract_pdf_text(data: bytes, label: str) -> str:
|
||
try:
|
||
from pypdf import PdfReader
|
||
except ImportError as exc:
|
||
raise HTTPException(500, "PDF support requires the pypdf package") from exc
|
||
try:
|
||
reader = PdfReader(io.BytesIO(data))
|
||
text = "\n\n".join(page.extract_text() or "" for page in reader.pages)
|
||
except Exception as exc:
|
||
raise HTTPException(400, f"Could not extract text from PDF {label}: {exc}") from exc
|
||
return text[:MAX_ATTACHMENT_TEXT]
|
||
|
||
|
||
class _PageTextParser(HTMLParser):
|
||
def __init__(self) -> None:
|
||
super().__init__()
|
||
self.parts: list[str] = []
|
||
self._skip = 0
|
||
|
||
def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None:
|
||
if tag.lower() in {"script", "style", "noscript", "svg"}:
|
||
self._skip += 1
|
||
|
||
def handle_endtag(self, tag: str) -> None:
|
||
if tag.lower() in {"script", "style", "noscript", "svg"} and self._skip:
|
||
self._skip -= 1
|
||
|
||
def handle_data(self, data: str) -> None:
|
||
if not self._skip and data.strip():
|
||
self.parts.append(data.strip())
|
||
|
||
|
||
def _extract_page_text(data: bytes, content_type: str) -> str:
|
||
decoded = data.decode("utf-8", errors="replace")
|
||
if "html" in content_type.lower() or re.search(r"<html|<body|<article", decoded, re.I):
|
||
parser = _PageTextParser()
|
||
parser.feed(decoded)
|
||
decoded = "\n".join(parser.parts)
|
||
return unescape(decoded)[:MAX_ATTACHMENT_TEXT]
|
||
|
||
|
||
def _decode_data_url(data_url: str, fallback_mime: str, label: str) -> tuple[bytes, str]:
|
||
match = re.match(r"data:([^;,]+)?;base64,(.*)", data_url or "", re.S)
|
||
if not match:
|
||
raise HTTPException(400, f"Attachment {label} is not a valid base64 data URL")
|
||
mime = (match.group(1) or fallback_mime or "application/octet-stream").lower()
|
||
try:
|
||
data = base64.b64decode(match.group(2), validate=True)
|
||
except (binascii.Error, ValueError) as exc:
|
||
raise HTTPException(400, f"Attachment {label} has invalid base64 data") from exc
|
||
if len(data) > MAX_ATTACHMENT_BYTES:
|
||
raise HTTPException(413, f"Attachment {label} is larger than 20 MiB")
|
||
return data, mime
|
||
|
||
|
||
def _attachment_parts(attachment: "ChatAttachment") -> tuple[str | None, str | None]:
|
||
label = attachment.name or attachment.url or "attachment"
|
||
if attachment.url:
|
||
data, mime, _ = _fetch_attachment_url(attachment.url)
|
||
elif attachment.data_url:
|
||
data, mime = _decode_data_url(attachment.data_url, attachment.mime_type, label)
|
||
else:
|
||
raise HTTPException(400, f"Attachment {label} has no data or URL")
|
||
if mime == "application/pdf" or label.lower().endswith(".pdf"):
|
||
return f"[PDF: {label}]\n{_extract_pdf_text(data, label)}", None
|
||
if mime.startswith("image/"):
|
||
return f"[Image attached: {label}]", base64.b64encode(data).decode("ascii")
|
||
return f"[Text attachment: {label}]\n{_extract_page_text(data, mime)}", None
|
||
|
||
|
||
def _is_mlx(raw: dict[str, Any] | str) -> bool:
|
||
text = str(raw if isinstance(raw, str) else {
|
||
"name": raw.get("name") or raw.get("model"),
|
||
"format": (raw.get("details") or {}).get("format"),
|
||
"capabilities": raw.get("capabilities"),
|
||
}).lower()
|
||
return "mlx" in text or bool(re.search(r"(?:^|[-:])mlx(?:$|[-:])", text))
|
||
|
||
|
||
def _family_key(name: str) -> str:
|
||
return str(name or "").split(":", 1)[0].strip()
|
||
|
||
|
||
def _known_ram_fit(model: dict[str, Any]) -> bool:
|
||
"""Return True when known size/RAM estimates fit installed host RAM."""
|
||
size_gb = model.get("size_gb")
|
||
ram_gb = model.get("expected_ram_gb")
|
||
host_ram_gib = _host_ram_gib()
|
||
return (
|
||
isinstance(size_gb, (int, float))
|
||
and isinstance(ram_gb, (int, float))
|
||
and host_ram_gib is not None
|
||
and float(ram_gb) <= host_ram_gib
|
||
)
|
||
|
||
|
||
def _popular_fit_models(
|
||
raw_rows: list[dict[str, Any]],
|
||
family_rows: dict[str, list[dict[str, Any]]],
|
||
catalog_view,
|
||
installed_names: set[str],
|
||
loaded: dict[str, Any],
|
||
) -> list[dict[str, Any]]:
|
||
"""Select popular models that are usable within this host's RAM budget.
|
||
|
||
Oversized popular entries may be represented by the largest known smaller
|
||
family variant that fits. Unknown entries are never shown.
|
||
"""
|
||
result: list[dict[str, Any]] = []
|
||
seen: set[str] = set()
|
||
seen_footprints: set[tuple[str, float, float]] = set()
|
||
seen_families: set[str] = set()
|
||
for raw in raw_rows:
|
||
if _is_mlx(raw):
|
||
continue
|
||
original = catalog_view(raw, "popular")
|
||
if _known_ram_fit(original):
|
||
candidate = original
|
||
else:
|
||
original_size = original.get("size_gb")
|
||
if not isinstance(original_size, (int, float)):
|
||
continue
|
||
family = _family_key(original["name"])
|
||
variants: list[dict[str, Any]] = []
|
||
for variant_raw in family_rows.get(family, []):
|
||
if _is_mlx(variant_raw):
|
||
continue
|
||
variant = _model_view(variant_raw, loaded.get(str(variant_raw.get("name") or variant_raw.get("model"))), source="popular")
|
||
variant["installed"] = variant["name"] in installed_names
|
||
variant["loaded"] = variant["name"] in loaded
|
||
if _known_ram_fit(variant) and float(variant.get("size_gb", 0)) < float(original_size):
|
||
variants.append(variant)
|
||
if not variants:
|
||
continue
|
||
candidate = max(variants, key=lambda item: (float(item.get("size_gb", 0)), item["name"]))
|
||
candidate["popular_origin"] = original["name"]
|
||
candidate["popular_note"] = f"Smaller fit variant for {original['name']}"
|
||
if candidate["name"] in seen:
|
||
continue
|
||
candidate_family = _family_key(candidate["name"])
|
||
if candidate.get("popular_origin") and candidate_family in seen_families:
|
||
continue
|
||
footprint = (
|
||
candidate_family,
|
||
round(float(candidate.get("size_gb", 0)), 2),
|
||
round(float(candidate.get("expected_ram_gb", 0)), 1),
|
||
)
|
||
if footprint in seen_footprints:
|
||
continue
|
||
seen.add(candidate["name"])
|
||
seen_families.add(candidate_family)
|
||
seen_footprints.add(footprint)
|
||
result.append(candidate)
|
||
for rank, row in enumerate(result, 1):
|
||
row["popularity_rank"] = rank
|
||
return result
|
||
|
||
|
||
def _parse_number(text: Any) -> float | None:
|
||
match = re.search(r"([0-9]+(?:\.[0-9]+)?)", str(text or ""))
|
||
return float(match.group(1)) if match else None
|
||
|
||
|
||
def _ram_estimate(size_bytes: Any, parameter_size: Any, quantization: Any, context_length: Any) -> tuple[float | None, str]:
|
||
size = float(size_bytes or 0)
|
||
if size > 0:
|
||
base = size / (1024 ** 3)
|
||
# Ollama's runtime needs allocator/graph overhead beyond the GGUF blob.
|
||
estimate = base * 1.15 + 0.5
|
||
basis = "disk size × 1.15 + 0.5 GiB runtime overhead"
|
||
else:
|
||
params = _parse_number(parameter_size)
|
||
if params is None:
|
||
return None, "Unavailable: source did not publish a model size"
|
||
bits = 4.5 if "Q4" in str(quantization).upper() else 8.0
|
||
estimate = params * (bits / 8.0) + 0.8
|
||
basis = "parameter/quantization estimate; source size unavailable"
|
||
context = int(context_length or 0)
|
||
if context >= 524288:
|
||
estimate += 2.0
|
||
basis += "; includes large-context allowance"
|
||
elif context >= 131072:
|
||
estimate += 1.0
|
||
basis += "; includes long-context allowance"
|
||
return round(estimate, 1), basis
|
||
|
||
|
||
def _infer_capabilities(name: str, family: str, details: dict[str, Any], advertised: Any) -> list[str]:
|
||
caps = [str(item).lower() for item in advertised or [] if item]
|
||
text = f"{name} {family}".lower()
|
||
if not caps or "completion" not in caps:
|
||
caps.append("completion")
|
||
if any(token in text for token in ("vision", "vl", "gemma4", "muse-glimmer", "qwen3.8")) and "vision" not in caps:
|
||
caps.append("vision")
|
||
if any(token in text for token in ("audio", "omni")) and "audio" not in caps:
|
||
caps.append("audio")
|
||
if "video" in text and "video" not in caps:
|
||
caps.append("video")
|
||
if any(token in text for token in ("thinking", "reasoning", "nemotron", "deepseek", "qwen", "kimi")) and "thinking" not in caps:
|
||
caps.append("thinking")
|
||
if any(token in text for token in ("tool", "agent", "qwen", "gemma", "nemotron", "deepseek", "gpt-oss")) and "tools" not in caps:
|
||
caps.append("tools")
|
||
order = ["completion", "tools", "thinking", "vision", "audio", "video"]
|
||
return [cap for cap in order if cap in set(caps)]
|
||
|
||
|
||
def _architecture(name: str, family: str, details: dict[str, Any]) -> tuple[str, bool]:
|
||
text = f"{name} {family} {details.get('parent_model', '')}".lower()
|
||
moe = any(token in text for token in ("moe", "mixture", "_h_", "a3b"))
|
||
label = "Mixture of Experts (MoE)" if moe else "Dense / single-expert"
|
||
if family:
|
||
label += f" · {family}"
|
||
return label, moe
|
||
|
||
|
||
def _strengths(name: str, family: str, capabilities: list[str], context_length: Any) -> list[str]:
|
||
text = f"{name} {family}".lower()
|
||
result: list[str] = []
|
||
for key, values in FAMILY_STRENGTHS.items():
|
||
if key in text:
|
||
result.extend(values)
|
||
break
|
||
if "tools" in capabilities:
|
||
result.append("tool-enabled automation")
|
||
if "vision" in capabilities:
|
||
result.append("image-aware tasks")
|
||
if int(context_length or 0) >= 131072:
|
||
result.append("long documents and large codebases")
|
||
if not result:
|
||
result = ["general local inference"]
|
||
return list(dict.fromkeys(result))
|
||
|
||
|
||
def _model_view(raw: dict[str, Any], loaded: dict[str, Any] | None = None, source: str = "local") -> dict[str, Any]:
|
||
details = raw.get("details") if isinstance(raw.get("details"), dict) else {}
|
||
name = str(raw.get("name") or raw.get("model") or "")
|
||
family = str(details.get("family") or (details.get("families") or [""])[0] or "")
|
||
capabilities = _infer_capabilities(name, family, details, raw.get("capabilities"))
|
||
context_length = details.get("context_length") or raw.get("context_length")
|
||
architecture, is_moe = _architecture(name, family, details)
|
||
size_bytes = raw.get("size") or 0
|
||
ram_gb, ram_basis = _ram_estimate(size_bytes, details.get("parameter_size"), details.get("quantization_level"), context_length)
|
||
loaded = loaded or {}
|
||
return {
|
||
"name": name,
|
||
"source": source,
|
||
"downloadable": source == "catalog",
|
||
"installed": source == "local",
|
||
"loaded": bool(loaded),
|
||
"size_bytes": int(size_bytes or 0),
|
||
"size_gb": round(float(size_bytes or 0) / (1024 ** 3), 2) if size_bytes else None,
|
||
"size_label": raw.get("size_label") or (f"{round(float(size_bytes or 0) / (1024 ** 3), 2)} GiB" if size_bytes else "Unknown"),
|
||
"loaded_bytes": int(loaded.get("size") or 0),
|
||
"loaded_vram_bytes": int(loaded.get("size_vram") or 0),
|
||
"digest": raw.get("digest", ""),
|
||
"modified_at": raw.get("modified_at"),
|
||
"family": family or "unknown",
|
||
"architecture": architecture,
|
||
"is_moe": is_moe,
|
||
"parameter_size": details.get("parameter_size") or "unknown",
|
||
"quantization": details.get("quantization_level") or "unknown",
|
||
"format": details.get("format") or "unknown",
|
||
"context_length": context_length,
|
||
"input_modalities": raw.get("input_modalities") or (["Text", "Image"] if "vision" in capabilities else ["Text"]),
|
||
"embedding_length": details.get("embedding_length"),
|
||
"capabilities": capabilities,
|
||
"capability_breakdown": {cap: CAPABILITY_INFO[cap] for cap in capabilities if cap in CAPABILITY_INFO},
|
||
"strengths": _strengths(name, family, capabilities, context_length),
|
||
"expected_ram_gb": ram_gb,
|
||
"expected_ram_label": f"{ram_gb:.1f} GiB baseline" if ram_gb is not None else "Unknown",
|
||
"expected_ram_basis": ram_basis,
|
||
}
|
||
|
||
|
||
class _VariantPageParser(HTMLParser):
|
||
"""Extract the public Ollama tag rows without depending on third-party HTML packages."""
|
||
|
||
def __init__(self) -> None:
|
||
super().__init__()
|
||
self._depth = 0
|
||
self._parts: list[str] = []
|
||
self._href = ""
|
||
self.rows: list[dict[str, Any]] = []
|
||
|
||
def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None:
|
||
attrs_map = dict(attrs)
|
||
classes = attrs_map.get("class") or ""
|
||
if tag == "div" and "group" in classes.split() and "px-4" in classes.split():
|
||
self._depth = 1
|
||
self._parts = []
|
||
self._href = ""
|
||
return
|
||
if self._depth:
|
||
if tag == "div":
|
||
self._depth += 1
|
||
if tag == "a" and (attrs_map.get("href") or "").startswith("/library/") and not self._href:
|
||
self._href = attrs_map["href"] or ""
|
||
|
||
def handle_data(self, data: str) -> None:
|
||
if self._depth:
|
||
self._parts.append(data)
|
||
|
||
def handle_endtag(self, tag: str) -> None:
|
||
if not self._depth or tag != "div":
|
||
return
|
||
self._depth -= 1
|
||
if self._depth == 0 and self._href:
|
||
self.rows.append({"href": self._href, "text": " ".join("".join(self._parts).split())})
|
||
self._parts = []
|
||
self._href = ""
|
||
|
||
|
||
def _text_request(url: str, timeout: int = 30) -> str:
|
||
request = Request(url, headers={"Accept": "text/html"})
|
||
with urlopen(request, timeout=timeout) as response:
|
||
return response.read().decode("utf-8", errors="replace")
|
||
|
||
|
||
def _size_bytes(label: str) -> int:
|
||
match = re.search(r"([0-9]+(?:\.[0-9]+)?)\s*(KB|MB|GB|TB)", label.upper())
|
||
if not match:
|
||
return 0
|
||
multipliers = {"KB": 10**3, "MB": 10**6, "GB": 10**9, "TB": 10**12}
|
||
return int(float(match.group(1)) * multipliers[match.group(2)])
|
||
|
||
|
||
def _parse_variant_page(html: str, family: str) -> list[dict[str, Any]]:
|
||
parser = _VariantPageParser()
|
||
parser.feed(html)
|
||
rows: list[dict[str, Any]] = []
|
||
seen: set[str] = set()
|
||
for item in parser.rows:
|
||
href = unescape(item["href"])
|
||
name = unquote(href.rsplit("/", 1)[-1])
|
||
if not name or name in seen or not name.startswith(family + ":"):
|
||
continue
|
||
seen.add(name)
|
||
text = item["text"]
|
||
size_match = re.search(r"•\s*([0-9]+(?:\.[0-9]+)?(?:KB|MB|GB|TB))\s*•", text, re.I)
|
||
context_match = re.search(r"([0-9]+)K\s+context window", text, re.I)
|
||
input_match = re.search(r"([^•]+?)\s+input\s+•", text, re.I)
|
||
size_label = size_match.group(1).upper() if size_match else ("cloud" if ("-cloud" in name or ":cloud" in name) else "Unknown")
|
||
modalities = [part.strip() for part in (input_match.group(1).split(",") if input_match else []) if part.strip()]
|
||
caps = ["vision"] if any(part.lower() == "image" for part in modalities) else []
|
||
rows.append({
|
||
"name": name,
|
||
"model": name,
|
||
"size": _size_bytes(size_label),
|
||
"size_label": size_label,
|
||
"modified_at": None,
|
||
"digest": "",
|
||
"details": {"family": family, "context_length": int(context_match.group(1)) * 1024 if context_match else None, "format": "gguf"},
|
||
"capabilities": caps,
|
||
"input_modalities": modalities or ["Text"],
|
||
"is_mlx": _is_mlx(name) or bool(re.search(r"\bMLX\b", text, re.I)),
|
||
})
|
||
return [row for row in rows if not row["is_mlx"]]
|
||
|
||
|
||
def _fetch_family_variants(family: str) -> list[dict[str, Any]]:
|
||
try:
|
||
html = _text_request(f"{REMOTE_OLLAMA}/library/{family}/tags", timeout=30)
|
||
return _parse_variant_page(html, family)
|
||
except (HTTPError, URLError, OSError, ValueError):
|
||
return []
|
||
|
||
|
||
def _catalog_path() -> Path:
|
||
return _home() / CATALOG_FILE
|
||
|
||
|
||
def _read_catalog() -> dict[str, Any]:
|
||
try:
|
||
value = json.loads(_catalog_path().read_text(encoding="utf-8"))
|
||
return value if isinstance(value, dict) else {}
|
||
except (OSError, ValueError):
|
||
return {}
|
||
|
||
|
||
def _write_catalog(value: dict[str, Any]) -> None:
|
||
path = _catalog_path()
|
||
temp = path.with_suffix(".tmp")
|
||
temp.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8")
|
||
temp.replace(path)
|
||
|
||
|
||
def _refresh_family_variants(rows: list[dict[str, Any]]) -> dict[str, list[dict[str, Any]]]:
|
||
families = sorted({_family_key(str(row.get("name") or row.get("model"))) for row in rows if not _is_mlx(row)})
|
||
return {family: _fetch_family_variants(family) for family in families if family}
|
||
|
||
|
||
def refresh_catalog(force: bool = True) -> dict[str, Any]:
|
||
with _catalog_lock:
|
||
try:
|
||
query = urlencode({"limit": 100, "sort": "popular"})
|
||
payload = _json_request(f"{REMOTE_OLLAMA}/api/tags?{query}", timeout=30)
|
||
rows = [item for item in payload.get("models", []) if isinstance(item, dict) and not _is_mlx(item)]
|
||
catalog = {
|
||
"fetched_at": datetime.now(timezone.utc).isoformat(),
|
||
"source": f"{REMOTE_OLLAMA}/api/tags",
|
||
"models": rows,
|
||
"families": _refresh_family_variants(_local_tags()),
|
||
}
|
||
_write_catalog(catalog)
|
||
return catalog
|
||
except (HTTPError, URLError, OSError, ValueError) as exc:
|
||
cached = _read_catalog()
|
||
if cached:
|
||
cached["last_error"] = str(exc)
|
||
return cached
|
||
return {"fetched_at": None, "source": f"{REMOTE_OLLAMA}/api/tags", "models": [], "last_error": str(exc)}
|
||
|
||
|
||
def _catalog_stale(catalog: dict[str, Any]) -> bool:
|
||
raw = catalog.get("fetched_at")
|
||
if not raw:
|
||
return True
|
||
try:
|
||
fetched = datetime.fromisoformat(str(raw).replace("Z", "+00:00"))
|
||
return datetime.now(timezone.utc) - fetched > timedelta(hours=20)
|
||
except ValueError:
|
||
return True
|
||
|
||
|
||
def _ensure_catalog() -> dict[str, Any]:
|
||
catalog = _read_catalog()
|
||
if not catalog.get("models"):
|
||
return refresh_catalog()
|
||
if _catalog_stale(catalog) and not _catalog_lock.locked():
|
||
threading.Thread(target=refresh_catalog, kwargs={"force": True}, daemon=True, name="ollama-catalog-refresh").start()
|
||
return catalog
|
||
|
||
|
||
def _next_refresh() -> str:
|
||
now = datetime.now(MELBOURNE)
|
||
target = now.replace(hour=1, minute=0, second=0, microsecond=0)
|
||
if now >= target:
|
||
target += timedelta(days=1)
|
||
return target.isoformat()
|
||
|
||
|
||
def _job_snapshot() -> list[dict[str, Any]]:
|
||
with _jobs_lock:
|
||
return [dict(item) for item in _jobs.values()]
|
||
|
||
|
||
def _set_job(job_id: str, **values: Any) -> None:
|
||
with _jobs_lock:
|
||
if job_id in _jobs:
|
||
_jobs[job_id].update(values, updated_at=time.time())
|
||
|
||
|
||
def _run_pull(job_id: str, name: str, action: str) -> None:
|
||
try:
|
||
payload = json.dumps({"name": name, "stream": True}).encode("utf-8")
|
||
request = Request(LOCAL_OLLAMA + "/api/pull", data=payload, headers={"Content-Type": "application/json"}, method="POST")
|
||
with urlopen(request, timeout=3600) as response:
|
||
for raw_line in response:
|
||
try:
|
||
event = json.loads(raw_line.decode("utf-8"))
|
||
except ValueError:
|
||
continue
|
||
status = str(event.get("status") or "working")
|
||
completed = int(event.get("completed") or 0)
|
||
total = int(event.get("total") or 0)
|
||
percent = round(completed * 100 / total, 1) if total else None
|
||
_set_job(job_id, status=status, completed=completed, total=total, percent=percent, digest=event.get("digest"))
|
||
if event.get("error"):
|
||
raise RuntimeError(str(event["error"]))
|
||
_set_job(job_id, state="completed", status="success", percent=100)
|
||
except Exception as exc:
|
||
_set_job(job_id, state="failed", status="error", error=str(exc))
|
||
|
||
|
||
def _run_delete(job_id: str, name: str) -> None:
|
||
try:
|
||
_json_request(LOCAL_OLLAMA + "/api/delete", method="DELETE", payload={"name": name}, timeout=120)
|
||
_set_job(job_id, state="completed", status="deleted", percent=100)
|
||
except Exception as exc:
|
||
_set_job(job_id, state="failed", status="error", error=str(exc))
|
||
|
||
|
||
def _new_job(name: str, action: str) -> str:
|
||
job_id = uuid.uuid4().hex
|
||
with _jobs_lock:
|
||
_jobs[job_id] = {"id": job_id, "name": name, "action": action, "state": "running", "status": "starting", "percent": 0, "created_at": time.time(), "updated_at": time.time()}
|
||
target = _run_delete if action == "delete" else _run_pull
|
||
args = (job_id, name) if action == "delete" else (job_id, name, action)
|
||
threading.Thread(target=target, args=args, daemon=True, name=f"ollama-{action}-{job_id[:8]}").start()
|
||
return job_id
|
||
|
||
|
||
class ModelRequest(BaseModel):
|
||
name: str
|
||
|
||
|
||
class ModelsRequest(BaseModel):
|
||
names: list[str] = Field(default_factory=list)
|
||
|
||
|
||
class ChatAttachment(BaseModel):
|
||
name: str = ""
|
||
mime_type: str = ""
|
||
data_url: str | None = None
|
||
url: str | None = None
|
||
|
||
|
||
class ChatRequest(BaseModel):
|
||
model: str = ""
|
||
models: list[str] = Field(default_factory=list)
|
||
message: str = ""
|
||
history: list[dict[str, Any]] = Field(default_factory=list)
|
||
attachments: list[ChatAttachment] = Field(default_factory=list)
|
||
request_id: str = ""
|
||
conversation_id: str = ""
|
||
|
||
|
||
class ChatStopRequest(BaseModel):
|
||
request_id: str
|
||
|
||
def _installed_model_names() -> set[str]:
|
||
return {
|
||
str(row.get("name") or row.get("model"))
|
||
for row in _local_tags()
|
||
if not _is_mlx(row)
|
||
}
|
||
|
||
|
||
def _require_installed_model(name: str) -> str:
|
||
name = _valid_name(name)
|
||
if name not in _installed_model_names():
|
||
raise HTTPException(400, f"Model '{name}' is not installed locally")
|
||
return name
|
||
|
||
|
||
def _ollama_error(exc: HTTPError) -> HTTPException:
|
||
try:
|
||
detail = exc.read().decode("utf-8", errors="replace")[:1000]
|
||
payload = json.loads(detail)
|
||
detail = str(payload.get("error") or detail)
|
||
except (OSError, ValueError):
|
||
detail = str(exc)
|
||
return HTTPException(502, f"Ollama request failed: {detail}")
|
||
|
||
|
||
def _load_model(name: str) -> dict[str, Any]:
|
||
name = _require_installed_model(name)
|
||
try:
|
||
result = _json_request(
|
||
LOCAL_OLLAMA + "/api/generate",
|
||
method="POST",
|
||
payload={"model": name, "prompt": "", "stream": False, "keep_alive": CHAT_KEEP_ALIVE, "options": {"num_predict": 1}},
|
||
timeout=900,
|
||
)
|
||
except HTTPError as exc:
|
||
raise _ollama_error(exc) from exc
|
||
return {"ok": True, "model": name, "response": result.get("response", ""), "runtime": _runtime_snapshot()}
|
||
|
||
|
||
def _chat_payload(body: ChatRequest, model_name: str | None = None) -> dict[str, Any]:
|
||
model = _require_installed_model(model_name or body.model)
|
||
messages: list[dict[str, Any]] = []
|
||
for item in body.history[-24:]:
|
||
role = str(item.get("role") or "")
|
||
content = str(item.get("content") or "").strip()
|
||
if role in {"user", "assistant"} and content:
|
||
messages.append({"role": role, "content": content[:MAX_ATTACHMENT_TEXT]})
|
||
text_parts = [body.message.strip()] if body.message.strip() else []
|
||
images: list[str] = []
|
||
for attachment in body.attachments[:12]:
|
||
text, image = _attachment_parts(attachment)
|
||
if text:
|
||
text_parts.append(text)
|
||
if image:
|
||
images.append(image)
|
||
if not text_parts and not images:
|
||
raise HTTPException(400, "Enter a message or attach a file/URL")
|
||
user_message: dict[str, Any] = {"role": "user", "content": "\n\n".join(text_parts)[:MAX_ATTACHMENT_TEXT]}
|
||
if images:
|
||
user_message["images"] = images
|
||
messages.append(user_message)
|
||
return {"model": model, "messages": messages, "stream": False, "keep_alive": CHAT_KEEP_ALIVE}
|
||
|
||
|
||
class _ChatStopped(Exception):
|
||
pass
|
||
|
||
|
||
def _valid_chat_request_id(value: str) -> str:
|
||
value = str(value or "").strip()
|
||
if not re.fullmatch(r"[A-Za-z0-9._-]{1,80}", value):
|
||
raise HTTPException(400, "Invalid chat request id")
|
||
return value
|
||
|
||
|
||
def _chat_state(request_id: str, **values: Any) -> dict[str, Any]:
|
||
with _chat_requests_lock:
|
||
state = _chat_requests.setdefault(request_id, {"request_id": request_id, "cancel": threading.Event(), "state": "starting", "stage": "Preparing request", "started_at": time.time(), "chunks": 0, "thinking_chars": 0, "response_chars": 0, "first_token_at": None, "prompt_eval_count": None, "eval_count": None, "total_duration": None, "load_duration": None, "prompt_eval_duration": None, "eval_duration": None})
|
||
state.update(values, updated_at=time.time())
|
||
return {key: value for key, value in state.items() if key not in {"cancel", "response"}}
|
||
|
||
|
||
def _stream_chat_request(payload: dict[str, Any], request_id: str, cancel_event: threading.Event | None = None, parent_id: str | None = None) -> dict[str, Any]:
|
||
if cancel_event is not None or parent_id is not None:
|
||
with _chat_requests_lock:
|
||
state = _chat_requests.get(request_id)
|
||
if state:
|
||
if cancel_event is not None: state["cancel"] = cancel_event
|
||
if parent_id is not None: state["parent_id"] = parent_id
|
||
data = json.dumps({**payload, "stream": True}).encode("utf-8")
|
||
request = Request(LOCAL_OLLAMA + "/api/chat", data=data, headers={"Accept": "application/x-ndjson", "Content-Type": "application/json"}, method="POST")
|
||
response_text: list[str] = []
|
||
thinking_text: list[str] = []
|
||
_chat_state(request_id, state="connecting", stage="Connecting to Ollama")
|
||
try:
|
||
with urlopen(request, timeout=1800) as response:
|
||
_chat_state(request_id, response=response, state="generating", stage="Ollama is generating")
|
||
while True:
|
||
with _chat_requests_lock:
|
||
cancelled = bool(_chat_requests.get(request_id, {}).get("cancel", threading.Event()).is_set())
|
||
if cancelled:
|
||
response.close()
|
||
raise _ChatStopped()
|
||
try:
|
||
raw_line = response.readline()
|
||
except (socket.timeout, TimeoutError, OSError, ValueError) as exc:
|
||
with _chat_requests_lock:
|
||
cancelled = bool(_chat_requests.get(request_id, {}).get("cancel", threading.Event()).is_set())
|
||
if cancelled:
|
||
raise _ChatStopped() from exc
|
||
raise
|
||
if not raw_line:
|
||
break
|
||
try:
|
||
event = json.loads(raw_line.decode("utf-8", errors="replace"))
|
||
except ValueError:
|
||
continue
|
||
message = event.get("message") if isinstance(event.get("message"), dict) else {}
|
||
chunk = str(message.get("content") or event.get("response") or "")
|
||
thinking_chunk = str(message.get("thinking") or event.get("thinking") or "")
|
||
if chunk:
|
||
response_text.append(chunk)
|
||
if thinking_chunk:
|
||
thinking_text.append(thinking_chunk)
|
||
if chunk and not _chat_requests.get(request_id, {}).get("first_token_at"):
|
||
_chat_state(request_id, first_token_at=time.time())
|
||
_chat_state(
|
||
request_id,
|
||
state="generating",
|
||
stage="Ollama is generating the response" if chunk else "Ollama is processing model thinking",
|
||
chunks=int(_chat_requests.get(request_id, {}).get("chunks", 0)) + 1,
|
||
thinking_chars=sum(map(len, thinking_text)),
|
||
response_chars=sum(map(len, response_text)),
|
||
prompt_eval_count=event.get("prompt_eval_count", _chat_requests.get(request_id, {}).get("prompt_eval_count")),
|
||
eval_count=event.get("eval_count", _chat_requests.get(request_id, {}).get("eval_count")),
|
||
total_duration=event.get("total_duration", _chat_requests.get(request_id, {}).get("total_duration")),
|
||
load_duration=event.get("load_duration", _chat_requests.get(request_id, {}).get("load_duration")),
|
||
prompt_eval_duration=event.get("prompt_eval_duration", _chat_requests.get(request_id, {}).get("prompt_eval_duration")),
|
||
eval_duration=event.get("eval_duration", _chat_requests.get(request_id, {}).get("eval_duration")),
|
||
)
|
||
if event.get("done"):
|
||
break
|
||
except _ChatStopped:
|
||
_chat_state(request_id, state="stopped", stage="Stopped by user", finished_at=time.time())
|
||
raise
|
||
except Exception:
|
||
_chat_state(request_id, state="failed", stage="Ollama request failed", finished_at=time.time())
|
||
raise
|
||
_chat_state(request_id, state="completed", stage="Response complete", finished_at=time.time())
|
||
with _chat_requests_lock:
|
||
final_state = dict(_chat_requests.get(request_id, {}))
|
||
return {"message": {"role": "assistant", "content": "".join(response_text)}, "done": True, "metrics": _metric_values(final_state, status="completed")}
|
||
|
||
|
||
@router.get("/chat/status/{request_id}")
|
||
def chat_status(request_id: str) -> dict[str, Any]:
|
||
request_id = _valid_chat_request_id(request_id)
|
||
with _chat_requests_lock:
|
||
state = _chat_requests.get(request_id)
|
||
if not state:
|
||
raise HTTPException(404, "Chat request not found")
|
||
result = {key: value for key, value in state.items() if key not in {"cancel", "response"}}
|
||
result["elapsed"] = round(max(0.0, time.time() - float(state.get("started_at") or time.time())), 1)
|
||
return result
|
||
|
||
|
||
@router.post("/chat/stop")
|
||
def chat_stop(body: ChatStopRequest) -> dict[str, Any]:
|
||
request_id = _valid_chat_request_id(body.request_id)
|
||
with _chat_requests_lock:
|
||
state = _chat_requests.get(request_id)
|
||
if not state:
|
||
return {"ok": True, "request_id": request_id, "state": "not_found"}
|
||
state["cancel"].set()
|
||
responses = []
|
||
for child_id, child in _chat_requests.items():
|
||
if child_id == request_id or child.get("parent_id") == request_id:
|
||
child["cancel"].set()
|
||
response = child.get("response")
|
||
if response is not None: responses.append(response)
|
||
child["state"] = "stopping"
|
||
child["stage"] = "Stopping Ollama request"
|
||
child["updated_at"] = time.time()
|
||
state["state"] = "stopping"
|
||
state["stage"] = "Stopping Ollama request"
|
||
state["updated_at"] = time.time()
|
||
for response in responses:
|
||
try: response.close()
|
||
except Exception: pass
|
||
return {"ok": True, "request_id": request_id, "state": "stopping"}
|
||
|
||
|
||
@router.get("/conversations")
|
||
def conversations() -> dict[str, Any]:
|
||
db = _chat_db()
|
||
try:
|
||
rows = db.execute("SELECT id,title,model,models_json,created_at,updated_at,(SELECT COUNT(*) FROM messages m WHERE m.conversation_id=c.id) AS message_count FROM conversations c ORDER BY updated_at DESC LIMIT 100").fetchall()
|
||
result = []
|
||
for row in rows:
|
||
item = dict(row)
|
||
item["models"] = json.loads(item.pop("models_json") or "[]")
|
||
result.append(item)
|
||
return {"conversations": result}
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
@router.get("/conversations/{conversation_id}")
|
||
def conversation(conversation_id: str) -> dict[str, Any]:
|
||
conversation_id = _conversation_id(conversation_id)
|
||
db = _chat_db()
|
||
try:
|
||
row = db.execute("SELECT id,title,model,models_json,created_at,updated_at FROM conversations WHERE id=?", (conversation_id,)).fetchone()
|
||
if not row:
|
||
raise HTTPException(404, "Conversation not found")
|
||
item = dict(row)
|
||
item["models"] = json.loads(item.pop("models_json") or "[]")
|
||
messages = []
|
||
for message in db.execute("SELECT id,request_id,role,content,model,attachments_json,created_at FROM messages WHERE conversation_id=? ORDER BY id", (conversation_id,)).fetchall():
|
||
value = dict(message)
|
||
value["attachments"] = json.loads(value.pop("attachments_json") or "[]")
|
||
messages.append(value)
|
||
metrics = [_row_metric(metric) for metric in db.execute("SELECT * FROM chat_metrics WHERE conversation_id=? ORDER BY id", (conversation_id,)).fetchall()]
|
||
return {"conversation": item, "messages": messages, "metrics": metrics}
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
@router.delete("/conversations/{conversation_id}")
|
||
def delete_conversation(conversation_id: str) -> dict[str, Any]:
|
||
conversation_id = _conversation_id(conversation_id)
|
||
db = _chat_db()
|
||
try:
|
||
db.execute("DELETE FROM conversations WHERE id=?", (conversation_id,))
|
||
db.commit()
|
||
return {"ok": True, "conversation_id": conversation_id}
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
@router.get("/metrics")
|
||
def metrics(limit: int = 100) -> dict[str, Any]:
|
||
limit = max(1, min(int(limit), 500))
|
||
db = _chat_db()
|
||
try:
|
||
rows = [_row_metric(row) for row in db.execute("SELECT * FROM chat_metrics ORDER BY id DESC LIMIT ?", (limit,)).fetchall()]
|
||
completed = [row for row in rows if row["status"] == "completed"]
|
||
def average(key: str) -> float | None:
|
||
values = [float(row[key]) for row in completed if row.get(key) is not None]
|
||
return round(sum(values) / len(values), 2) if values else None
|
||
aggregate = {
|
||
"sample_count": len(rows),
|
||
"completed_count": len(completed),
|
||
"error_count": sum(row["status"] == "failed" for row in rows),
|
||
"stopped_count": sum(row["status"] == "stopped" for row in rows),
|
||
"avg_time_to_first_token_ms": average("time_to_first_token_ms"),
|
||
"avg_total_latency_ms": average("total_latency_ms"),
|
||
"avg_eval_tokens_per_second": average("eval_tokens_per_second"),
|
||
"avg_prompt_tokens_per_second": average("prompt_tokens_per_second"),
|
||
"total_output_tokens": sum(int(row["eval_count"]) for row in completed if row.get("eval_count") is not None),
|
||
}
|
||
return {"metrics": rows, "aggregate": aggregate}
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
@router.get("/runtime")
|
||
def runtime() -> dict[str, Any]:
|
||
return _runtime_snapshot()
|
||
|
||
|
||
@router.post("/chat/load")
|
||
def chat_load(body: ModelRequest) -> dict[str, Any]:
|
||
return _load_model(body.name)
|
||
|
||
|
||
@router.post("/models/load")
|
||
def models_load(body: ModelsRequest) -> dict[str, Any]:
|
||
names = list(dict.fromkeys(_valid_name(name) for name in body.names if str(name).strip()))[:12]
|
||
if not names:
|
||
raise HTTPException(400, "Select at least one model to load")
|
||
results = []
|
||
for name in names:
|
||
try:
|
||
results.append({"name": name, "ok": True, "result": _load_model(name)})
|
||
except Exception as exc:
|
||
results.append({"name": name, "ok": False, "error": str(exc)})
|
||
return {"ok": all(item["ok"] for item in results), "results": results, "runtime": _runtime_snapshot(), "keep_alive": "permanent"}
|
||
|
||
|
||
@router.post("/models/unload")
|
||
def models_unload(body: ModelsRequest) -> dict[str, Any]:
|
||
names = list(dict.fromkeys(_valid_name(name) for name in body.names if str(name).strip()))[:12]
|
||
if not names:
|
||
raise HTTPException(400, "Select at least one model to unload")
|
||
results = []
|
||
for name in names:
|
||
try:
|
||
_require_installed_model(name)
|
||
_json_request(LOCAL_OLLAMA + "/api/generate", method="POST", payload={"model": name, "prompt": "", "stream": False, "keep_alive": 0}, timeout=120)
|
||
results.append({"name": name, "ok": True})
|
||
except Exception as exc:
|
||
results.append({"name": name, "ok": False, "error": str(exc)})
|
||
return {"ok": all(item["ok"] for item in results), "results": results, "runtime": _runtime_snapshot()}
|
||
|
||
|
||
@router.post("/chat")
|
||
def chat(body: ChatRequest) -> dict[str, Any]:
|
||
request_id = _valid_chat_request_id(body.request_id or uuid.uuid4().hex)
|
||
conversation_id = _conversation_id(body.conversation_id)
|
||
selected = list(dict.fromkeys(_valid_name(name) for name in (body.models or ([body.model] if body.model else [])) if str(name).strip()))[:12]
|
||
if not selected:
|
||
raise HTTPException(400, "Select at least one loaded model")
|
||
_ensure_conversation(conversation_id, selected[0], selected, body.message or "New conversation")
|
||
attachment_meta = [{"name": item.name, "mime_type": item.mime_type, "url": item.url} for item in body.attachments[:12]]
|
||
_persist_message(conversation_id, request_id, "user", body.message.strip() or "[Attachments]", selected[0], attachment_meta)
|
||
_chat_state(request_id, state="preparing", stage="Preparing attachments", models=selected, conversation_id=conversation_id)
|
||
with _chat_requests_lock:
|
||
cancel_event = _chat_requests[request_id]["cancel"]
|
||
if len(selected) == 1:
|
||
payload = _chat_payload(body, selected[0])
|
||
try:
|
||
result = _stream_chat_request(payload, request_id, cancel_event=cancel_event)
|
||
except _ChatStopped as exc:
|
||
with _chat_requests_lock:
|
||
state = dict(_chat_requests.get(request_id, {}))
|
||
_persist_metric(conversation_id, request_id, selected[0], state, status="stopped")
|
||
raise HTTPException(499, "Chat stopped by user") from exc
|
||
except HTTPError as exc:
|
||
with _chat_requests_lock:
|
||
state = dict(_chat_requests.get(request_id, {}))
|
||
_persist_metric(conversation_id, request_id, selected[0], state, status="failed", error=str(exc))
|
||
raise _ollama_error(exc) from exc
|
||
except Exception as exc:
|
||
with _chat_requests_lock:
|
||
state = dict(_chat_requests.get(request_id, {}))
|
||
_persist_metric(conversation_id, request_id, selected[0], state, status="failed", error=str(exc))
|
||
raise
|
||
message = result.get("message") if isinstance(result.get("message"), dict) else {}
|
||
content = str(message.get("content") or "")
|
||
_persist_message(conversation_id, request_id, "assistant", content, selected[0])
|
||
with _chat_requests_lock:
|
||
state = dict(_chat_requests.get(request_id, {}))
|
||
persisted_metrics = _persist_metric(conversation_id, request_id, selected[0], state, status="completed")
|
||
return {"ok": True, "request_id": request_id, "conversation_id": conversation_id, "model": selected[0], "models": selected, "message": {"role": "assistant", "content": content}, "done": True, "metrics": persisted_metrics, "runtime": _runtime_snapshot()}
|
||
|
||
_chat_state(request_id, state="generating", stage=f"Querying {len(selected)} models in parallel")
|
||
results: dict[str, dict[str, Any]] = {}
|
||
errors: dict[str, str] = {}
|
||
|
||
def run_model(index: int, name: str):
|
||
child_id = f"{request_id}-{index}"
|
||
_chat_state(child_id, state="preparing", stage=f"Preparing {name}", model=name, parent_id=request_id, conversation_id=conversation_id)
|
||
payload = _chat_payload(body, name)
|
||
return name, _stream_chat_request(payload, child_id, cancel_event=cancel_event, parent_id=request_id)
|
||
|
||
try:
|
||
with ThreadPoolExecutor(max_workers=len(selected), thread_name_prefix="ollama-chat") as pool:
|
||
futures = [pool.submit(run_model, index, name) for index, name in enumerate(selected)]
|
||
for future in as_completed(futures):
|
||
try:
|
||
name, result = future.result()
|
||
results[name] = result
|
||
child_id = f"{request_id}-{selected.index(name)}"
|
||
with _chat_requests_lock:
|
||
child_state = dict(_chat_requests.get(child_id, {}))
|
||
_persist_metric(conversation_id, child_id, name, child_state, status="completed")
|
||
_chat_state(request_id, stage=f"Received response from {len(results)} of {len(selected)} models", response_chars=sum(len(str((r.get("message") or {}).get("content") or "")) for r in results.values()))
|
||
except _ChatStopped:
|
||
raise
|
||
except HTTPError as exc:
|
||
errors[str(exc)] = str(exc)
|
||
except Exception as exc:
|
||
errors[type(exc).__name__] = str(exc)
|
||
except _ChatStopped as exc:
|
||
raise HTTPException(499, "Chat stopped by user") from exc
|
||
|
||
if not results and errors:
|
||
raise HTTPException(502, "All selected Ollama models failed: " + "; ".join(errors.values()))
|
||
sections = []
|
||
for name in selected:
|
||
if name in results:
|
||
message = results[name].get("message") if isinstance(results[name].get("message"), dict) else {}
|
||
sections.append(f"[{name}]\n{str(message.get('content') or '').strip()}")
|
||
else:
|
||
sections.append(f"[{name}]\nModel failed: {errors.get(name, 'No response received')}")
|
||
combined = "\n\n".join(sections)
|
||
_chat_state(request_id, state="completed", stage="Combined model responses", finished_at=time.time(), response_chars=len(combined))
|
||
_persist_message(conversation_id, request_id, "assistant", combined, selected[0])
|
||
return {"ok": True, "request_id": request_id, "conversation_id": conversation_id, "model": selected[0], "models": selected, "message": {"role": "assistant", "content": combined}, "model_responses": {name: str((results.get(name, {}).get("message") or {}).get("content") or "") for name in selected if name in results}, "metrics": [results[name].get("metrics") for name in selected if name in results], "errors": errors, "done": True, "runtime": _runtime_snapshot()}
|
||
|
||
|
||
@router.get("/status")
|
||
def status() -> dict[str, Any]:
|
||
tags = _local_tags()
|
||
ps_rows = _local_ps()
|
||
loaded = {str(row.get("name") or row.get("model")): row for row in ps_rows}
|
||
local = [
|
||
_model_view(row, loaded.get(str(row.get("name") or row.get("model"))))
|
||
for row in tags
|
||
if not _is_mlx(row)
|
||
]
|
||
catalog = _ensure_catalog()
|
||
installed_names = {row["name"] for row in local}
|
||
|
||
def catalog_view(raw: dict[str, Any], source: str) -> dict[str, Any]:
|
||
name = str(raw.get("name") or raw.get("model"))
|
||
view = _model_view(raw, loaded.get(name), source=source)
|
||
view["installed"] = name in installed_names
|
||
view["loaded"] = name in loaded
|
||
return view
|
||
|
||
family_rows = catalog.get("families") if isinstance(catalog.get("families"), dict) else {}
|
||
catalog_rows = list(catalog.get("models", []))
|
||
catalog_names = {str(row.get("name") or row.get("model") or "") for row in catalog_rows}
|
||
candidate_rows = catalog_rows + [
|
||
variant
|
||
for rows in family_rows.values()
|
||
for variant in rows
|
||
if str(variant.get("name") or variant.get("model") or "") not in catalog_names
|
||
]
|
||
downloadable = []
|
||
seen_downloads: set[str] = set()
|
||
for rank, row in enumerate(candidate_rows, 1):
|
||
name = str(row.get("name") or row.get("model") or "")
|
||
if _is_mlx(row) or not name or name in installed_names or name in seen_downloads:
|
||
continue
|
||
view = catalog_view(row, "catalog")
|
||
if not _known_ram_fit(view):
|
||
continue
|
||
if name in catalog_names:
|
||
view["popularity_rank"] = rank
|
||
seen_downloads.add(name)
|
||
downloadable.append(view)
|
||
|
||
popular = _popular_fit_models(
|
||
[row for row in catalog.get("models", []) if not _is_mlx(row)],
|
||
family_rows,
|
||
catalog_view,
|
||
installed_names,
|
||
loaded,
|
||
)
|
||
for row in local:
|
||
family = _family_key(row["name"])
|
||
variants = []
|
||
for raw in family_rows.get(family, []):
|
||
if _is_mlx(raw):
|
||
continue
|
||
variant = _model_view(raw, loaded.get(raw["name"]), source="variant")
|
||
variant["installed"] = variant["name"] in installed_names
|
||
variant["current"] = variant["name"] == row["name"]
|
||
variants.append(variant)
|
||
row["variants"] = sorted(variants, key=lambda item: (item.get("size_bytes") or 0, item["name"]))
|
||
|
||
ollama_version = _ollama_version()
|
||
return {
|
||
"ollama": {
|
||
"available": bool(tags or ps_rows or ollama_version),
|
||
"version": ollama_version,
|
||
"endpoint": LOCAL_OLLAMA,
|
||
"containerized": _running_in_container(),
|
||
"configured_endpoint": bool(os.environ.get("OLLAMA_HOST", "").strip()),
|
||
"connection_hint": (
|
||
"Set OLLAMA_HOST to a reachable Ollama service, such as "
|
||
"http://ollama:11434 or http://host.docker.internal:11434."
|
||
if _running_in_container() and not (tags or ps_rows or ollama_version)
|
||
else ""
|
||
),
|
||
},
|
||
"models": local,
|
||
"popular": popular,
|
||
"popular_filter": {
|
||
"max_expected_ram_gib": _host_ram_gib(),
|
||
"basis": "detected MemTotal from the running host",
|
||
"requires_known_size": True,
|
||
"requires_known_ram": True,
|
||
"smaller_fit_variants_substituted": True,
|
||
},
|
||
"catalog": downloadable,
|
||
"catalog_filter": {
|
||
"max_expected_ram_gib": _host_ram_gib(),
|
||
"basis": "detected MemTotal from the running host",
|
||
"requires_known_size": True,
|
||
"requires_known_ram": True,
|
||
},
|
||
"catalog_filter_options": {
|
||
"types": ["all", "moe", "dense"],
|
||
"capabilities": sorted({cap for row in downloadable for cap in row.get("capabilities", [])}),
|
||
},
|
||
"catalog_source": catalog.get("source"),
|
||
"catalog_updated_at": catalog.get("fetched_at"),
|
||
"catalog_error": catalog.get("last_error"),
|
||
"next_catalog_refresh": _next_refresh(),
|
||
"jobs": _job_snapshot(),
|
||
"generated_at": time.time(),
|
||
}
|
||
|
||
|
||
def _ollama_version() -> str | None:
|
||
try:
|
||
return str(_json_request(LOCAL_OLLAMA + "/api/version", timeout=5).get("version") or "unknown")
|
||
except Exception:
|
||
return None
|
||
|
||
|
||
@router.post("/catalog/refresh")
|
||
def catalog_refresh() -> dict[str, Any]:
|
||
catalog = refresh_catalog()
|
||
return {"ok": bool(catalog.get("models")), "updated_at": catalog.get("fetched_at"), "count": len(catalog.get("models", [])), "error": catalog.get("last_error")}
|
||
|
||
|
||
@router.post("/pull")
|
||
def pull_model(body: ModelRequest) -> dict[str, Any]:
|
||
name = _valid_name(body.name)
|
||
return {"ok": True, "job_id": _new_job(name, "download"), "message": f"Downloading or updating {name}"}
|
||
|
||
|
||
@router.post("/redownload")
|
||
def redownload_model(body: ModelRequest) -> dict[str, Any]:
|
||
name = _valid_name(body.name)
|
||
return {"ok": True, "job_id": _new_job(name, "redownload"), "message": f"Re-downloading or updating {name}"}
|
||
|
||
|
||
@router.delete("/model")
|
||
def delete_model(body: ModelRequest) -> dict[str, Any]:
|
||
name = _valid_name(body.name)
|
||
return {"ok": True, "job_id": _new_job(name, "delete"), "message": f"Removing {name}"}
|
||
|
||
|
||
def create_ollama_routes(app) -> None:
|
||
app.include_router(router, prefix="/api/plugins/ollama-manager")
|