SPB Git forge

spb/llm-api

Public
0commits 0branches 0releases
0 Bsize
maindefault branch
—last push
23.8 KB · 507 lines python
Raw Blame History
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