#!/usr/bin/env python3
"""
Measure how Ollama degrades as another process eats the same GPU.

    python3 contention.py --model qwen3.5:9b-q4_K_M --steps 0,1024,2048,3072,4096

The question `bench.py` does not answer: a local-AI box usually runs more than
one tool, and the second tool does not politely wait. What happens to the model
when something else is already holding half the card?

The co-tenant here is `ballast.py` rather than a real image-generation job, on
purpose. A diffusion job's footprint swings by gigabytes between steps, so using
one as the independent variable means the variable is not actually controlled.
The ballast holds a flat number, so the only thing changing between runs is how
much of the card was already gone when Ollama went to load.

At each step this records whether the load succeeded at all, how much of the
model Ollama put on the GPU versus spilled to CPU, what that did to throughput,
and whether the co-tenant survived. Run it on your own card and open a PR with
the results file.
"""

from __future__ import annotations

import argparse
import json
import re
import subprocess
import sys
import time
from datetime import datetime, timezone
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))

from bench import api, evict, host_info, ollama_vram_mb, processes, vram_used_mb, wait_until

# The ballast needs a torch that can see the GPU. On this box that is ComfyUI's
# venv; anywhere else, the interpreter running this script will usually do.
PYTHON_CANDIDATES = ["/home/tomo/ComfyUI/.venv/bin/python", sys.executable]

PROMPT = ("Explain in three sentences why memory bandwidth, rather than raw "
          "FLOPs, usually limits transformer inference on consumer GPUs.")


def torch_python() -> str:
    for p in PYTHON_CANDIDATES:
        try:
            r = subprocess.run([p, "-c", "import torch;assert torch.cuda.is_available()"],
                               capture_output=True, timeout=90)
            if r.returncode == 0:
                return p
        except (subprocess.SubprocessError, OSError):
            continue
    sys.exit("no Python with a working CUDA torch found; edit PYTHON_CANDIDATES")


class Ballast:
    """A co-tenant holding a fixed slice of the card."""

    def __init__(self, python: str, mib: int):
        self.python, self.mib, self.proc, self.granted = python, mib, None, 0

    def __enter__(self) -> "Ballast":
        if self.mib <= 0:
            return self
        here = Path(__file__).resolve().parent / "ballast.py"
        self.proc = subprocess.Popen([self.python, str(here), str(self.mib)],
                                     stdout=subprocess.PIPE, text=True)
        line = self.proc.stdout.readline().strip()
        if not line.startswith("READY"):
            raise RuntimeError(f"ballast failed to start: {line!r}")
        self.granted = int(line.split()[1])
        time.sleep(2)
        return self

    def alive(self) -> bool:
        return self.mib <= 0 or (self.proc is not None and self.proc.poll() is None)

    def __exit__(self, *_):
        if self.proc and self.proc.poll() is None:
            self.proc.terminate()
            try:
                self.proc.wait(timeout=30)
            except subprocess.TimeoutExpired:
                self.proc.kill()
        wait_until(lambda: not any("ballast" in k for k in processes()), timeout=30)
        time.sleep(2)


def ps_split(model: str) -> str:
    """Ollama's own view of where the layers went: `100% GPU`, `43%/57% CPU/GPU`."""
    try:
        out = subprocess.run(["ollama", "ps"], capture_output=True, text=True,
                             timeout=20).stdout
    except (subprocess.SubprocessError, OSError):
        return "unknown"
    for line in out.splitlines():
        if line.startswith(model.split(":")[0]):
            m = re.search(r"(\d+%(?:/\d+%)?\s+(?:CPU|GPU)(?:/GPU)?)", line)
            if m:
                return m.group(1).strip()
    return "not loaded"


def one_step(model: str, python: str, mib: int, total_mb: int) -> dict:
    evict(model)
    wait_until(lambda: ollama_vram_mb() < 200, timeout=60)
    time.sleep(2)

    with Ballast(python, mib) as b:
        occupied = vram_used_mb()
        row: dict = {
            "ballast_requested_mb": mib,
            "ballast_granted_mb": b.granted,
            "vram_in_use_before_load_mb": occupied,
            "vram_free_before_load_mb": max(total_mb - occupied, 0),
        }

        t = time.time()
        try:
            r = api("/api/generate",
                    {"model": model, "prompt": PROMPT, "stream": False,
                     "keep_alive": "5m", "options": {"num_predict": 96}},
                    timeout=600)
        except Exception as e:                      # a refused load is a result
            row.update({"loaded": False, "error": f"{type(e).__name__}: {e}",
                        "ballast_survived": b.alive()})
            evict(model)
            return row

        elapsed = round(time.time() - t, 2)
        eval_ct, eval_ns = r.get("eval_count", 0), r.get("eval_duration", 0)
        row.update({
            "loaded": True,
            "total_seconds": elapsed,
            "load_seconds": round(r.get("load_duration", 0) / 1e9, 2),
            "tokens_generated": eval_ct,
            "tokens_per_second": round(eval_ct / (eval_ns / 1e9), 2) if eval_ns else None,
            "ollama_vram_mb": ollama_vram_mb(),
            "processor_split": ps_split(model),
            "ballast_survived": b.alive(),
        })
        evict(model)
        return row


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--steps", default="0,1024,2048,3072,4096,5120",
                    help="MiB of VRAM the co-tenant holds, comma separated")
    ap.add_argument("--out", default=None)
    a = ap.parse_args()

    host = host_info()
    total = int(host["vram_total_mb"] or 0)
    python = torch_python()
    steps = [int(s) for s in a.steps.split(",") if s.strip()]

    others = {k: v for k, v in processes().items() if "ollama" not in k.lower()}
    print(f"\n{host['gpu']}  {total} MiB total   model={a.model}")
    print(f"ballast interpreter: {python}")
    if others:
        print(f"other processes already on the card: {others}")
    print()

    rows = []
    for mib in steps:
        print(f"  co-tenant holding {mib:>5} MiB ...", end="", flush=True)
        row = one_step(a.model, python, mib, total)
        rows.append(row)
        if row.get("loaded"):
            print(f" {row['processor_split']:>16}  "
                  f"{row['tokens_per_second']} tok/s  "
                  f"load {row['load_seconds']}s")
        else:
            print(f" FAILED: {row.get('error')}")

    data = {"host": host, "model": a.model, "prompt": PROMPT,
            "other_processes_at_start": others,
            "measured_at": datetime.now(timezone.utc).isoformat(),
            "steps": rows}
    slug = (host["gpu"] or "gpu").lower().replace(" ", "-").replace("/", "-")
    out = Path(a.out or
               f"results/contention--{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}")


if __name__ == "__main__":
    main()
