#!/usr/bin/env python3
"""Time to complete: wall-clock time to a finished answer on a fixed question set, with an output cap high enough not to
clip reasoning (32,768 tokens). Companion to the RigMark receipts in this folder.

Questions: ttc-questions.json (30 MMLU-Pro + 30 BIG-Bench Hard, fixed selection). Requests use the given reasoning
effort and Z.AI's recommended sampling (temperature 1.0, top_p 0.95) with a fixed seed per question, streamed, at the
given concurrency. Per question it records end-to-end latency, time to first token, completion and reasoning tokens,
finish reason and correctness (last \\boxed{...} answer, as below). Python standard library only.

usage: python3 ttc.py --base-url http://HOST:8888 --model NAME --label L --out OUT.json [--concurrency 1] [--effort max]
"""
import argparse
import json
import os
import re
import statistics as st
import time
import urllib.request
from concurrent.futures import ThreadPoolExecutor

HERE = os.path.dirname(os.path.abspath(__file__))


def last_boxed(s):
    i = (s or "").rfind("\\boxed")
    if i < 0:
        return None
    j = s.find("{", i)
    if j < 0:
        return None
    depth, k = 0, j
    while k < len(s):
        if s[k] == "{":
            depth += 1
        elif s[k] == "}":
            depth -= 1
            if depth == 0:
                return s[j + 1:k]
        k += 1
    return None


def norm_bbh(x):
    if x is None:
        return None
    x = re.sub(r"\\text\{([^}]*)\}", r"\1", x).strip().strip(".").strip()
    m = re.fullmatch(r"\(?([A-R])\)?(?:[.:)]\s.*)?", x, re.S)
    if m:
        return m.group(1)
    return re.sub(r"\s+", " ", x).lower()


def letter(s):
    if s is None:
        return None
    s = re.sub(r"\\text\{([^}]*)\}", r"\1", s).strip().strip("()").strip()
    m = re.match(r"([A-Ja-j])(?![A-Za-z])", s)
    return m.group(1).upper() if m else None


def grade(row, content, reasoning):
    got = last_boxed(content) or last_boxed(reasoning)
    if row["suite"] == "bbh":
        return norm_bbh(got) == norm_bbh(row["answer"])
    return letter(got) == row["answer"]


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--base-url", required=True); ap.add_argument("--model", required=True)
    ap.add_argument("--label", required=True); ap.add_argument("--out", required=True)
    ap.add_argument("--questions", default=os.path.join(HERE, "ttc-questions.json"))
    ap.add_argument("--concurrency", type=int, default=1)
    ap.add_argument("--effort", default="max"); ap.add_argument("--max-tokens", type=int, default=32768)
    a = ap.parse_args()
    qs = json.load(open(a.questions))

    def ask(item):
        i, r = item
        body = {"model": a.model, "messages": [{"role": "user", "content": r["prompt"]}], "max_tokens": a.max_tokens,
                "temperature": 1.0, "top_p": 0.95, "seed": 1000 + i, "stream": True,
                "stream_options": {"include_usage": True}, "chat_template_kwargs": {"reasoning_effort": a.effort}}
        req = urllib.request.Request(f"{a.base_url}/v1/chat/completions", data=json.dumps(body).encode(),
                                     headers={"Content-Type": "application/json"})
        t0 = time.time(); ttft = None; text = ""; reas = ""; usage = {}; finish = None
        with urllib.request.urlopen(req, timeout=7200) as resp:
            for line in resp:
                line = line.decode().strip()
                if not line.startswith("data: ") or line == "data: [DONE]":
                    continue
                d = json.loads(line[6:])
                if d.get("usage"):
                    usage = d["usage"]
                for c in d.get("choices", []):
                    delta = c.get("delta") or {}
                    piece = (delta.get("content") or "") + (delta.get("reasoning_content") or delta.get("reasoning") or "")
                    if piece and ttft is None:
                        ttft = time.time() - t0
                    text += delta.get("content") or ""; reas += delta.get("reasoning_content") or delta.get("reasoning") or ""
                    finish = c.get("finish_reason") or finish
        return {"id": r.get("id", i), "suite": r["suite"], "latency_s": time.time() - t0, "ttft_s": ttft,
                "completion_tokens": usage.get("completion_tokens"), "finish": finish,
                "reasoning_tokens": (usage.get("completion_tokens_details") or {}).get("reasoning_tokens"),
                "correct": bool(grade(r, text, reas))}

    t0 = time.time()
    with ThreadPoolExecutor(a.concurrency) as ex:
        out = list(ex.map(ask, enumerate(qs)))
    wall = time.time() - t0
    lat = [x["latency_s"] for x in out]
    res = {"label": a.label, "model": a.model, "concurrency": a.concurrency, "effort": a.effort, "n": len(out),
           "wall_s": wall, "latency_median_s": st.median(lat), "latency_mean_s": st.mean(lat),
           "latency_p90_s": sorted(lat)[int(0.9 * len(lat)) - 1], "correct": sum(x["correct"] for x in out),
           "completion_tokens_mean": st.mean(x["completion_tokens"] or 0 for x in out),
           "length_stops": sum(x["finish"] == "length" for x in out), "items": out}
    json.dump(res, open(a.out, "w"), indent=1)
    print(f"TTC {a.label} c{a.concurrency}: {res['correct']}/{len(out)} correct, median {res['latency_median_s']:.1f}s, "
          f"mean {res['latency_mean_s']:.1f}s, wall {wall / 60:.1f} min, tokens {res['completion_tokens_mean']:.0f}, "
          f"length stops {res['length_stops']}")


if __name__ == "__main__":
    main()
