1"""Model Manager: load/unload/evict local inference workers with a strict memory policy."""23from __future__ import annotations45import asyncio6import contextlib7import json8import logging9import socket10import time11from dataclasses import dataclass, field12from pathlib import Path13from typing import Any1415import httpx16import psutil1718from .config import Settings19from .db import Database20from .errors import (Conflict, ModelIncompatible, ModelLoadError, ModelLoadTimeout, ModelNotFound, ModelTooLarge,21 RuntimeUnsupported, WorkerCrashed)22from .events import bus23from .hardware import GB, sample_telemetry_async24from .models.registry import Registry25from .runtimes import LlamaCppAdapter, MLXAdapter, RuntimeAdapter, WorkerHandle2627log = logging.getLogger("llm_api.manager")2829STATUS_UNLOADED = "unloaded"30STATUS_QUEUED = "queued"31STATUS_UNLOADING_PREVIOUS = "unloading_previous"32STATUS_LOADING = "loading"33STATUS_WARMING = "warming"34STATUS_READY = "ready"35STATUS_ERROR = "error"36STATUS_UNLOADING = "unloading"373839@dataclass40class LoadedModel:41 model: dict42 handle: WorkerHandle43 adapter: RuntimeAdapter44 status: str = STATUS_LOADING45 started_at: float = field(default_factory=time.time)46 ready_at: float | None = None47 last_used: float = field(default_factory=time.time)48 in_flight: int = 049 estimate_gb: float = 0.050 measured_gb: float = 0.051 warm: dict = field(default_factory=dict)52 error: str | None = None53 requests: int = 05455 def to_dict(self) -> dict:56 return {57 "model_id": self.model["id"], "name": self.model["name"], "runtime": self.handle.runtime,58 "status": self.status, "port": self.handle.port, "pid": self.handle.pid, "context": self.handle.context,59 "started_at": self.started_at, "ready_at": self.ready_at, "last_used": self.last_used,60 "in_flight": self.in_flight, "estimate_gb": round(self.estimate_gb, 2),61 "measured_gb": round(self.measured_gb, 2), "warm": self.warm, "error": self.error,62 "requests": self.requests, "pinned": bool(self.model.get("pinned")),63 "elapsed_seconds": round(time.time() - self.started_at, 1),64 "load_ms": self.handle.extra.get("load_ms"),65 }666768class ModelManager:69 def __init__(self, settings: Settings, db: Database, registry: Registry):70 self.settings = settings71 self.db = db72 self.registry = registry73 self.adapters: dict[str, RuntimeAdapter] = {74 "mlx": MLXAdapter(settings, settings.logs_path / "workers"),75 "llamacpp": LlamaCppAdapter(settings, settings.logs_path / "workers"),76 }77 self.loaded: dict[str, LoadedModel] = {}78 self.switch_lock = asyncio.Lock()79 self.progress: dict[str, dict] = {} # model_id -> transient status while switching80 self.waiting: dict[str, int] = {}81 self.client = httpx.AsyncClient(timeout=httpx.Timeout(30.0, read=None))82 self._tasks: list[asyncio.Task] = []83 self._workers_file = settings.data_path / "workers.json"84 self.stats = {"requests": 0, "tokens": 0, "loads": 0, "unloads": 0, "evictions": 0, "errors": 0}8586 # ------------------------------------------------------------------ lifecycle87 async def start(self) -> None:88 await self._recover_stale_workers()89 self._tasks.append(asyncio.create_task(self._monitor_loop(), name="worker-monitor"))90 self._tasks.append(asyncio.create_task(self._idle_loop(), name="idle-unload"))91 pre = (await self.db.get_setting("preload_model", self.settings.preload_model)) or "none"92 if pre and pre != "none":93 async def _pre():94 try:95 await self.ensure_loaded(pre, reason="preload")96 except Exception as e:97 log.warning("preload of %s failed: %s", pre, e)98 self._tasks.append(asyncio.create_task(_pre(), name="preload"))99100 async def stop(self) -> None:101 for t in self._tasks:102 t.cancel()103 for mid in list(self.loaded):104 with contextlib.suppress(Exception):105 await self.unload(mid, reason="shutdown", wait_inflight=True)106 await self.client.aclose()107108 async def _recover_stale_workers(self) -> None:109 """Kill workers left over from a previous (crashed) server and reset registry state."""110 if self._workers_file.exists():111 try:112 data = json.loads(self._workers_file.read_text())113 except Exception:114 data = {}115 for mid, w in data.items():116 pid = w.get("pid")117 if not pid:118 continue119 try:120 p = psutil.Process(pid)121 cmd = " ".join(p.cmdline())122 if "mlx_worker" in cmd or "llama-server" in cmd:123 log.warning("killing stale worker pid %s for %s", pid, mid)124 for c in p.children(recursive=True):125 with contextlib.suppress(psutil.Error):126 c.kill()127 p.kill()128 except psutil.Error:129 pass130 await self.db.model_event(None, "recovered_stale_workers", {"count": len(data)})131 self._persist_workers()132133 def _persist_workers(self) -> None:134 data = {mid: {"pid": lm.handle.pid, "port": lm.handle.port, "runtime": lm.handle.runtime}135 for mid, lm in self.loaded.items()}136 try:137 self._workers_file.parent.mkdir(parents=True, exist_ok=True)138 self._workers_file.write_text(json.dumps(data))139 except OSError:140 pass141142 # ------------------------------------------------------------------ queries143 def get_ready(self, model_id: str) -> LoadedModel | None:144 lm = self.loaded.get(model_id)145 return lm if lm and lm.status == STATUS_READY else None146147 def status_of(self, model_id: str) -> str:148 lm = self.loaded.get(model_id)149 if lm:150 return lm.status151 p = self.progress.get(model_id)152 return p["status"] if p else STATUS_UNLOADED153154 def current_model(self) -> LoadedModel | None:155 # Largest ready text model, else any ready156 ready = [lm for lm in self.loaded.values() if lm.status == STATUS_READY]157 if not ready:158 return None159 text = [lm for lm in ready if not (lm.model.get("embedding") or lm.model.get("reranker"))]160 pool = text or ready161 return max(pool, key=lambda lm: lm.estimate_gb)162163 def snapshot(self) -> dict:164 return {165 "loaded": [lm.to_dict() for lm in self.loaded.values()],166 "progress": self.progress,167 "waiting": self.waiting,168 "switching": self.switch_lock.locked(),169 "stats": self.stats,170 "resident_gb": round(self.resident_gb(), 2),171 "runtimes": {k: v.available() for k, v in self.adapters.items()},172 }173174 def resident_gb(self) -> float:175 return sum(max(lm.estimate_gb, lm.measured_gb) for lm in self.loaded.values())176177 async def budgets(self) -> tuple[float, float, int]:178 b = float(await self.db.get_setting("max_model_memory_gb", self.settings.max_model_memory_gb))179 a = float(await self.db.get_setting("absolute_max_memory_gb", self.settings.absolute_max_memory_gb))180 n = int(await self.db.get_setting("max_simultaneous_models", self.settings.max_simultaneous_models))181 return b, a, n182183 # ------------------------------------------------------------------ public ops184 async def ensure_loaded(self, name: str, *, reason: str = "request", context: int | None = None,185 force: bool = False) -> LoadedModel:186 model = await self.registry.resolve(name)187 if not model:188 raise ModelNotFound(f"The model '{name}' does not exist. Use GET /v1/models to list available models.")189 mid = model["id"]190 lm = self.get_ready(mid)191 if lm:192 lm.last_used = time.time()193 return lm194 self.waiting[mid] = self.waiting.get(mid, 0) + 1195 self._set_progress(mid, STATUS_QUEUED)196 try:197 async with self.switch_lock:198 lm = self.get_ready(mid)199 if lm:200 return lm201 return await self._load_locked(model, reason=reason, context=context, force=force)202 finally:203 self.waiting[mid] = max(0, self.waiting.get(mid, 1) - 1)204 if not self.waiting[mid]:205 self.waiting.pop(mid, None)206207 async def load(self, name: str, *, context: int | None = None, force: bool = False) -> LoadedModel:208 return await self.ensure_loaded(name, reason="manual", context=context, force=force)209210 async def unload(self, model_id: str, *, reason: str = "manual", wait_inflight: bool = True) -> bool:211 lm = self.loaded.get(model_id)212 if not lm:213 return False214 lm.status = STATUS_UNLOADING215 self._publish_state()216 if wait_inflight:217 t0 = time.time()218 while lm.in_flight > 0 and time.time() - t0 < 120:219 await asyncio.sleep(0.2)220 before = psutil.virtual_memory().available221 await lm.adapter.stop(lm.handle, self.client)222 self.loaded.pop(model_id, None)223 self.progress.pop(model_id, None)224 self._persist_workers()225 self.stats["unloads"] += 1226 released = await self._wait_memory_release(before, lm.measured_gb or lm.estimate_gb)227 await self.registry.update(model_id) # touch updated_at228 await self.db.model_event(model_id, "unloaded", {"reason": reason, "released_gb": released})229 bus.publish("model", {"model_id": model_id, "status": STATUS_UNLOADED, "reason": reason, "released_gb": released})230 self._publish_state()231 log.info("unloaded %s (%s), released ~%.1f GB", model_id, reason, released)232 return True233234 @contextlib.asynccontextmanager235 async def use(self, lm: LoadedModel):236 lm.in_flight += 1237 lm.requests += 1238 self.stats["requests"] += 1239 try:240 yield lm241 finally:242 lm.in_flight = max(0, lm.in_flight - 1)243 lm.last_used = time.time()244245 # ------------------------------------------------------------------ internals246 def _set_progress(self, mid: str, status: str, **extra: Any) -> None:247 p = self.progress.get(mid) or {"started": time.time()}248 p.update({"status": status, "elapsed_seconds": round(time.time() - p["started"], 1), **extra})249 self.progress[mid] = p250 bus.publish("model", {"model_id": mid, **p})251252 def _publish_state(self) -> None:253 bus.publish("manager", self.snapshot())254255 def _free_port(self) -> int:256 used = {lm.handle.port for lm in self.loaded.values()}257 for port in range(self.settings.worker_port_start, self.settings.worker_port_end + 1):258 if port in used:259 continue260 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:261 try:262 s.bind(("127.0.0.1", port))263 return port264 except OSError:265 continue266 raise ModelLoadError("No free worker port available.")267268 async def _load_locked(self, model: dict, *, reason: str, context: int | None, force: bool) -> LoadedModel:269 mid = model["id"]270 if not model.get("installed"):271 raise ModelNotFound(f"Model '{mid}' files are missing on disk.")272 if not model.get("enabled"):273 raise ModelIncompatible(f"Model '{mid}' is disabled.")274 rt = model["runtime"]275 adapter = self.adapters.get(rt)276 if not adapter or not adapter.available():277 raise RuntimeUnsupported(f"Runtime '{rt}' is not available on this machine.")278 if not model.get("compatible") and not force:279 raise ModelIncompatible(280 f"Model '{mid}' is marked {model.get('compatibility_status')}: {model.get('compatibility_reason')}",281 extra={"compatibility_status": model.get("compatibility_status")})282283 budget, absolute, max_models = await self.budgets()284 ctx = int(context or (model.get("overrides") or {}).get("context") or model.get("recommended_context")285 or min(self.settings.default_context, model.get("max_context") or self.settings.default_context))286 if model.get("max_context"):287 ctx = min(ctx, int(model["max_context"]))288 est = adapter.estimate_memory(model, ctx)289 limit = absolute if force else budget290 if est.total_gb > limit:291 raise ModelTooLarge(292 f"Model '{mid}' requires approximately {est.total_gb:.1f} GB at a {ctx} context but the safe limit is "293 f"{limit:.0f} GB.", extra={"estimate": est.to_dict(), "limit_gb": limit})294295 # ---- eviction ---------------------------------------------------296 self._set_progress(mid, STATUS_QUEUED, context=ctx, estimate_gb=round(est.total_gb, 2))297 small = bool(model.get("embedding") or model.get("reranker")) and est.total_gb <= self.settings.small_model_resident_gb298 others = [lm for lm in self.loaded.values() if lm.model["id"] != mid]299 big_others = [lm for lm in others if not (300 (lm.model.get("embedding") or lm.model.get("reranker")) and lm.estimate_gb <= self.settings.small_model_resident_gb)]301 to_evict: list[LoadedModel] = []302 # count policy: at most `max_models` large models303 if not small:304 while len(big_others) - len(to_evict) >= max_models:305 victim = self._pick_victim([lm for lm in big_others if lm not in to_evict])306 if not victim:307 break308 to_evict.append(victim)309 # memory policy: resident + new <= budget310 def resident_after() -> float:311 return sum(max(lm.estimate_gb, lm.measured_gb) for lm in others if lm not in to_evict)312 while resident_after() + est.total_gb > limit:313 victim = self._pick_victim([lm for lm in others if lm not in to_evict])314 if not victim:315 break316 to_evict.append(victim)317 if resident_after() + est.total_gb > limit:318 raise ModelTooLarge(f"Not enough memory budget for '{mid}' ({est.total_gb:.1f} GB) alongside resident models.")319 if to_evict:320 self._set_progress(mid, STATUS_UNLOADING_PREVIOUS, evicting=[lm.model["id"] for lm in to_evict])321 for lm in to_evict:322 self.stats["evictions"] += 1323 await self.unload(lm.model["id"], reason=f"evicted for {mid}")324325 # ---- real free-memory check -------------------------------------326 tel = await sample_telemetry_async(self.settings.models_dir)327 needed = est.total_gb + 1.5328 if tel.mem_available_gb < needed:329 # give the OS a moment to reclaim330 for _ in range(20):331 await asyncio.sleep(0.5)332 tel = await sample_telemetry_async(self.settings.models_dir)333 if tel.mem_available_gb >= needed:334 break335 if tel.mem_available_gb < needed and not force:336 raise ModelTooLarge(337 f"Only {tel.mem_available_gb:.1f} GB of memory is available but '{mid}' needs about {est.total_gb:.1f} GB. "338 f"Memory pressure: {tel.mem_pressure_level}.", extra={"estimate": est.to_dict(), "available_gb": tel.mem_available_gb})339340 # ---- spawn --------------------------------------------------------341 port = self._free_port()342 used_before = psutil.virtual_memory().total - psutil.virtual_memory().available343 t0 = time.time()344 self._set_progress(mid, STATUS_LOADING, port=port)345 try:346 handle = adapter.spawn(model, port, ctx)347 except Exception as e:348 self.progress.pop(mid, None)349 raise ModelLoadError(f"Failed to start worker for '{mid}': {e}")350 lm = LoadedModel(model=model, handle=handle, adapter=adapter, estimate_gb=est.total_gb)351 self.loaded[mid] = lm352 self._persist_workers()353 await self.db.model_event(mid, "loading", {"reason": reason, "context": ctx, "port": port, "estimate": est.to_dict()})354 try:355 await self._wait_ready(lm)356 load_ms = (time.time() - t0) * 1000357 handle.extra["load_ms"] = round(load_ms)358 lm.status = STATUS_WARMING359 self._set_progress(mid, STATUS_WARMING, load_ms=round(load_ms))360 warm = await asyncio.wait_for(adapter.warmup(handle, self.client, model), timeout=self.settings.load_timeout_seconds)361 lm.warm = warm362 lm.measured_gb = await self._measure(lm, used_before)363 lm.status = STATUS_READY364 lm.ready_at = time.time()365 self.progress.pop(mid, None)366 self.stats["loads"] += 1367 await self.registry.record_load(mid, load_ms)368 if warm.get("ttft_ms"):369 await self.registry.update(mid, first_token_latency_ms=warm["ttft_ms"])370 await self.registry.update(mid, verified=True)371 await self.db.model_event(mid, "ready", {"load_ms": round(load_ms), "warm": warm, "measured_gb": round(lm.measured_gb, 2)})372 bus.publish("model", {"model_id": mid, "status": STATUS_READY, "load_ms": round(load_ms), "warm": warm})373 self._publish_state()374 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)375 return lm376 except Exception as e:377 self.stats["errors"] += 1378 err = str(e)379 tail = self._log_tail(handle.log_path)380 lm.status = STATUS_ERROR381 lm.error = err382 await adapter.stop(handle, None)383 self.loaded.pop(mid, None)384 self.progress.pop(mid, None)385 self._persist_workers()386 await self.db.model_event(mid, "load_failed", {"error": err, "log": tail})387 bus.publish("model", {"model_id": mid, "status": STATUS_ERROR, "error": err})388 self._publish_state()389 if isinstance(e, asyncio.TimeoutError):390 raise ModelLoadTimeout(f"Model '{mid}' did not become ready within {self.settings.load_timeout_seconds}s.",391 extra={"log": tail})392 if isinstance(e, (ModelLoadError, WorkerCrashed)):393 e.extra.setdefault("log", tail)394 raise395 raise ModelLoadError(f"Model '{mid}' failed to load: {err}", extra={"log": tail})396397 async def _measure(self, lm: LoadedModel, used_before: int | None = None) -> float:398 """Best available estimate of the worker's real memory: process footprint, worker-reported Metal memory,399 or the system-wide used-memory delta since spawn (Metal buffers of llama-server are not in RSS)."""400 vals = [lm.handle.memory_bytes() / GB]401 rep = getattr(lm.adapter, "memory_gb", None)402 if rep is not None:403 v = await rep(lm.handle, self.client)404 if v:405 vals.append(v)406 if used_before is not None:407 vm = psutil.virtual_memory()408 vals.append(max(0.0, ((vm.total - vm.available) - used_before) / GB))409 return round(max(vals), 2)410411 def _pick_victim(self, candidates: list[LoadedModel]) -> LoadedModel | None:412 if not candidates:413 return None414 unpinned = [lm for lm in candidates if not lm.model.get("pinned")]415 pool = unpinned or candidates # pinned models are evicted only when nothing else can be416 return min(pool, key=lambda lm: lm.last_used)417418 async def _wait_ready(self, lm: LoadedModel) -> None:419 deadline = time.time() + self.settings.load_timeout_seconds420 while time.time() < deadline:421 if not lm.handle.alive():422 raise WorkerCrashed(f"Worker for '{lm.model['id']}' exited during load (code {lm.handle.process.returncode}).")423 status, err = await lm.adapter.is_ready(lm.handle, self.client)424 if status == "ready":425 return426 if status == "error":427 raise ModelLoadError(f"Model '{lm.model['id']}' failed to load: {err}")428 lm.measured_gb = lm.handle.memory_bytes() / GB429 self._set_progress(lm.model["id"], STATUS_LOADING, measured_gb=round(lm.measured_gb, 2))430 await asyncio.sleep(0.5)431 raise asyncio.TimeoutError()432433 async def _wait_memory_release(self, before_available: int, expected_gb: float) -> float:434 """Wait until the OS reports the freed memory (up to ~8 s). Returns GB released."""435 best = 0.0436 for _ in range(16):437 await asyncio.sleep(0.5)438 now = psutil.virtual_memory().available439 best = max(best, (now - before_available) / GB)440 if expected_gb and best >= expected_gb * 0.7:441 break442 return round(best, 2)443444 @staticmethod445 def _log_tail(path: Path, n: int = 40) -> str:446 try:447 lines = path.read_text(errors="replace").splitlines()448 return "\n".join(lines[-n:])449 except Exception:450 return ""451452 # ------------------------------------------------------------------ loops453 async def _monitor_loop(self) -> None:454 while True:455 try:456 await asyncio.sleep(5)457 for mid, lm in list(self.loaded.items()):458 if lm.status in (STATUS_LOADING, STATUS_WARMING, STATUS_UNLOADING):459 continue460 if not lm.handle.alive():461 log.error("worker for %s died (code %s)", mid, lm.handle.process.returncode)462 self.loaded.pop(mid, None)463 self._persist_workers()464 self.stats["errors"] += 1465 await self.db.model_event(mid, "worker_crashed", {"code": lm.handle.process.returncode,466 "log": self._log_tail(lm.handle.log_path)})467 bus.publish("model", {"model_id": mid, "status": STATUS_ERROR, "error": "worker crashed"})468 self._publish_state()469 continue470 lm.measured_gb = max(lm.measured_gb, await self._measure(lm))471 status, err = await lm.adapter.is_ready(lm.handle, self.client)472 if status == "error":473 log.error("worker for %s reports error: %s", mid, err)474 await self.unload(mid, reason=f"worker error: {err}", wait_inflight=False)475 except asyncio.CancelledError:476 return477 except Exception:478 log.exception("monitor loop error")479480 async def _idle_loop(self) -> None:481 while True:482 try:483 await asyncio.sleep(30)484 minutes = int(await self.db.get_setting("model_idle_timeout_minutes", self.settings.model_idle_timeout_minutes))485 if minutes <= 0:486 continue487 now = time.time()488 for mid, lm in list(self.loaded.items()):489 if lm.status != STATUS_READY or lm.model.get("pinned") or lm.in_flight:490 continue491 if now - lm.last_used > minutes * 60:492 log.info("idle timeout: unloading %s", mid)493 await self.unload(mid, reason="idle timeout")494 except asyncio.CancelledError:495 return496 except Exception:497 log.exception("idle loop error")498499 # ------------------------------------------------------------------ pin / refresh500 async def refresh_model(self, model_id: str) -> None:501 """Re-read registry row for a loaded model (pin/favorite changes)."""502 lm = self.loaded.get(model_id)503 if lm:504 m = await self.registry.get(model_id)505 if m:506 lm.model = m507