#!/usr/bin/env python3
"""
Measure what actually happens to VRAM when two local-AI tools share one GPU.

Written because a supposedly-idle Ollama was holding 8.3 GB of a 10 GB card
after `ollama stop`, which starved a ComfyUI job. The documented knobs did not
behave the way the docs imply, and nobody seems to have written down what they
actually do. So: measure it.

Everything here is reproducible. Run it on your own card and open a PR with the
results file — the interesting question is how much of this is Ollama's
behaviour versus a property of a particular GPU.

Usage:
    python3 bench.py --model qwen3.5:9b-q4_K_M
    python3 bench.py --model qwen3.5:9b-q4_K_M --skip-slow
"""

from __future__ import annotations

import argparse
import json
import platform
import subprocess
import time
import urllib.request
from datetime import datetime, timezone
from pathlib import Path

OLLAMA = "http://127.0.0.1:11434"


def nvidia(query: str) -> list[str]:
    try:
        out = subprocess.run(
            ["nvidia-smi", f"--query-{query}", "--format=csv,noheader,nounits"],
            capture_output=True, text=True, timeout=15,
        ).stdout.strip()
        return [l.strip() for l in out.splitlines() if l.strip()]
    except (subprocess.SubprocessError, FileNotFoundError, OSError):
        return []


def vram_used_mb() -> int:
    """Total VRAM in use on the card, all processes."""
    v = nvidia("gpu=memory.used")
    return int(v[0]) if v else -1


def processes() -> dict[str, int]:
    """Per-process VRAM, keyed by process name."""
    out: dict[str, int] = {}
    for line in nvidia("compute-apps=process_name,used_memory"):
        parts = [p.strip() for p in line.split(",")]
        if len(parts) == 2 and parts[1].isdigit():
            out[parts[0]] = out.get(parts[0], 0) + int(parts[1])
    return out


def ollama_vram_mb() -> int:
    return sum(v for k, v in processes().items()
               if "ollama" in k.lower() or "llama" in k.lower())


def api(path: str, payload: dict | None = None, timeout: int = 180):
    req = urllib.request.Request(
        f"{OLLAMA}{path}",
        data=json.dumps(payload).encode() if payload is not None else None,
        headers={"Content-Type": "application/json"},
    )
    with urllib.request.urlopen(req, timeout=timeout) as r:
        return json.loads(r.read())


def wait_until(predicate, timeout: float = 120, interval: float = 0.5) -> float:
    """Seconds until predicate() is true, or -1 on timeout."""
    start = time.time()
    while time.time() - start < timeout:
        if predicate():
            return round(time.time() - start, 2)
        time.sleep(interval)
    return -1.0


def generate(model: str, prompt: str = "Say OK.", keep_alive=None) -> float:
    body = {"model": model, "prompt": prompt, "stream": False}
    if keep_alive is not None:
        body["keep_alive"] = keep_alive
    t = time.time()
    api("/api/generate", body)
    return round(time.time() - t, 2)


def evict(model: str) -> None:
    try:
        api("/api/generate", {"model": model, "keep_alive": 0}, timeout=60)
    except Exception:
        pass


def measure(name: str, fn) -> dict:
    print(f"  {name} ...", end="", flush=True)
    r = fn()
    print(f" {r}")
    return {"test": name, **r}


def host_info() -> dict:
    gpu = nvidia("gpu=name,memory.total,driver_version")
    name, total, driver = (gpu[0].split(", ") + ["", "", ""])[:3] if gpu else ("", "", "")
    try:
        ver = subprocess.run(["ollama", "--version"], capture_output=True,
                             text=True, timeout=10).stdout.strip()
    except Exception:
        ver = "unknown"
    return {
        "gpu": name, "vram_total_mb": total, "driver": driver,
        "ollama_version": ver, "platform": platform.platform(),
        "measured_at": datetime.now(timezone.utc).isoformat(),
    }


def run(model: str, skip_slow: bool) -> dict:
    results = []
    print(f"\nbenchmarking {model}\n")

    baseline = vram_used_mb()
    others = {k: v for k, v in processes().items()
              if "ollama" not in k.lower() and "llama" not in k.lower()}
    print(f"  baseline VRAM in use: {baseline} MiB")
    if others:
        print(f"  other processes present: {others}")

    # 1. What does a cold load cost, in time and VRAM?
    evict(model); time.sleep(3)
    before = ollama_vram_mb()
    cold = generate(model, keep_alive="5m")
    resident = ollama_vram_mb()
    results.append({"test": "cold_load", "seconds": cold,
                    "vram_before_mb": before, "vram_after_mb": resident,
                    "model_footprint_mb": resident - before})

    # 2. Warm call, model already resident.
    warm = generate(model, keep_alive="5m")
    results.append({"test": "warm_call", "seconds": warm,
                    "vram_mb": ollama_vram_mb()})

    # 3. Does `ollama stop` actually free the card?
    subprocess.run(["ollama", "stop", model], capture_output=True, timeout=30)
    time.sleep(4)
    after_stop = ollama_vram_mb()
    results.append({"test": "ollama_stop_frees_vram",
                    "vram_after_mb": after_stop,
                    "freed": after_stop < 200})

    # 4. Does the API keep_alive:0 free it?
    if after_stop > 200:
        generate(model, keep_alive="5m")
    evict(model)
    freed_in = wait_until(lambda: ollama_vram_mb() < 200, timeout=60)
    results.append({"test": "api_keep_alive_0_frees_vram",
                    "seconds_to_free": freed_in,
                    "vram_after_mb": ollama_vram_mb(),
                    "freed": freed_in >= 0})

    # 5. Cost of keep_alive:0 — every call pays the cold-load price.
    if not skip_slow:
        times = [generate(model, keep_alive=0) for _ in range(3)]
        results.append({"test": "keep_alive_0_repeat_calls",
                        "seconds_each": times,
                        "mean": round(sum(times) / len(times), 2)})
        evict(model); time.sleep(2)
        generate(model, keep_alive="5m")
        warm_times = [generate(model, keep_alive="5m") for _ in range(3)]
        results.append({"test": "keep_alive_5m_repeat_calls",
                        "seconds_each": warm_times,
                        "mean": round(sum(warm_times) / len(warm_times), 2)})
        evict(model)

    # 6. How long does idle eviction actually take at a short keep_alive?
    if not skip_slow:
        generate(model, keep_alive="10s")
        gone = wait_until(lambda: ollama_vram_mb() < 200, timeout=90)
        results.append({"test": "idle_eviction_at_10s_keep_alive",
                        "seconds_to_free": gone,
                        "note": "time from last call until VRAM released"})
        evict(model)

    return {"host": host_info(), "model": model, "baseline_vram_mb": baseline,
            "other_processes_at_start": others, "results": results}


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--skip-slow", action="store_true")
    ap.add_argument("--out", default=None)
    a = ap.parse_args()

    data = run(a.model, a.skip_slow)
    slug = (data["host"]["gpu"] or "gpu").lower().replace(" ", "-").replace("/", "-")
    out = Path(a.out or f"results/{slug}--{a.model.replace(':','-').replace('/','-')}.json")
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps(data, indent=2))
    print(f"\nwrote {out}")
    for r in data["results"]:
        print(" ", json.dumps(r))


if __name__ == "__main__":
    main()
