"""Model Manager: load/unload/evict local inference workers with a strict memory policy.""" from __future__ import annotations import asyncio import contextlib import json import logging import socket import time from dataclasses import dataclass, field from pathlib import Path from typing import Any import httpx import psutil from .config import Settings from .db import Database from .errors import (Conflict, ModelIncompatible, ModelLoadError, ModelLoadTimeout, ModelNotFound, ModelTooLarge, RuntimeUnsupported, WorkerCrashed) from .events import bus from .hardware import GB, sample_telemetry_async from .models.registry import Registry from .runtimes import LlamaCppAdapter, MLXAdapter, RuntimeAdapter, WorkerHandle log = logging.getLogger("llm_api.manager") STATUS_UNLOADED = "unloaded" STATUS_QUEUED = "queued" STATUS_UNLOADING_PREVIOUS = "unloading_previous" STATUS_LOADING = "loading" STATUS_WARMING = "warming" STATUS_READY = "ready" STATUS_ERROR = "error" STATUS_UNLOADING = "unloading" @dataclass class LoadedModel: model: dict handle: WorkerHandle adapter: RuntimeAdapter status: str = STATUS_LOADING started_at: float = field(default_factory=time.time) ready_at: float | None = None last_used: float = field(default_factory=time.time) in_flight: int = 0 estimate_gb: float = 0.0 measured_gb: float = 0.0 warm: dict = field(default_factory=dict) error: str | None = None requests: int = 0 def to_dict(self) -> dict: return { "model_id": self.model["id"], "name": self.model["name"], "runtime": self.handle.runtime, "status": self.status, "port": self.handle.port, "pid": self.handle.pid, "context": self.handle.context, "started_at": self.started_at, "ready_at": self.ready_at, "last_used": self.last_used, "in_flight": self.in_flight, "estimate_gb": round(self.estimate_gb, 2), "measured_gb": round(self.measured_gb, 2), "warm": self.warm, "error": self.error, "requests": self.requests, "pinned": bool(self.model.get("pinned")), "elapsed_seconds": round(time.time() - self.started_at, 1), "load_ms": self.handle.extra.get("load_ms"), } class ModelManager: def __init__(self, settings: Settings, db: Database, registry: Registry): self.settings = settings self.db = db self.registry = registry self.adapters: dict[str, RuntimeAdapter] = { "mlx": MLXAdapter(settings, settings.logs_path / "workers"), "llamacpp": LlamaCppAdapter(settings, settings.logs_path / "workers"), } self.loaded: dict[str, LoadedModel] = {} self.switch_lock = asyncio.Lock() self.progress: dict[str, dict] = {} # model_id -> transient status while switching self.waiting: dict[str, int] = {} self.client = httpx.AsyncClient(timeout=httpx.Timeout(30.0, read=None)) self._tasks: list[asyncio.Task] = [] self._workers_file = settings.data_path / "workers.json" self.stats = {"requests": 0, "tokens": 0, "loads": 0, "unloads": 0, "evictions": 0, "errors": 0} # ------------------------------------------------------------------ lifecycle async def start(self) -> None: await self._recover_stale_workers() self._tasks.append(asyncio.create_task(self._monitor_loop(), name="worker-monitor")) self._tasks.append(asyncio.create_task(self._idle_loop(), name="idle-unload")) pre = (await self.db.get_setting("preload_model", self.settings.preload_model)) or "none" if pre and pre != "none": async def _pre(): try: await self.ensure_loaded(pre, reason="preload") except Exception as e: log.warning("preload of %s failed: %s", pre, e) self._tasks.append(asyncio.create_task(_pre(), name="preload")) async def stop(self) -> None: for t in self._tasks: t.cancel() for mid in list(self.loaded): with contextlib.suppress(Exception): await self.unload(mid, reason="shutdown", wait_inflight=True) await self.client.aclose() async def _recover_stale_workers(self) -> None: """Kill workers left over from a previous (crashed) server and reset registry state.""" if self._workers_file.exists(): try: data = json.loads(self._workers_file.read_text()) except Exception: data = {} for mid, w in data.items(): pid = w.get("pid") if not pid: continue try: p = psutil.Process(pid) cmd = " ".join(p.cmdline()) if "mlx_worker" in cmd or "llama-server" in cmd: log.warning("killing stale worker pid %s for %s", pid, mid) for c in p.children(recursive=True): with contextlib.suppress(psutil.Error): c.kill() p.kill() except psutil.Error: pass await self.db.model_event(None, "recovered_stale_workers", {"count": len(data)}) self._persist_workers() def _persist_workers(self) -> None: data = {mid: {"pid": lm.handle.pid, "port": lm.handle.port, "runtime": lm.handle.runtime} for mid, lm in self.loaded.items()} try: self._workers_file.parent.mkdir(parents=True, exist_ok=True) self._workers_file.write_text(json.dumps(data)) except OSError: pass # ------------------------------------------------------------------ queries def get_ready(self, model_id: str) -> LoadedModel | None: lm = self.loaded.get(model_id) return lm if lm and lm.status == STATUS_READY else None def status_of(self, model_id: str) -> str: lm = self.loaded.get(model_id) if lm: return lm.status p = self.progress.get(model_id) return p["status"] if p else STATUS_UNLOADED def current_model(self) -> LoadedModel | None: # Largest ready text model, else any ready ready = [lm for lm in self.loaded.values() if lm.status == STATUS_READY] if not ready: return None text = [lm for lm in ready if not (lm.model.get("embedding") or lm.model.get("reranker"))] pool = text or ready return max(pool, key=lambda lm: lm.estimate_gb) def snapshot(self) -> dict: return { "loaded": [lm.to_dict() for lm in self.loaded.values()], "progress": self.progress, "waiting": self.waiting, "switching": self.switch_lock.locked(), "stats": self.stats, "resident_gb": round(self.resident_gb(), 2), "runtimes": {k: v.available() for k, v in self.adapters.items()}, } def resident_gb(self) -> float: return sum(max(lm.estimate_gb, lm.measured_gb) for lm in self.loaded.values()) async def budgets(self) -> tuple[float, float, int]: b = float(await self.db.get_setting("max_model_memory_gb", self.settings.max_model_memory_gb)) a = float(await self.db.get_setting("absolute_max_memory_gb", self.settings.absolute_max_memory_gb)) n = int(await self.db.get_setting("max_simultaneous_models", self.settings.max_simultaneous_models)) return b, a, n # ------------------------------------------------------------------ public ops async def ensure_loaded(self, name: str, *, reason: str = "request", context: int | None = None, force: bool = False) -> LoadedModel: model = await self.registry.resolve(name) if not model: raise ModelNotFound(f"The model '{name}' does not exist. Use GET /v1/models to list available models.") mid = model["id"] lm = self.get_ready(mid) if lm: lm.last_used = time.time() return lm self.waiting[mid] = self.waiting.get(mid, 0) + 1 self._set_progress(mid, STATUS_QUEUED) try: async with self.switch_lock: lm = self.get_ready(mid) if lm: return lm return await self._load_locked(model, reason=reason, context=context, force=force) finally: self.waiting[mid] = max(0, self.waiting.get(mid, 1) - 1) if not self.waiting[mid]: self.waiting.pop(mid, None) async def load(self, name: str, *, context: int | None = None, force: bool = False) -> LoadedModel: return await self.ensure_loaded(name, reason="manual", context=context, force=force) async def unload(self, model_id: str, *, reason: str = "manual", wait_inflight: bool = True) -> bool: lm = self.loaded.get(model_id) if not lm: return False lm.status = STATUS_UNLOADING self._publish_state() if wait_inflight: t0 = time.time() while lm.in_flight > 0 and time.time() - t0 < 120: await asyncio.sleep(0.2) before = psutil.virtual_memory().available await lm.adapter.stop(lm.handle, self.client) self.loaded.pop(model_id, None) self.progress.pop(model_id, None) self._persist_workers() self.stats["unloads"] += 1 released = await self._wait_memory_release(before, lm.measured_gb or lm.estimate_gb) await self.registry.update(model_id) # touch updated_at await self.db.model_event(model_id, "unloaded", {"reason": reason, "released_gb": released}) bus.publish("model", {"model_id": model_id, "status": STATUS_UNLOADED, "reason": reason, "released_gb": released}) self._publish_state() log.info("unloaded %s (%s), released ~%.1f GB", model_id, reason, released) return True @contextlib.asynccontextmanager async def use(self, lm: LoadedModel): lm.in_flight += 1 lm.requests += 1 self.stats["requests"] += 1 try: yield lm finally: lm.in_flight = max(0, lm.in_flight - 1) lm.last_used = time.time() # ------------------------------------------------------------------ internals def _set_progress(self, mid: str, status: str, **extra: Any) -> None: p = self.progress.get(mid) or {"started": time.time()} p.update({"status": status, "elapsed_seconds": round(time.time() - p["started"], 1), **extra}) self.progress[mid] = p bus.publish("model", {"model_id": mid, **p}) def _publish_state(self) -> None: bus.publish("manager", self.snapshot()) def _free_port(self) -> int: used = {lm.handle.port for lm in self.loaded.values()} for port in range(self.settings.worker_port_start, self.settings.worker_port_end + 1): if port in used: continue with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: try: s.bind(("127.0.0.1", port)) return port except OSError: continue raise ModelLoadError("No free worker port available.") async def _load_locked(self, model: dict, *, reason: str, context: int | None, force: bool) -> LoadedModel: mid = model["id"] if not model.get("installed"): raise ModelNotFound(f"Model '{mid}' files are missing on disk.") if not model.get("enabled"): raise ModelIncompatible(f"Model '{mid}' is disabled.") rt = model["runtime"] adapter = self.adapters.get(rt) if not adapter or not adapter.available(): raise RuntimeUnsupported(f"Runtime '{rt}' is not available on this machine.") if not model.get("compatible") and not force: raise ModelIncompatible( f"Model '{mid}' is marked {model.get('compatibility_status')}: {model.get('compatibility_reason')}", extra={"compatibility_status": model.get("compatibility_status")}) budget, absolute, max_models = await self.budgets() ctx = int(context or (model.get("overrides") or {}).get("context") or model.get("recommended_context") or min(self.settings.default_context, model.get("max_context") or self.settings.default_context)) if model.get("max_context"): ctx = min(ctx, int(model["max_context"])) est = adapter.estimate_memory(model, ctx) limit = absolute if force else budget if est.total_gb > limit: raise ModelTooLarge( f"Model '{mid}' requires approximately {est.total_gb:.1f} GB at a {ctx} context but the safe limit is " f"{limit:.0f} GB.", extra={"estimate": est.to_dict(), "limit_gb": limit}) # ---- eviction --------------------------------------------------- self._set_progress(mid, STATUS_QUEUED, context=ctx, estimate_gb=round(est.total_gb, 2)) small = bool(model.get("embedding") or model.get("reranker")) and est.total_gb <= self.settings.small_model_resident_gb others = [lm for lm in self.loaded.values() if lm.model["id"] != mid] big_others = [lm for lm in others if not ( (lm.model.get("embedding") or lm.model.get("reranker")) and lm.estimate_gb <= self.settings.small_model_resident_gb)] to_evict: list[LoadedModel] = [] # count policy: at most `max_models` large models if not small: while len(big_others) - len(to_evict) >= max_models: victim = self._pick_victim([lm for lm in big_others if lm not in to_evict]) if not victim: break to_evict.append(victim) # memory policy: resident + new <= budget def resident_after() -> float: return sum(max(lm.estimate_gb, lm.measured_gb) for lm in others if lm not in to_evict) while resident_after() + est.total_gb > limit: victim = self._pick_victim([lm for lm in others if lm not in to_evict]) if not victim: break to_evict.append(victim) if resident_after() + est.total_gb > limit: raise ModelTooLarge(f"Not enough memory budget for '{mid}' ({est.total_gb:.1f} GB) alongside resident models.") if to_evict: self._set_progress(mid, STATUS_UNLOADING_PREVIOUS, evicting=[lm.model["id"] for lm in to_evict]) for lm in to_evict: self.stats["evictions"] += 1 await self.unload(lm.model["id"], reason=f"evicted for {mid}") # ---- real free-memory check ------------------------------------- tel = await sample_telemetry_async(self.settings.models_dir) needed = est.total_gb + 1.5 if tel.mem_available_gb < needed: # give the OS a moment to reclaim for _ in range(20): await asyncio.sleep(0.5) tel = await sample_telemetry_async(self.settings.models_dir) if tel.mem_available_gb >= needed: break if tel.mem_available_gb < needed and not force: raise ModelTooLarge( f"Only {tel.mem_available_gb:.1f} GB of memory is available but '{mid}' needs about {est.total_gb:.1f} GB. " f"Memory pressure: {tel.mem_pressure_level}.", extra={"estimate": est.to_dict(), "available_gb": tel.mem_available_gb}) # ---- spawn -------------------------------------------------------- port = self._free_port() used_before = psutil.virtual_memory().total - psutil.virtual_memory().available t0 = time.time() self._set_progress(mid, STATUS_LOADING, port=port) try: handle = adapter.spawn(model, port, ctx) except Exception as e: self.progress.pop(mid, None) raise ModelLoadError(f"Failed to start worker for '{mid}': {e}") lm = LoadedModel(model=model, handle=handle, adapter=adapter, estimate_gb=est.total_gb) self.loaded[mid] = lm self._persist_workers() await self.db.model_event(mid, "loading", {"reason": reason, "context": ctx, "port": port, "estimate": est.to_dict()}) try: await self._wait_ready(lm) load_ms = (time.time() - t0) * 1000 handle.extra["load_ms"] = round(load_ms) lm.status = STATUS_WARMING self._set_progress(mid, STATUS_WARMING, load_ms=round(load_ms)) warm = await asyncio.wait_for(adapter.warmup(handle, self.client, model), timeout=self.settings.load_timeout_seconds) lm.warm = warm lm.measured_gb = await self._measure(lm, used_before) lm.status = STATUS_READY lm.ready_at = time.time() self.progress.pop(mid, None) self.stats["loads"] += 1 await self.registry.record_load(mid, load_ms) if warm.get("ttft_ms"): await self.registry.update(mid, first_token_latency_ms=warm["ttft_ms"]) await self.registry.update(mid, verified=True) await self.db.model_event(mid, "ready", {"load_ms": round(load_ms), "warm": warm, "measured_gb": round(lm.measured_gb, 2)}) bus.publish("model", {"model_id": mid, "status": STATUS_READY, "load_ms": round(load_ms), "warm": warm}) self._publish_state() log.info("ready %s in %.0f ms (ctx %d, est %.1f GB, measured %.1f GB)", mid, load_ms, ctx, est.total_gb, lm.measured_gb) return lm except Exception as e: self.stats["errors"] += 1 err = str(e) tail = self._log_tail(handle.log_path) lm.status = STATUS_ERROR lm.error = err await adapter.stop(handle, None) self.loaded.pop(mid, None) self.progress.pop(mid, None) self._persist_workers() await self.db.model_event(mid, "load_failed", {"error": err, "log": tail}) bus.publish("model", {"model_id": mid, "status": STATUS_ERROR, "error": err}) self._publish_state() if isinstance(e, asyncio.TimeoutError): raise ModelLoadTimeout(f"Model '{mid}' did not become ready within {self.settings.load_timeout_seconds}s.", extra={"log": tail}) if isinstance(e, (ModelLoadError, WorkerCrashed)): e.extra.setdefault("log", tail) raise raise ModelLoadError(f"Model '{mid}' failed to load: {err}", extra={"log": tail}) async def _measure(self, lm: LoadedModel, used_before: int | None = None) -> float: """Best available estimate of the worker's real memory: process footprint, worker-reported Metal memory, or the system-wide used-memory delta since spawn (Metal buffers of llama-server are not in RSS).""" vals = [lm.handle.memory_bytes() / GB] rep = getattr(lm.adapter, "memory_gb", None) if rep is not None: v = await rep(lm.handle, self.client) if v: vals.append(v) if used_before is not None: vm = psutil.virtual_memory() vals.append(max(0.0, ((vm.total - vm.available) - used_before) / GB)) return round(max(vals), 2) def _pick_victim(self, candidates: list[LoadedModel]) -> LoadedModel | None: if not candidates: return None unpinned = [lm for lm in candidates if not lm.model.get("pinned")] pool = unpinned or candidates # pinned models are evicted only when nothing else can be return min(pool, key=lambda lm: lm.last_used) async def _wait_ready(self, lm: LoadedModel) -> None: deadline = time.time() + self.settings.load_timeout_seconds while time.time() < deadline: if not lm.handle.alive(): raise WorkerCrashed(f"Worker for '{lm.model['id']}' exited during load (code {lm.handle.process.returncode}).") status, err = await lm.adapter.is_ready(lm.handle, self.client) if status == "ready": return if status == "error": raise ModelLoadError(f"Model '{lm.model['id']}' failed to load: {err}") lm.measured_gb = lm.handle.memory_bytes() / GB self._set_progress(lm.model["id"], STATUS_LOADING, measured_gb=round(lm.measured_gb, 2)) await asyncio.sleep(0.5) raise asyncio.TimeoutError() async def _wait_memory_release(self, before_available: int, expected_gb: float) -> float: """Wait until the OS reports the freed memory (up to ~8 s). Returns GB released.""" best = 0.0 for _ in range(16): await asyncio.sleep(0.5) now = psutil.virtual_memory().available best = max(best, (now - before_available) / GB) if expected_gb and best >= expected_gb * 0.7: break return round(best, 2) @staticmethod def _log_tail(path: Path, n: int = 40) -> str: try: lines = path.read_text(errors="replace").splitlines() return "\n".join(lines[-n:]) except Exception: return "" # ------------------------------------------------------------------ loops async def _monitor_loop(self) -> None: while True: try: await asyncio.sleep(5) for mid, lm in list(self.loaded.items()): if lm.status in (STATUS_LOADING, STATUS_WARMING, STATUS_UNLOADING): continue if not lm.handle.alive(): log.error("worker for %s died (code %s)", mid, lm.handle.process.returncode) self.loaded.pop(mid, None) self._persist_workers() self.stats["errors"] += 1 await self.db.model_event(mid, "worker_crashed", {"code": lm.handle.process.returncode, "log": self._log_tail(lm.handle.log_path)}) bus.publish("model", {"model_id": mid, "status": STATUS_ERROR, "error": "worker crashed"}) self._publish_state() continue lm.measured_gb = max(lm.measured_gb, await self._measure(lm)) status, err = await lm.adapter.is_ready(lm.handle, self.client) if status == "error": log.error("worker for %s reports error: %s", mid, err) await self.unload(mid, reason=f"worker error: {err}", wait_inflight=False) except asyncio.CancelledError: return except Exception: log.exception("monitor loop error") async def _idle_loop(self) -> None: while True: try: await asyncio.sleep(30) minutes = int(await self.db.get_setting("model_idle_timeout_minutes", self.settings.model_idle_timeout_minutes)) if minutes <= 0: continue now = time.time() for mid, lm in list(self.loaded.items()): if lm.status != STATUS_READY or lm.model.get("pinned") or lm.in_flight: continue if now - lm.last_used > minutes * 60: log.info("idle timeout: unloading %s", mid) await self.unload(mid, reason="idle timeout") except asyncio.CancelledError: return except Exception: log.exception("idle loop error") # ------------------------------------------------------------------ pin / refresh async def refresh_model(self, model_id: str) -> None: """Re-read registry row for a loaded model (pin/favorite changes).""" lm = self.loaded.get(model_id) if lm: m = await self.registry.get(model_id) if m: lm.model = m