Files
gpu-program-swapper/server.py
drjones 172d812820 Harden .gitignore and drop hardcoded home paths before publishing
The repo is about to be pushed to a remote, so this covers what should never
travel with it and what should not be baked into the source.

.gitignore now covers credentials (.env, keys, tokens, .netrc), host-local
config (*.local.json), the SQLite telemetry store and its WAL sidecars, logs,
benchmark and sweep output, and the timestamped .bak files this project has
accumulated before. Verified that no currently tracked file is caught by the
new patterns.

Also removed /home/drjones from tracked source, which an ignore file cannot
help with. BASE_DIR now derives from the module's own location, the ComfyUI
model directory falls back to ~/ComfyUI/models, and start_manager.sh resolves
its interpreter through $HOME with a python3 fallback. All three still resolve
to exactly the same paths on this machine; they just no longer hardcode one
user's home directory into a published repository.

Note for the record: the git history was scanned across all refs and contains
no credentials. The password visible in `git remote -v` lives only in
.git/config, which is never pushed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-08-29 10:30:37 -07:00

552 lines
25 KiB
Python

"""FastAPI Backend Server with SSE Real-Time Telemetry and Model Orchestration API."""
import asyncio
import contextlib
import json
import logging
import os
import signal
import time
from contextlib import asynccontextmanager
from typing import Dict, Any, Optional, List, Set
from fastapi import FastAPI, Request, HTTPException, Query
from fastapi.responses import HTMLResponse, StreamingResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
import autotune
import overclock_manager
import ram_optimizer
import telemetry_store
import thermal_governor
import vram_arbitrator
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(name)s: %(message)s")
logger = logging.getLogger("model_manager_server")
# The sampler makes four HTTP calls a second; at INFO, httpx narrates every one of them.
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
BASE_DIR = os.environ.get("HYPERSWAP_BASE_DIR", os.path.dirname(os.path.abspath(__file__)))
# ==========================================
# TELEMETRY BROKER
# ==========================================
class TelemetryBroker:
"""One sampler, many subscribers.
Every SSE client used to run its own copy of the full snapshot once per second:
NVML queries, /proc/meminfo, an HTTP round-trip each to Ollama and ComfyUI, and — the
expensive one — a recursive walk of the ComfyUI models tree with a stat() per
checkpoint. Opening the dashboard in three tabs tripled the load on the very thing it
was measuring. Now a single background task samples at 1 Hz and fans the snapshot out.
The sampler is also the natural feed for the thermal governor and the persistence
layer, so neither needs to poll the GPU on its own.
"""
def __init__(self, interval_s: float = 1.0) -> None:
self.interval_s = interval_s
self.snapshot: Dict[str, Any] = {}
self.subscribers: Set[asyncio.Queue] = set()
self.task: Optional[asyncio.Task] = None
self.running = False
self.samples = 0
self.last_sample_ms = 0.0
# Set on shutdown so open SSE generators finish instead of holding the server up.
self.closing = False
async def start(self) -> None:
if self.running:
return
self.running = True
self.task = asyncio.create_task(self._loop())
def begin_shutdown(self) -> None:
"""Release every SSE subscriber. Safe to call from a signal handler."""
self.closing = True
for q in list(self.subscribers):
with contextlib.suppress(asyncio.QueueFull):
q.put_nowait(None)
async def stop(self) -> None:
self.running = False
self.closing = True
# Wake every subscriber so their generator can return. Without this, uvicorn waits
# on the open SSE responses during graceful shutdown and systemd eventually
# SIGKILLs the unit -- which skips the in-process GPU restore hook entirely.
for q in list(self.subscribers):
with contextlib.suppress(asyncio.QueueFull):
q.put_nowait(None)
if self.task:
self.task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await self.task
def subscribe(self) -> asyncio.Queue:
q: asyncio.Queue = asyncio.Queue(maxsize=2)
self.subscribers.add(q)
return q
def unsubscribe(self, q: asyncio.Queue) -> None:
self.subscribers.discard(q)
async def _loop(self) -> None:
while self.running:
t0 = time.perf_counter()
try:
snap = await self._sample()
self.snapshot = snap
self.samples += 1
self.last_sample_ms = round((time.perf_counter() - t0) * 1000, 2)
# Feed the governor and the durable store from the sample we already have.
thermal_governor.governor.observe(snap.get("gpu", {}),
overclock_manager.ACTIVE_PROFILE)
telemetry_store.record_telemetry(
snap.get("gpu", {}), snap.get("ram", {}),
profile=overclock_manager.ACTIVE_PROFILE,
throttle_reasons=",".join(snap.get("gpu", {}).get("throttle_reasons") or []),
)
for q in list(self.subscribers):
if q.full():
# Slow client: drop the stale frame rather than stalling the sampler.
with contextlib.suppress(asyncio.QueueEmpty):
q.get_nowait()
with contextlib.suppress(asyncio.QueueFull):
q.put_nowait(snap)
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(f"telemetry sampler error: {e}")
await asyncio.sleep(max(self.interval_s - (time.perf_counter() - t0), 0.05))
async def _sample(self) -> Dict[str, Any]:
gpu_stats = vram_arbitrator.get_gpu_hardware_stats()
mem_stats = ram_optimizer.get_detailed_meminfo()
ollama_state, comfy_state = await asyncio.gather(
vram_arbitrator.get_ollama_live_state(),
vram_arbitrator.get_comfyui_live_state(),
)
return {
"timestamp": time.time(), # wall clock, not the event loop's monotonic clock
"monotonic": asyncio.get_running_loop().time(),
"gpu": gpu_stats,
"ram": mem_stats,
"ollama": ollama_state,
"comfyui": comfy_state,
"arbitrator": vram_arbitrator.arbitrator.get_status(),
"governor": thermal_governor.governor.get_status(),
"overclock": {"active_profile": overclock_manager.ACTIVE_PROFILE},
"history": vram_arbitrator.get_switch_history(),
"comfy_models_count": len(ram_optimizer.find_comfy_model_files()),
"sampler": {"samples": self.samples, "last_sample_ms": self.last_sample_ms,
"subscribers": len(self.subscribers)},
}
async def get(self) -> Dict[str, Any]:
"""Latest snapshot, sampling on demand if the loop has not produced one yet."""
if not self.snapshot:
self.snapshot = await self._sample()
return self.snapshot
broker = TelemetryBroker()
def _install_shutdown_hook() -> None:
"""Close SSE streams the moment a shutdown signal arrives.
uvicorn runs the lifespan shutdown only after it has finished waiting on open
connections, so releasing subscribers from there is too late: the streams keep the
server busy until the graceful timeout expires and every one of them is force
cancelled, which logs a CancelledError traceback apiece. Chaining onto the existing
signal handler lets us drain them first and leaves uvicorn's own shutdown intact.
"""
loop = asyncio.get_running_loop()
for sig in (signal.SIGTERM, signal.SIGINT):
previous = signal.getsignal(sig)
def handler(signum, frame, _prev=previous):
broker.begin_shutdown()
if callable(_prev):
_prev(signum, frame)
try:
signal.signal(sig, handler)
except (ValueError, OSError):
pass # not on the main thread; the lifespan path still cleans up
@asynccontextmanager
async def lifespan(app: FastAPI):
telemetry_store.start()
await broker.start()
await vram_arbitrator.arbitrator.start()
_install_shutdown_hook()
yield
await vram_arbitrator.arbitrator.stop()
await broker.stop()
# Never leave the card with locked clocks and pinned fans after we exit.
try:
overclock_manager.restore_safe("server shutdown")
except Exception as e:
logger.error(f"restore_safe on shutdown failed: {e}")
telemetry_store.stop()
app = FastAPI(
title="HyperSwap // GPU Program Swapper & Telemetry API",
version="2.0.0",
description="High-performance VRAM arbitration and 64GB RAM cache orchestrator for simultaneous Ollama and ComfyUI workloads on Linux.",
docs_url="/docs",
redoc_url="/redoc",
lifespan=lifespan,
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Pydantic Request Models
class SwitchRequest(BaseModel):
model: str = Field(..., description="Name of the Ollama model to hot-swap to in VRAM", example="qwen3.8fast:latest")
keep_alive: Optional[str] = Field("30m", description="Keep-alive duration in VRAM (e.g. 5m, 30m, 0)", example="30m")
free_comfy_first: bool = Field(False, description="Purge ComfyUI VRAM first if it is holding memory")
class WarmRequest(BaseModel):
model_name: Optional[str] = Field(None, description="Ollama model name to warm", example="gemma4:26b")
filepath: Optional[str] = Field(None, description="Absolute file path of Safetensors/GGUF to warm into RAM")
blob_only: bool = Field(False, description="Warm the model's weights into page cache without loading VRAM")
force: bool = Field(False, description="Warm even if residency sampling thinks it is already resident")
class WarmAllRequest(BaseModel):
budget_gb: Optional[float] = Field(None, description="Byte budget for warming; defaults to 70% of MemAvailable", example=24.0)
class BenchmarkRequest(BaseModel):
iterations: Optional[int] = Field(2, description="Number of back-and-forth switch iterations to measure", example=2)
models: Optional[List[str]] = Field(None, description="Optional pair of models to benchmark between")
class OverclockApplyRequest(BaseModel):
profile: str = Field(..., description="Profile name: ollama | comfy | balanced", example="ollama")
class OverclockProfileUpdate(BaseModel):
config: Dict[str, Any] = Field(..., description="Profile settings dict", example={"power_limit_w": 370, "core_offset_mhz": 100})
class FanRequest(BaseModel):
mode: str = Field("auto", description="'auto' or 'manual'", example="manual")
percent: Optional[int] = Field(None, description="Fan speed 30-100 when mode=manual", example=70)
speed_pct: Optional[int] = Field(None, description="Alias for percent (30-100)", example=70)
class GovernorRequest(BaseModel):
enabled: Optional[bool] = Field(None, description="Enable or disable the thermal governor")
reset: bool = Field(False, description="Clear any active derate and reapply the full profile")
class SweepRequest(BaseModel):
knob: str = Field("mem_offset_mhz", description="mem_offset_mhz | core_offset_mhz | lock_mem_mhz | lock_core_max")
profile: str = Field("ollama", description="Profile to tune")
workload: str = Field("auto", description="ollama (decode tok/s) | comfy (diffusion it/s) | auto")
model: Optional[str] = Field(None, description="Model to benchmark with; defaults to the loaded one")
start: Optional[int] = Field(None, description="First offset value")
stop: Optional[int] = Field(None, description="Last offset value")
step: Optional[int] = Field(None, description="Offset increment")
repeats: int = Field(1, description="Benchmark runs per step")
max_steps: int = Field(6, description="Cap on swept values for discrete clock knobs")
include_unlocked: bool = Field(True, description="Include an unlocked (0) control step")
apply_best: bool = Field(False, description="Write the winning value into the profile")
class RequestVramRequest(BaseModel):
needed_gb: float = Field(0.0, description="How much free VRAM Ollama needs", example=12.0)
# ==========================================
# REST API ENDPOINTS
# ==========================================
@app.get("/api/stats", summary="Full System Snapshot", tags=["Telemetry"])
async def get_all_stats() -> Dict[str, Any]:
"""Latest unified snapshot of GPU hardware, host RAM, Ollama, ComfyUI and swap history."""
return await broker.get()
@app.get("/api/gpu", summary="GPU Sensors and VRAM Breakdown", tags=["Telemetry"])
async def get_gpu_metrics() -> Dict[str, Any]:
"""Detailed NVML sensors (utilization, temp, power, fan, clocks, throttle reasons, per-process VRAM)."""
return vram_arbitrator.get_gpu_hardware_stats()
@app.get("/api/memory", summary="Host RAM and Page Cache Breakdown", tags=["Telemetry"])
async def get_ram_metrics() -> Dict[str, Any]:
"""Precise host RAM breakdown, active cache size and cache ratios."""
return ram_optimizer.get_detailed_meminfo()
@app.get("/api/stream", summary="Real-Time SSE Telemetry Stream", tags=["Telemetry"])
async def sse_telemetry_stream(request: Request):
"""Server-Sent Events stream of the shared 1Hz snapshot."""
async def event_generator():
q = broker.subscribe()
try:
snap = await broker.get()
yield f"data: {json.dumps(snap)}\n\n"
while not broker.closing:
if await request.is_disconnected():
break
try:
snap = await asyncio.wait_for(q.get(), timeout=5.0)
if snap is None: # shutdown sentinel
break
yield f"data: {json.dumps(snap)}\n\n"
except asyncio.TimeoutError:
yield ": keepalive\n\n"
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(f"SSE stream error: {e}")
yield f"data: {json.dumps({'error': str(e)})}\n\n"
finally:
broker.unsubscribe(q)
return StreamingResponse(
event_generator(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
}
)
@app.post("/api/switch-model", summary="Hot-Swap Ollama LLM in VRAM", tags=["Orchestration"])
async def api_switch_model(req: SwitchRequest):
"""Hot-swap the active Ollama model, measuring real load bandwidth and token throughput."""
if req.free_comfy_first:
await vram_arbitrator.arbitrator.request_vram_for_ollama()
res = await vram_arbitrator.switch_ollama_model(req.model, keep_alive=req.keep_alive or "30m")
if not res.get("success"):
raise HTTPException(status_code=500, detail=res.get("error"))
return res
@app.post("/api/free-vram", summary="Soft-Yield Ollama VRAM", tags=["Orchestration"])
async def api_free_vram(confirm: bool = Query(True, description="Wait for the driver to actually release the allocation")):
"""Yield Ollama's VRAM and wait for the release to be confirmed by NVML."""
return await vram_arbitrator.instant_free_ollama_vram(confirm=confirm)
@app.post("/api/comfy-free", summary="Purge ComfyUI VRAM Cache", tags=["Orchestration"])
async def api_comfy_free():
"""Purge loaded diffusion models and VRAM cache from the ComfyUI pipeline."""
return await vram_arbitrator.instant_free_comfyui_vram()
@app.post("/api/request-vram", summary="Ask for VRAM on Ollama's behalf", tags=["Orchestration"])
async def api_request_vram(req: RequestVramRequest):
"""Force an immediate ComfyUI purge if there is not enough free VRAM for Ollama."""
return await vram_arbitrator.arbitrator.request_vram_for_ollama(req.needed_gb)
@app.post("/api/warm-all", summary="Pre-warm Models into RAM Cache", tags=["Memory Optimization"])
async def api_warm_all(req: Optional[WarmAllRequest] = None):
"""Warm the highest-value models into the page cache within a byte budget."""
return await ram_optimizer.warm_all_models(budget_gb=req.budget_gb if req else None)
@app.get("/api/warm-plan", summary="Preview the Warm Plan", tags=["Memory Optimization"])
async def api_warm_plan(budget_gb: Optional[float] = Query(None, description="Override the byte budget")):
"""Show what warming would read, in what order, and what it would skip — without doing it."""
return ram_optimizer.build_warm_plan(budget_gb)
@app.post("/api/warm-model", summary="Pre-warm Single Model or File", tags=["Memory Optimization"])
async def api_warm_model(req: WarmRequest):
"""Pre-warm a specific Ollama model or file path into the Linux page cache."""
if req.model_name and req.blob_only:
return ram_optimizer.warm_ollama_blob(req.model_name, force=req.force)
if req.model_name:
return await ram_optimizer.warm_ollama_model(req.model_name, keep_alive="1m")
if req.filepath:
return ram_optimizer.warm_file_to_ram(req.filepath, force=req.force)
raise HTTPException(status_code=400, detail="model_name or filepath required")
@app.get("/api/cache/report", summary="Measured Page-Cache Residency", tags=["Memory Optimization"])
async def api_cache_report(files: bool = Query(True), refresh: bool = Query(False)):
"""Measured (not assumed) page-cache residency for every model on disk."""
report = ram_optimizer.get_cache_report(include_files=files, force_refresh=refresh)
report["capability"] = ram_optimizer.residency_capability()
return report
@app.get("/api/models", summary="List All Installed Models", tags=["Catalog"])
async def api_get_models(refresh: bool = Query(False)):
"""All installed Ollama models (with their on-disk blobs) and ComfyUI checkpoints."""
catalog = ram_optimizer.get_model_catalog(force_refresh=refresh)
ollama_state = await vram_arbitrator.get_ollama_live_state()
return {
"ollama_models": ollama_state.get("installed_models", []),
"ollama_blobs": catalog["ollama"],
"comfy_models": catalog["comfy"],
"cached_at": catalog["cached_at"],
}
@app.get("/api/history", summary="Model Switch History Log", tags=["Analytics"])
async def api_get_history(limit: int = Query(20, description="Max history items to return"),
durable: bool = Query(False, description="Read from the persistent store instead of the in-memory ring")):
"""Recent swap events, durations, achieved bandwidth and cache status."""
if durable:
return telemetry_store.recent_events(limit)
return vram_arbitrator.get_switch_history()[:limit]
@app.post("/api/benchmark", summary="Run Latency Benchmark", tags=["Analytics"])
async def api_run_benchmark(req: BenchmarkRequest):
"""Automated round-trip switch benchmark measuring latency and cache effectiveness."""
from mcp_server import run_model_switch_benchmark
res_str = await run_model_switch_benchmark(iterations=req.iterations or 2)
return json.loads(res_str)
# ==========================================
# ANALYTICS (persisted)
# ==========================================
@app.get("/api/analytics/profiles", summary="Which Overclock Profile Is Actually Faster", tags=["Analytics"])
async def api_analytics_profiles(days: float = Query(7.0)):
"""Decode throughput and thermals grouped by the profile that was active at the time."""
return {"window_days": days, "profiles": telemetry_store.profile_comparison(days)}
@app.get("/api/analytics/swaps", summary="Swap Statistics", tags=["Analytics"])
async def api_analytics_swaps(days: float = Query(7.0)):
"""Aggregated swap/yield/purge latencies, cache-hit split and per-model throughput."""
return telemetry_store.swap_stats(days)
@app.get("/api/analytics/timeseries", summary="Downsampled Telemetry History", tags=["Analytics"])
async def api_analytics_timeseries(hours: float = Query(6.0), buckets: int = Query(240)):
"""Long-range history for charts that outlive a page refresh."""
return {"hours": hours, "points": telemetry_store.timeseries(hours, buckets)}
@app.get("/api/analytics/models", summary="Model Usage Ranking", tags=["Analytics"])
async def api_analytics_models(days: float = Query(30.0)):
"""Recency/frequency ranking used to prioritise the RAM warm budget."""
return {"window_days": days, "models": telemetry_store.model_usage_ranking(days)}
@app.get("/api/db", summary="Telemetry Store Info", tags=["Analytics"])
async def api_db_info():
"""Where the persistent store lives and how much history it holds."""
return telemetry_store.db_info()
# ==========================================
# OVERCLOCK MANAGEMENT
# ==========================================
@app.get("/api/overclock", summary="Overclock Status & Profiles", tags=["Overclock"])
async def api_overclock_status():
"""Live GPU overclock state, active profile, governor state and all per-app profiles."""
status = overclock_manager.get_status()
status["governor"] = thermal_governor.governor.get_status()
return status
@app.post("/api/overclock/apply", summary="Apply Overclock Profile", tags=["Overclock"])
async def api_overclock_apply(req: OverclockApplyRequest):
"""Apply a named overclock profile (ollama | comfy | balanced) immediately."""
res = overclock_manager.apply_profile(req.profile)
if not res.get("success"):
raise HTTPException(status_code=400, detail=res.get("error"))
return res
@app.post("/api/overclock/restore", summary="Restore Stock GPU State", tags=["Overclock"])
async def api_overclock_restore():
"""Drop all clock locks and offsets, restore default power limit and automatic fans."""
return overclock_manager.restore_safe("manual request")
@app.get("/api/overclock/profiles", summary="List Overclock Profiles", tags=["Overclock"])
async def api_overclock_profiles():
"""All overclock profiles with their current settings."""
return overclock_manager.get_profiles()
@app.post("/api/overclock/profiles/{name}", summary="Update Overclock Profile", tags=["Overclock"])
async def api_overclock_update_profile(name: str, req: OverclockProfileUpdate):
"""Update a profile's settings (persisted to disk)."""
res = overclock_manager.set_profile(name, req.config)
if not res.get("success"):
raise HTTPException(status_code=400, detail=res.get("error"))
return {"success": True, "profile": name, "profiles": res.get("profiles")}
@app.get("/api/overclock/fan", summary="Get GPU Fan Status", tags=["Overclock"])
@app.get("/api/gpu/fan", summary="Get GPU Fan Status", tags=["Overclock"])
async def api_get_fan_status():
"""Current GPU fan control mode and speed."""
return overclock_manager.get_fan_status()
@app.post("/api/overclock/fan", summary="Set GPU Fan Speed", tags=["Overclock"])
@app.post("/api/gpu/fan", summary="Set GPU Fan Speed", tags=["Overclock"])
async def api_set_fan(req: FanRequest):
"""Set the GPU fan to manual speed (30-100%) or back to automatic control."""
pct = req.percent if req.percent is not None else req.speed_pct
if req.mode == "manual" and pct is not None:
return overclock_manager.set_fan_speed(pct)
return overclock_manager.set_fan_auto()
# ==========================================
# THERMAL GOVERNOR
# ==========================================
@app.get("/api/governor", summary="Thermal Governor State", tags=["Governor"])
async def api_governor_status():
"""Current derate level, why it was applied, and the escalation history."""
return thermal_governor.governor.get_status()
@app.post("/api/governor", summary="Control the Thermal Governor", tags=["Governor"])
async def api_governor_control(req: GovernorRequest):
"""Enable/disable the governor, or clear an active derate."""
if req.enabled is not None:
thermal_governor.governor.set_enabled(req.enabled)
if req.reset:
thermal_governor.governor.reset()
return thermal_governor.governor.get_status()
# ==========================================
# AUTOTUNE
# ==========================================
@app.get("/api/autotune", summary="Autotune Status & History", tags=["Autotune"])
async def api_autotune_status():
"""Sweep progress, the last result, and every recorded autotune step."""
return autotune.get_status()
@app.post("/api/autotune/sweep", summary="Run an Overclock Sweep", tags=["Autotune"])
async def api_autotune_sweep(req: SweepRequest):
"""Walk a clock offset upward, measuring tok/s and watching for instability at each step."""
res = await autotune.sweep(
knob=req.knob, profile=req.profile, workload=req.workload, model=req.model,
start=req.start, stop=req.stop, step=req.step,
repeats=req.repeats, max_steps=req.max_steps,
include_unlocked=req.include_unlocked, apply_best=req.apply_best,
)
if not res.get("success"):
raise HTTPException(status_code=400, detail=res.get("error"))
return res
@app.post("/api/autotune/cancel", summary="Cancel a Running Sweep", tags=["Autotune"])
async def api_autotune_cancel():
"""Stop the current sweep after the step in flight; the profile is restored either way."""
return autotune.cancel()
# Mount static web UI files
app.mount("/static", StaticFiles(directory=f"{BASE_DIR}/static"), name="static")
@app.get("/", summary="Dashboard Web UI", tags=["UI"])
async def root_index():
with open(f"{BASE_DIR}/static/index.html", "r") as f:
content = f.read()
return HTMLResponse(content=content)
if __name__ == "__main__":
import uvicorn
uvicorn.run("server:app", host="0.0.0.0", port=9090, reload=False, log_level="info",
timeout_graceful_shutdown=10)