#!/usr/bin/env python3
"""DaniiT GPU inference benchmark — Ollama. Dependencies: Python 3.10+, requests.
Usage: python3 bench.py [--models m1,m2] [--tests 1,2,3] [--host http://localhost:11434]
"""
import argparse, base64, datetime, html.parser, json, os, re, shutil, socket, statistics
import subprocess, sys, tempfile, threading, time, urllib.request
import requests

HERE = os.path.dirname(os.path.abspath(__file__))
MODELS = ["qwen2.5vl:7b", "gemma3:12b", "qwen2.5-coder:14b", "qwen3:14b"]
VISION = {"qwen2.5vl:7b", "gemma3:12b"}
OPTIONS = {"temperature": 0, "seed": 42, "num_ctx": 8192, "num_predict": 6144}
WARM_RUNS = 3
EXPECTED1 = {"animal": "cat", "count": 1, "eye_color": "green", "coat_pattern": "tabby",
             "posture": "sitting", "background_sky": True}

def prompt(name):
    with open(os.path.join(HERE, "prompts", name), encoding="utf-8") as f:
        return f.read()

# ---------- nvidia-smi sampling ----------
class GpuSampler(threading.Thread):
    def __init__(self):
        super().__init__(daemon=True); self.samples = []; self.stop = threading.Event()
        self.ok = shutil.which("nvidia-smi") is not None
    def run(self):
        while self.ok and not self.stop.is_set():
            try:
                out = subprocess.run(["nvidia-smi", "--query-gpu=memory.used,power.draw,utilization.gpu",
                                      "--format=csv,noheader,nounits"], capture_output=True, text=True, timeout=5)
                if out.returncode != 0: self.ok = False; break
                m, p, u = [x.strip() for x in out.stdout.splitlines()[0].split(",")]
                self.samples.append({"mem_mib": float(m), "power_w": float(p), "util": float(u)})
            except Exception:
                self.ok = False; break
            self.stop.wait(0.5)
    def summary(self):
        if not self.samples: return None
        return {"vram_max_mib": max(s["mem_mib"] for s in self.samples),
                "power_max_w": max(s["power_w"] for s in self.samples),
                "util_avg": round(statistics.mean(s["util"] for s in self.samples), 1)}

def gpu_info():
    try:
        out = subprocess.run(["nvidia-smi", "--query-gpu=name,memory.total,driver_version",
                              "--format=csv,noheader"], capture_output=True, text=True, timeout=10)
        if out.returncode == 0: return out.stdout.strip()
        return "nvidia-smi error: " + (out.stdout + out.stderr).strip()[:200]
    except Exception as e:
        return f"nvidia-smi unavailable: {e}"

# ---------- Ollama ----------
def unload(host, model):
    requests.post(f"{host}/api/generate", json={"model": model, "keep_alive": 0}, timeout=120)
    time.sleep(2)

def generate(host, model, text, images=None):
    body = {"model": model, "prompt": text, "stream": True, "options": OPTIONS, "keep_alive": "10m"}
    if images: body["images"] = images
    sampler = GpuSampler(); sampler.start()
    t0 = time.perf_counter(); ttft = None; chunks = []; final = {}
    with requests.post(f"{host}/api/generate", json=body, stream=True, timeout=3600) as r:
        r.raise_for_status()
        for line in r.iter_lines():
            if not line: continue
            d = json.loads(line)
            if d.get("response"):
                if ttft is None: ttft = time.perf_counter() - t0
                chunks.append(d["response"])
            if d.get("done"): final = d
    wall = time.perf_counter() - t0
    sampler.stop.set(); sampler.join(2)
    ns = 1e9
    m = {"ttft_s": round(ttft, 3) if ttft else None, "wall_s": round(wall, 3),
         "total_s": final.get("total_duration", 0) / ns, "load_s": final.get("load_duration", 0) / ns,
         "prompt_tokens": final.get("prompt_eval_count"), "gen_tokens": final.get("eval_count"),
         "prompt_tps": (final.get("prompt_eval_count") or 0) / (final.get("prompt_eval_duration") or 1) * ns,
         "gen_tps": (final.get("eval_count") or 0) / (final.get("eval_duration") or 1) * ns,
         "gpu": sampler.summary()}
    return "".join(chunks), m

# ---------- extraction & checks ----------
def strip_think(text):
    return re.sub(r"<think>.*?</think>", "", text, flags=re.S)

FILE_RE = re.compile(r"FILE:\s*`?([\w./-]+)`?\s*\n```[\w+-]*\n(.*?)\n```", re.S)
def extract_files(text):
    return {p.strip().lstrip("./"): c for p, c in FILE_RE.findall(strip_think(text))}

def check_test1(text):
    t = strip_think(text); m = re.search(r"\{.*\}", t, re.S)
    if not m: return {"score": 0, "max": len(EXPECTED1), "error": "no JSON"}
    try: d = json.loads(m.group(0))
    except Exception as e: return {"score": 0, "max": len(EXPECTED1), "error": f"bad JSON: {e}"}
    ok = {}
    for k, v in EXPECTED1.items():
        got = d.get(k)
        ok[k] = (str(got).lower().strip() == str(v).lower()) if not isinstance(v, bool) else got is v
    return {"score": sum(ok.values()), "max": len(EXPECTED1), "fields": ok, "answer": d}

class _Collector(html.parser.HTMLParser):
    def __init__(self):
        super().__init__(); self.tags = []; self.ids = set(); self.attrs = []; self.text = []
    def handle_starttag(self, tag, attrs):
        self.tags.append(tag); a = dict(attrs); self.attrs.append((tag, a))
        if "id" in a: self.ids.add(a["id"])
    def handle_data(self, d): self.text.append(d)

def check_test2(files):
    res = {}
    h = files.get("index.html"); c = files.get("style.css")
    res["index.html exists"] = h is not None; res["style.css exists"] = c is not None
    if h:
        p = _Collector()
        try: p.feed(h); p.close(); res["html parseable"] = True
        except Exception: res["html parseable"] = False
        txt = " ".join(p.text)
        res["doctype"] = h.lstrip().lower().startswith("<!doctype html>")
        res["title Khleb & Co"] = bool(re.search(r"<title>[^<]*Khleb (&amp;|&) Co", h))
        res["stylesheet link"] = any(t == "link" and a.get("href") == "style.css" for t, a in p.attrs)
        for t in ("header", "nav", "h1", "footer", "table"): res[f"<{t}>"] = t in p.tags
        for i in ("about", "menu", "contact"): res[f"#{i}"] = i in p.ids
        res[">=3 sections"] = p.tags.count("section") >= 3
        res[">=4 menu rows"] = p.tags.count("tr") >= 5 or (p.tags.count("tr") >= 4 and "thead" not in p.tags)
        res["BYN prices"] = "BYN" in txt
        res["address"] = "Nezavisimosti" in txt
        res["footer text"] = "2026 Khleb" in txt
        res["no script"] = "script" not in p.tags
    if c:
        res["@media 600px"] = bool(re.search(r"@media[^{]*max-width:\s*600px", c))
        res["css braces balanced"] = c.count("{") == c.count("}")
    return {"score": sum(res.values()), "max": len(res), "checks": res}

DJANGO_FILES = ["manage.py", "benchsite/settings.py", "benchsite/urls.py", "pages/views.py",
                "pages/urls.py", "pages/forms.py", "pages/templates/pages/base.html",
                "pages/templates/pages/home.html", "pages/templates/pages/contact.html"]

def check_test3(files, port=8765):
    res = {f"{f} exists": f in files for f in DJANGO_FILES}
    d = tempfile.mkdtemp(prefix="benchsite-")
    for path, content in files.items():
        if ".." in path or path.startswith("/"): continue
        full = os.path.join(d, path); os.makedirs(os.path.dirname(full), exist_ok=True)
        open(full, "w", encoding="utf-8").write(content)
    for pkg in ("benchsite", "pages"):
        f = os.path.join(d, pkg, "__init__.py")
        if os.path.isdir(os.path.dirname(f)) and not os.path.exists(f): open(f, "w").close()
    py = sys.executable
    try: import django  # noqa
    except ImportError:
        res["django installed"] = False
        return {"score": sum(res.values()), "max": len(res) + 4, "checks": res, "dir": d}
    chk = subprocess.run([py, "manage.py", "check"], cwd=d, capture_output=True, text=True, timeout=60)
    res["manage.py check"] = chk.returncode == 0
    srv = subprocess.Popen([py, "manage.py", "runserver", f"127.0.0.1:{port}", "--noreload"], cwd=d,
                           stdout=subprocess.DEVNULL, stderr=subprocess.PIPE, text=True)
    try:
        for _ in range(30):
            try: socket.create_connection(("127.0.0.1", port), 0.5).close(); break
            except OSError: time.sleep(0.5)
        def get(u):
            try:
                with urllib.request.urlopen(f"http://127.0.0.1:{port}{u}", timeout=10) as r:
                    return r.status, r.read().decode("utf-8", "replace")
            except Exception as e: return getattr(e, "code", 0), ""
        s1, b1 = get("/"); s2, b2 = get("/contact/")
        res["GET / 200"] = s1 == 200; res["home h1"] = "Welcome to BenchSite" in b1
        res["GET /contact/ 200"] = s2 == 200; res["contact form"] = "<form" in b2 and "csrfmiddlewaretoken" in b2
    finally:
        srv.terminate(); srv.wait(10)
    shutil.rmtree(d, ignore_errors=True)
    return {"score": sum(res.values()), "max": len(res), "checks": res}

# ---------- protocol ----------
TESTS = {
    1: ("vision", "test1_vision.txt", True),
    2: ("static_site", "test2_static_site.txt", False),
    3: ("django", "test3_django.txt", False),
}

def run_test(host, model, tid, outdir):
    name, pfile, needs_img = TESTS[tid]
    text = prompt(pfile)
    imgs = None
    if needs_img:
        with open(os.path.join(HERE, "test1.jpg"), "rb") as f:
            imgs = [base64.b64encode(f.read()).decode()]
    unload(host, model)
    runs = []
    for i in range(1 + WARM_RUNS):
        out, m = generate(host, model, text, imgs)
        m["cold"] = i == 0; runs.append(m)
        print(f"  {model} test{tid} run{i} gen={m['gen_tps']:.1f} tok/s ttft={m['ttft_s']}s", flush=True)
    with open(os.path.join(outdir, f"{model.replace(':', '_').replace('/', '_')}-test{tid}.txt"), "w") as f:
        f.write(out)
    if tid == 1: q = check_test1(out)
    elif tid == 2: q = check_test2(extract_files(out))
    else: q = check_test3(extract_files(out))
    warm = runs[1:]
    med = lambda k: statistics.median([r[k] for r in warm if r[k] is not None]) if warm else None
    gpus = [r["gpu"] for r in runs if r["gpu"]]
    return {"model": model, "test": tid, "name": name, "cold_load_s": runs[0]["load_s"],
            "median": {k: med(k) for k in ("gen_tps", "prompt_tps", "ttft_s", "total_s", "gen_tokens")},
            "vram_max_mib": max((g["vram_max_mib"] for g in gpus), default=None),
            "power_max_w": max((g["power_max_w"] for g in gpus), default=None),
            "quality": q, "runs": runs}

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--host", default=os.environ.get("OLLAMA_HOST", "http://localhost:11434"))
    ap.add_argument("--models", default=",".join(MODELS))
    ap.add_argument("--tests", default="1,2,3")
    a = ap.parse_args()
    host = a.host if a.host.startswith("http") else "http://" + a.host
    hn = socket.gethostname(); date = datetime.date.today().isoformat()
    outdir = os.path.join(HERE, "results"); os.makedirs(outdir, exist_ok=True)
    raw = os.path.join(outdir, f"raw-{hn}-{date}"); os.makedirs(raw, exist_ok=True)
    ver = requests.get(f"{host}/api/version", timeout=10).json().get("version")
    data = {"host": hn, "date": date, "gpu": gpu_info(), "ollama": ver, "options": OPTIONS,
            "warm_runs": WARM_RUNS, "results": []}
    for model in a.models.split(","):
        for tid in map(int, a.tests.split(",")):
            if tid == 1 and model not in VISION: continue
            try: data["results"].append(run_test(host, model, tid, raw))
            except Exception as e:
                print(f"  ERROR {model} test{tid}: {e}", flush=True)
                data["results"].append({"model": model, "test": tid, "error": str(e)})
    path = os.path.join(outdir, f"{hn}-{date}.json")
    with open(path, "w") as f: json.dump(data, f, indent=1)
    print("Results:", path)

if __name__ == "__main__":
    main()
