#!/usr/bin/env python3
"""Summarise per-task runs: medians over reps, speedups vs mlx-lm, identical-output check.

Usage: analyze.py <results-root>   (expects <root>/<model>/<config>/rep*.json)
Writes <root>/summary.json and prints Markdown tables.
"""
import glob
import json
import math
import os
import statistics as st
import sys

ROOT = sys.argv[1]
TASKS = {t["id"]: t for t in json.load(open(os.path.join(os.path.dirname(os.path.abspath(__file__)), "tasks.json")))["tasks"]}
ORDER = list(TASKS)


def med(xs):
    xs = [x for x in xs if x is not None]
    return st.median(xs) if xs else None


def gmean(xs):
    xs = [x for x in xs if x]
    return math.exp(sum(math.log(x) for x in xs) / len(xs)) if xs else None


def first_diff(a, b):
    n = min(len(a), len(b))
    for i in range(n):
        if a[i] != b[i]:
            return i
    return None if len(a) == len(b) else n


summary = {}
for mdir in sorted(glob.glob(os.path.join(ROOT, "*"))):
    if not os.path.isdir(mdir):
        continue
    model = os.path.basename(mdir)
    cfgs = {}
    for cdir in sorted(glob.glob(os.path.join(mdir, "*"))):
        reps = sorted(glob.glob(os.path.join(cdir, "rep*.json")))
        if not reps:
            continue
        runs = [json.load(open(r)) for r in reps]
        per = {}
        for tid in ORDER:
            rs = [next((x for x in run["results"] if x["id"] == tid), None) for run in runs]
            rs = [x for x in rs if x and not x.get("error")]
            if not rs:
                continue
            per[tid] = {
                "n": len(rs),
                "decode_tps": med([x["decode_tps"] for x in rs]),
                "ttft_s": med([x["ttft_s"] for x in rs]),
                "e2e_s": med([x["e2e_s"] for x in rs]),
                "completion_tokens": med([x["completion_tokens"] for x in rs]),
                "prompt_tokens": med([x["prompt_tokens"] for x in rs]),
                "texts": [x["text"] for x in rs],
                "self_consistent": len({x["text"] for x in rs}) == 1,
                "decode_tps_runs": [x["decode_tps"] for x in rs],
                "e2e_runs": [x["e2e_s"] for x in rs],
            }
        cfgs[os.path.basename(cdir)] = per
    if "mlxlm" not in cfgs:
        continue
    base = cfgs["mlxlm"]
    msum = {"configs": {}}
    for cfg, per in cfgs.items():
        if cfg == "mlxlm":
            continue
        rows = []
        for tid in ORDER:
            if tid not in per or tid not in base:
                continue
            b, r = base[tid], per[tid]
            dspd = (r["decode_tps"] / b["decode_tps"]) if (r["decode_tps"] and b["decode_tps"]) else None
            espd = b["e2e_s"] / r["e2e_s"]
            ta, tb = b["texts"][0], r["texts"][0]
            fd = first_diff(ta, tb)
            rows.append({
                "id": tid, "category": TASKS[tid]["category"], "copy_heavy": TASKS[tid].get("copy_heavy", False),
                "mlx_tps": b["decode_tps"], "rapid_tps": r["decode_tps"], "decode_speedup": dspd,
                "mlx_ttft": b["ttft_s"], "rapid_ttft": r["ttft_s"],
                "mlx_e2e": b["e2e_s"], "rapid_e2e": r["e2e_s"], "e2e_speedup": espd,
                "mlx_ctoks": b["completion_tokens"], "rapid_ctoks": r["completion_tokens"],
                "mlx_ptoks": b["prompt_tokens"], "rapid_ptoks": r["prompt_tokens"],
                "identical": fd is None, "first_diff_char": fd, "len_mlx": len(ta), "len_rapid": len(tb),
                "mlx_self_consistent": b["self_consistent"], "rapid_self_consistent": r["self_consistent"],
                "n_mlx": b["n"], "n_rapid": r["n"],
            })
        ds = [x["decode_speedup"] for x in rows if x["decode_speedup"]]
        es = [x["e2e_speedup"] for x in rows]
        msum["configs"][cfg] = {
            "rows": rows,
            "decode": {"max": max(ds) if ds else None, "median": med(ds), "geomean": gmean(ds), "min": min(ds) if ds else None, "n": len(ds)},
            "e2e": {"max": max(es), "median": med(es), "geomean": gmean(es), "min": min(es), "n": len(es)},
            "identical": sum(x["identical"] for x in rows), "tasks": len(rows),
        }
    summary[model] = msum

json.dump(summary, open(os.path.join(ROOT, "summary.json"), "w"), indent=1)


def f(x, p=1):
    return "—" if x is None else f"{x:.{p}f}"


for model, ms in summary.items():
    for cfg, cs in ms["configs"].items():
        print(f"\n### {model} — {cfg} vs mlx-lm\n")
        print("| Task | mlx-lm tok/s | Rapid tok/s | decode × | TTFT mlx / Rapid (s) | e2e mlx / Rapid (s) | e2e × | tokens | identical |")
        print("|---|---:|---:|---:|---:|---:|---:|---:|:---:|")
        for x in cs["rows"]:
            ident = "yes" if x["identical"] else f"no (char {x['first_diff_char']})"
            toks = f"{x['mlx_ctoks']}" if x["mlx_ctoks"] == x["rapid_ctoks"] else f"{x['mlx_ctoks']}/{x['rapid_ctoks']}"
            print(f"| {x['category']} | {f(x['mlx_tps'])} | {f(x['rapid_tps'])} | {f(x['decode_speedup'],2)} | {f(x['mlx_ttft'],2)} / {f(x['rapid_ttft'],2)} | {f(x['mlx_e2e'],1)} / {f(x['rapid_e2e'],1)} | {f(x['e2e_speedup'],2)} | {toks} | {ident} |")
        d, e = cs["decode"], cs["e2e"]
        print(f"\nDecode ×: max {f(d['max'],2)}, median {f(d['median'],2)}, geomean {f(d['geomean'],2)}, min {f(d['min'],2)} (n={d['n']}). "
              f"End-to-end ×: max {f(e['max'],2)}, median {f(e['median'],2)}, geomean {f(e['geomean'],2)}, min {f(e['min'],2)} (n={e['n']}). "
              f"Identical output: {cs['identical']}/{cs['tasks']}.")
