#!/usr/bin/env python3
"""Run Router Gauntlet: every contestant x every task x reps. Resumable (results.jsonl), hard spend cap."""
import json, sys, threading
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent))
from or_call import call
from grade import grade

HERE = Path(__file__).parent
CAP = float(sys.argv[1]) if len(sys.argv) > 1 else 12.0
ROUTERS = ["typesafe/jev-router", "nvidia/switchyard", "unbiased/pareto", "openrouter/auto"]
BASELINES = ["deepseek/deepseek-v4.1-flash", "google/gemini-3.8-flash", "openai/gpt-6.1-sol", "anthropic/claude-opus-5.5"]
REPS = {m: 3 for m in ROUTERS} | {m: 1 for m in BASELINES}
MAXTOK = {"easy": 3000, "medium": 12000, "hard": 16000, "extreme": 32000}

tasks = json.load(open(HERE / "tasks.json"))["tasks"]
out = HERE / "results.jsonl"
done, spent = set(), 0.0
if out.exists():
    for l in out.read_text().splitlines():
        r = json.loads(l); done.add((r["model"], r["task"], r["rep"])); spent += r.get("cost") or 0
lock = threading.Lock()
jobs = [(m, t, r) for m in REPS for r in range(REPS[m]) for t in tasks if (m, t["id"], r) not in done]
print(f"{len(jobs)} calls to run, ${spent:.3f} already spent, cap ${CAP}")

def work(job):
    global spent
    m, t, rep = job
    with lock:
        if spent >= CAP: return
    res = call(m, t["prompt"], max_tokens=MAXTOK[t["tier"]], timeout=600)
    ok, detail = (False, res["error"]) if "error" in res else grade(t, res["text"])
    rec = {"model": m, "task": t["id"], "tier": t["tier"], "rep": rep, "pass": ok, "detail": detail,
           **{k: res.get(k) for k in ("routed", "provider", "cost", "in_tok", "out_tok", "latency", "id")},
           "text": (res.get("text") or "")[-4000:]}
    with lock:
        spent += res.get("cost") or 0
        with open(out, "a") as f: f.write(json.dumps(rec) + "\n")
        print(f"{'PASS' if ok else 'fail'} {m:32} {t['id']} r{rep} -> {str(res.get('routed')):34} ${res.get('cost') or 0:.4f} {res.get('latency')}s  {str(detail)[:50]}  [total ${spent:.2f}]", flush=True)

with ThreadPoolExecutor(10) as ex: list(ex.map(work, jobs))
print(f"done. spent ${spent:.3f}")
