Add an end-to-end arbitration verifier; return honest HTTP status codes

The unit suite covers logic in isolation, but the promise this service exists to make
-- an LLM and a diffusion pipeline sharing one 16 GB card without either failing --
had only ever been checked by hand, piecemeal. verify_arbitration.py walks the whole
cycle against real hardware and reports what happened at each stage: load and its
bandwidth classification, the confirmed yield, a real SDXL graph, the deferred idle
purge, reclaim-and-retry, and finally whether VRAM attribution adds up and the
reported GPU state still matches the card. It restores what it changes and refuses
to start if ComfyUI is busy. Kept out of pytest deliberately: it moves real VRAM and
takes minutes.

Running it immediately found two bugs.

/api/switch-model reported every upstream failure as 500. Asking an embedding model
to generate makes Ollama return 400 -- the request is unusable, the service is fine
-- and calling that an Internal Server Error blames this service for the caller's
mistake. Failures now map to 400 for an upstream client error, 507 for a model that
will not fit (valid request, healthy service, no room), and 502 when Ollama itself
errors.

The verifier also picked the smallest installed model, which here is
nomic-embed-text -- an embedding model with no generate endpoint. It now filters
those out by family and name.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
drjones
2026-09-06 17:25:42 -07:00
parent 043d61722b
commit aed1c360f0
3 changed files with 264 additions and 1 deletions

View File

@@ -383,7 +383,18 @@ async def api_switch_model(req: SwitchRequest):
await vram_arbitrator.arbitrator.request_vram_for_ollama() await vram_arbitrator.arbitrator.request_vram_for_ollama()
res = await vram_arbitrator.switch_ollama_model(req.model, keep_alive=req.keep_alive or "30m") res = await vram_arbitrator.switch_ollama_model(req.model, keep_alive=req.keep_alive or "30m")
if not res.get("success"): if not res.get("success"):
raise HTTPException(status_code=500, detail=res.get("error")) # Reflect what actually went wrong. Ollama returns 400 for an unusable request --
# asking an embedding model to generate, say -- and reporting that as 500 blames
# this service for the caller's mistake. A model that will not fit is neither:
# the request is valid and the service is healthy, there is simply no room.
upstream = res.get("upstream_status")
if res.get("vram_oom"):
status = 507 # Insufficient Storage
elif isinstance(upstream, int) and 400 <= upstream < 500:
status = 400
else:
status = 502 if upstream else 500
raise HTTPException(status_code=status, detail=res.get("error"))
return res return res
@app.post("/api/free-vram", summary="Soft-Yield Ollama VRAM", tags=["Orchestration"]) @app.post("/api/free-vram", summary="Soft-Yield Ollama VRAM", tags=["Orchestration"])

251
verify_arbitration.py Executable file
View File

@@ -0,0 +1,251 @@
#!/usr/bin/env python
"""End-to-end verification of the arbitration cycle, against real hardware.
The unit suite covers logic in isolation; this exercises the promise the whole service
exists to make -- that an LLM and a diffusion pipeline can share one 16 GB card without
either failing -- and reports what actually happened at each stage.
It is deliberately not part of `pytest tests/`: it loads real models, runs a real
diffusion graph and moves real VRAM, taking a few minutes. Run it when you want proof
the system works on this machine:
python verify_arbitration.py # full cycle
python verify_arbitration.py --quick # skip the diffusion stages
Every stage restores what it changed, and the script refuses to start if ComfyUI is
already busy.
"""
import argparse
import asyncio
import sys
import time
from typing import Any, Dict, List, Optional
import httpx
BASE = "http://localhost:9090"
PASS, FAIL, SKIP, WARN = "PASS", "FAIL", "SKIP", "WARN"
_results: List[Dict[str, Any]] = []
def record(stage: str, status: str, detail: str, evidence: str = "") -> None:
_results.append({"stage": stage, "status": status, "detail": detail,
"evidence": evidence})
colour = {"PASS": "\033[32m", "FAIL": "\033[31m",
"SKIP": "\033[33m", "WARN": "\033[33m"}[status]
print(f" {colour}{status:<4}\033[0m {stage:<38} {detail}")
if evidence:
print(f" {evidence}")
async def api(client: httpx.AsyncClient, method: str, path: str, **kw) -> Any:
r = await client.request(method, f"{BASE}{path}", **kw)
r.raise_for_status()
return r.json()
async def stage_preflight(c: httpx.AsyncClient) -> Optional[Dict[str, Any]]:
health = await api(c, "GET", "/api/health")
failed = health.get("failed") or []
if failed:
record("preflight: dependencies", FAIL,
f"{len(failed)} dependency check(s) failing", ", ".join(failed))
return None
record("preflight: dependencies", PASS, health["summary"])
stats = await api(c, "GET", "/api/stats")
if stats["comfyui"].get("executing") or stats["comfyui"].get("queue_remaining"):
record("preflight: ComfyUI idle", FAIL, "ComfyUI is busy; refusing to interfere")
return None
record("preflight: ComfyUI idle", PASS, "queue empty")
return stats
async def stage_llm_load(c: httpx.AsyncClient, model: str) -> bool:
res = await api(c, "POST", "/api/switch-model",
json={"model": model, "keep_alive": "5m"}, timeout=300)
if not res.get("success"):
record("LLM loads", FAIL, res.get("error", "")[:90])
return False
gbps, status = res.get("load_gbps"), res.get("cache_status")
record("LLM loads", PASS, f"{res['model_size_gb']} GB in {res['load_duration_ms']:.0f} ms",
f"{gbps} GB/s -> {status}; {res['tokens_per_sec']} tok/s")
# The classification must follow from the measured bandwidth, not a fixed duration.
if gbps is not None:
expected = ("RAM Cache Hit" if gbps >= 2.0
else "Partial Cache" if gbps >= 0.8 else "Cold Disk Load")
ok = expected.split()[0] in (status or "")
record("load classified by bandwidth", PASS if ok else FAIL,
f"{gbps} GB/s reported as '{status}'",
"" if ok else f"expected something matching '{expected}'")
return True
async def stage_yield(c: httpx.AsyncClient) -> bool:
before = await api(c, "GET", "/api/gpu")
held = before["breakdown"]["ollama_gb"]
res = await api(c, "POST", "/api/free-vram", timeout=60)
outcome = res.get("outcome")
if outcome == "released":
after = await api(c, "GET", "/api/gpu")
freed = held - after["breakdown"]["ollama_gb"]
# The barrier's promise: on return, the VRAM is genuinely gone.
ok = after["breakdown"]["ollama_gb"] < 0.3
record("VRAM yield is confirmed", PASS if ok else FAIL,
f"released in {res['confirm_ms']} ms",
f"{freed:.2f} GB freed; NVML now reports "
f"{after['breakdown']['ollama_gb']} GB held by Ollama")
return ok
if outcome == "busy":
record("VRAM yield is confirmed", WARN,
"model was mid-generation, so the unload was deferred",
res.get("error", ""))
return True
record("VRAM yield is confirmed", FAIL, f"outcome={outcome}", res.get("error", ""))
return False
async def stage_diffusion(c: httpx.AsyncClient) -> bool:
sys.path.insert(0, "/home/drjones/unified-model-manager")
import autotune, vram_arbitrator # noqa: E402 (imported late; needs the service's deps)
t0 = time.perf_counter()
res = await autotune._diffusion_benchmark()
if not res.get("ok"):
record("diffusion runs", FAIL, res.get("error", "")[:90])
return False
record("diffusion runs", PASS,
f"SDXL 1024/20 steps in {res['exec_ms']} ms", f"{res['it_per_sec']} it/s")
snap = vram_arbitrator.get_process_vram_bytes()
comfy_gb = snap["comfyui_bytes"] / (1024 ** 3)
record("ComfyUI holds its checkpoint", PASS if comfy_gb > 0.5 else WARN,
f"{comfy_gb:.2f} GB retained",
"held for the idle window rather than purged between iterations")
return True
async def stage_idle_purge(c: httpx.AsyncClient) -> bool:
stats = await api(c, "GET", "/api/stats")
arb = stats["arbitrator"]
if not arb.get("pending_purge"):
record("purge is deferred, not immediate", WARN,
"no purge pending (ComfyUI may already be clean)")
return True
record("purge is deferred, not immediate", PASS,
f"holding checkpoints for {arb.get('idle_purge_after_s')} s",
f"idle {arb.get('comfy_idle_s')} s so far")
return True
async def stage_reclaim(c: httpx.AsyncClient, model: str) -> bool:
"""The direction that used to fail outright: an LLM that will not fit."""
gpu = await api(c, "GET", "/api/gpu")
comfy_gb = gpu["breakdown"]["comfyui_gb"]
if comfy_gb < 0.5:
record("reclaims VRAM for the LLM", SKIP,
f"ComfyUI only holds {comfy_gb} GB; nothing to reclaim")
return True
res = await api(c, "POST", "/api/switch-model",
json={"model": model, "keep_alive": "2m"}, timeout=300)
if not res.get("success"):
record("reclaims VRAM for the LLM", FAIL, res.get("error", "")[:90],
str(res.get("unmanaged_blockers", ""))[:120])
return False
if res.get("reclaimed_from_comfyui_gb"):
record("reclaims VRAM for the LLM", PASS,
f"reclaimed {res['reclaimed_from_comfyui_gb']} GB and retried",
f"loaded at {res.get('load_gbps')} GB/s")
else:
record("reclaims VRAM for the LLM", PASS,
"fit without needing a reclaim", f"{res.get('load_gbps')} GB/s")
return True
async def stage_accounting(c: httpx.AsyncClient) -> bool:
"""Reported VRAM must add up, and reported settings must match the hardware."""
gpu = await api(c, "GET", "/api/gpu")
bd = gpu["breakdown"]
parts = bd["ollama_gb"] + bd["comfyui_gb"] + bd["system_gb"]
used = gpu["vram_used_gb"]
# Driver overhead means the parts never sum exactly; a large gap means mis-attribution.
ok = abs(parts - used) < 1.5
record("VRAM attribution adds up", PASS if ok else FAIL,
f"parts {parts:.2f} GB vs NVML used {used:.2f} GB",
f"ollama {bd['ollama_gb']} + comfy {bd['comfyui_gb']} + system {bd['system_gb']} "
f"(desktop {bd.get('desktop_gb')}, unmanaged {bd.get('unmanaged_gb')})")
oc = await api(c, "GET", "/api/overclock")
drift = oc["drift"]
ok2 = not drift["drifted"]
record("reported GPU state matches hardware", PASS if ok2 else FAIL,
f"profile '{drift['profile']}' asks {drift['power_limit_intended_w']} W, "
f"card reports {drift['power_limit_actual_w']} W",
drift.get("reason") or "")
return ok and ok2
async def main() -> int:
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--quick", action="store_true",
help="skip the diffusion and reclaim stages")
ap.add_argument("--model", default=None,
help="Ollama model to test with (default: smallest installed)")
args = ap.parse_args()
print("\nHyperSwap arbitration verification\n" + "=" * 62)
async with httpx.AsyncClient(timeout=60.0) as c:
baseline = await stage_preflight(c)
if baseline is None:
print("\nPreflight failed; not continuing.\n")
return 2
model = args.model
if not model:
models = (await api(c, "GET", "/api/models"))["ollama_models"]
# Embedding models have no /api/generate endpoint, and the smallest model
# installed is often one of them.
EMBED_FAMILIES = {"bert", "nomic-bert", "gte", "jina-bert"}
usable = [m for m in models
if (m.get("details", {}).get("family") or "").lower()
not in EMBED_FAMILIES and "embed" not in m["name"].lower()]
if not usable:
record("choose a test model", FAIL,
"no generative Ollama models installed "
f"({len(models)} found, all embedding-only)")
return 2
model = min(usable, key=lambda m: m.get("size", 0))["name"]
print(f"\n using model: {model}\n")
await stage_llm_load(c, model)
await stage_yield(c)
if not args.quick:
if await stage_diffusion(c):
await stage_idle_purge(c)
await stage_reclaim(c, model)
else:
record("diffusion stages", SKIP, "--quick")
await stage_accounting(c)
# Leave the box as we found it.
await api(c, "POST", "/api/free-vram", timeout=60)
failed = [r for r in _results if r["status"] == FAIL]
warned = [r for r in _results if r["status"] == WARN]
print("=" * 62)
print(f" {len(_results) - len(failed) - len(warned)} passed, "
f"{len(warned)} warned, {len(failed)} failed")
for r in failed:
print(f" FAILED: {r['stage']} — {r['detail']}")
print()
return 1 if failed else 0
if __name__ == "__main__":
sys.exit(asyncio.run(main()))

View File

@@ -880,6 +880,7 @@ async def switch_ollama_model(target_model: str, keep_alive: str = "30m",
return retry return retry
return {"success": False, "error": f"HTTP {resp.status_code}: {body}", return {"success": False, "error": f"HTTP {resp.status_code}: {body}",
"duration_ms": total_duration_ms, "duration_ms": total_duration_ms,
"upstream_status": resp.status_code,
"vram_oom": looks_like_vram_oom(body)} "vram_oom": looks_like_vram_oom(body)}
except Exception as e: except Exception as e:
return {"success": False, "error": str(e), return {"success": False, "error": str(e),