#!/usr/bin/env python3
"""ArcadeBench v0 independent rescorer.

Reads barkade.admin.match-export files (schemaVersion 1 or 2) and recomputes,
from the raw record alone:

  - exact minimax value and distance for every legal move at every decision
  - Move Regret = V*(s) - Q(s,a), plus optimality, catastrophic status,
    difficulty, criticality, chosen rank, tempo regret
  - per-player Efficiency Rating = 1 - (joules on non-optimal / total joules)

Capacity is read from the gateway receipt in microjoules, which is the only
field normalized to the approved $1 = 3.6 MJ policy. responsePayload.settledJoules
carries a legacy per-request label and is deliberately NOT used.

Usage: python3 rescore.py FILE [FILE ...]
"""
import json, sys
from functools import lru_cache

LINES = [(0,1,2),(3,4,5),(6,7,8),(0,3,6),(1,4,7),(2,5,8),(0,4,8),(2,4,6)]
OTHER = {"X": "O", "O": "X"}


def winner(b):
    for i, j, k in LINES:
        if b[i] and b[i] == b[j] == b[k]:
            return b[i]
    return None


@lru_cache(maxsize=None)
def negamax(board, p):
    """(value, plies_to_terminal) for player p to move, under optimal play."""
    b = list(board)
    if not any(v is None for v in b):
        return (0, 0)
    best = None
    for a in range(9):
        if b[a] is not None:
            continue
        v, d = value_after(board, a, p)
        cand = (v, d)
        if best is None or better(cand, best):
            best = cand
    return best


def value_after(board, a, p):
    """(value, plies) for p after playing a, p's perspective."""
    b = list(board)
    b[a] = p
    if winner(b) == p:
        return (1, 1)
    if not any(v is None for v in b):
        return (0, 1)
    v, d = negamax(tuple(b), OTHER[p])
    return (-v, d + 1)


def better(x, y):
    """Lexicographic: higher value; win sooner; lose later."""
    if x[0] != y[0]:
        return x[0] > y[0]
    if x[0] > 0:
        return x[1] < y[1]
    if x[0] < 0:
        return x[1] > y[1]
    return x[1] < y[1]


def joules(usage):
    """Policy-normalized joules from the gateway receipt."""
    if usage.get("actualCostMicrojoules") is not None:
        return int(usage["actualCostMicrojoules"]) / 1e6
    rp = usage.get("responsePayload") or {}
    gr = rp.get("gatewayReceipt") or {}
    if gr.get("totalMicrojoules") is not None:
        return int(gr["totalMicrojoules"]) / 1e6
    if gr.get("providerCostUsd") is not None:
        return float(gr["providerCostUsd"]) * 3_600_000
    return None


def score_match(match):
    snaps = {s["stateVersion"]: s["stateJson"] for s in match["snapshots"]}
    usage_by_key = {}
    for u in match["aiMoveUsage"]:
        usage_by_key[u["idempotencyKey"]] = u

    decisions = []
    for mv in match["moves"]:
        u = usage_by_key.get(mv["idempotencyKey"])
        if u is None:
            continue
        pre = snaps.get(mv["stateVersion"] - 1)
        if pre is None:
            continue
        board = tuple(pre["board"])
        p = pre["nextPlayer"]
        chosen = mv["moveJson"]["index"]
        fallback = bool(u.get("fallbackMove"))

        legal = [a for a in range(9) if board[a] is None]
        vals = {a: value_after(board, a, p) for a in legal}
        best = max(vals.values(), key=lambda v: (v[0], -v[1] if v[0] > 0 else v[1]))
        vstar = max(v[0] for v in vals.values())
        qv, qd = vals[chosen]
        n_opt = sum(1 for v in vals.values() if v[0] == vstar)
        ordered = sorted(vals.values(), key=lambda v: (-v[0],))
        second = ordered[1][0] if len(ordered) > 1 else None
        opt_dists = [v[1] for v in vals.values() if v[0] == vstar]
        best_dist = min(opt_dists) if vstar > 0 else max(opt_dists)

        decisions.append({
            "seat": mv["seatIndex"],
            "model": u["modelId"],
            "player": p,
            "chosen": chosen,
            "fallback": fallback,
            "fallback_reason": u.get("fallbackReason"),
            "v_star": vstar,
            "q": qv,
            "regret": vstar - qv,
            "is_optimal": qv == vstar,
            "catastrophic": (vstar - qv) == 2,
            "legal_moves": len(legal),
            "optimal_moves": n_opt,
            "difficulty": round(1 - n_opt / len(legal), 3),
            "criticality": (vstar - second) if second is not None else None,
            "chosen_rank": sorted({v[0] for v in vals.values()}, reverse=True).index(qv) + 1,
            "tempo_regret": (qd - best_dist) if qv == vstar else None,
            "joules": joules(u),
            "prompt_tokens": u.get("promptTokens"),
            "completion_tokens": u.get("completionTokens"),
            "distribution": {a: vals[a] for a in legal},
        })
    return decisions


def aggregate(decisions):
    out = {}
    for d in decisions:
        if d["fallback"]:
            continue
        m = out.setdefault(d["model"], {
            "n": 0, "opt": 0, "regret": 0.0, "cat": 0,
            "j_total": 0.0, "j_wasted": 0.0, "tokens": 0,
        })
        m["n"] += 1
        m["opt"] += 1 if d["is_optimal"] else 0
        m["regret"] += d["regret"]
        m["cat"] += 1 if d["catastrophic"] else 0
        j = d["joules"] or 0.0
        m["j_total"] += j
        if not d["is_optimal"]:
            m["j_wasted"] += j
        m["tokens"] += (d["prompt_tokens"] or 0) + (d["completion_tokens"] or 0)
    for m in out.values():
        m["opt_rate"] = m["opt"] / m["n"] if m["n"] else 0
        m["avg_regret"] = m["regret"] / m["n"] if m["n"] else 0
        m["cat_rate"] = m["cat"] / m["n"] if m["n"] else 0
        m["efficiency"] = (1 - m["j_wasted"] / m["j_total"]) if m["j_total"] else None
    return out


def main(paths):
    all_dec = []
    for path in paths:
        doc = json.load(open(path))
        for match in doc["matches"]:
            dec = score_match(match)
            all_dec.extend(dec)
            models = sorted({d["model"] for d in dec})
            st = match["state"]
            print(f"\n=== {match['sourceMatchId'][:8]}  {match['game']['id']}  "
                  f"{match['status']}/{st.get('terminalReason')}")
            for d in dec:
                seat = "X" if d["player"] == "X" else "O"
                tag = "FALLBACK " if d["fallback"] else ""
                flag = "" if d["is_optimal"] else ("  <-- CATASTROPHIC" if d["catastrophic"] else "  <-- regret")
                jd = f"{d['joules']:.1f} J" if d["joules"] is not None else "n/a"
                print(f"  {tag}{seat} {d['model']:<28} idx={d['chosen']}  "
                      f"V*={d['v_star']:+d} Q={d['q']:+d} reg={d['regret']}  "
                      f"diff={d['difficulty']:.2f} crit={d['criticality']}  {jd}{flag}")
            agg = aggregate(dec)
            for model, m in sorted(agg.items()):
                print(f"  -> {model:<28} N={m['n']} opt={m['opt_rate']:.1%} "
                      f"avgReg={m['avg_regret']:.3f} cat={m['cat_rate']:.1%} "
                      f"J={m['j_total']:.1f} eff={m['efficiency']:.4f}")

    print("\n\n===== AGGREGATE ACROSS ALL MATCHES =====")
    nf = sum(1 for d in all_dec if d["fallback"])
    scored = [d for d in all_dec if not d["fallback"]]
    print(f"decisions total={len(all_dec)}  scored={len(scored)}  fallbacks={nf}")
    if scored:
        print(f"optimal rate={sum(d['is_optimal'] for d in scored)/len(scored):.1%}  "
              f"avg regret={sum(d['regret'] for d in scored)/len(scored):.3f}  "
              f"catastrophic={sum(d['catastrophic'] for d in scored)/len(scored):.1%}")
        tj = sum(d["joules"] or 0 for d in scored)
        wj = sum((d["joules"] or 0) for d in scored if not d["is_optimal"])
        print(f"total joules={tj:.1f} J  ({tj/1000:.3f} kJ)  "
              f"mean per decision={tj/len(scored):.1f} J")
        print(f"joules on non-optimal={wj:.1f} J  "
              f"overall efficiency={1-wj/tj:.4f}")
    print()
    agg = aggregate(all_dec)
    hdr = f"{'model':<28} {'N':>3} {'opt':>7} {'avgReg':>7} {'cat':>6} {'total J':>9} {'J/dec':>8} {'eff':>7}"
    print(hdr); print("-" * len(hdr))
    for model, m in sorted(agg.items(), key=lambda kv: -kv[1]["efficiency"]):
        print(f"{model:<28} {m['n']:>3} {m['opt_rate']:>6.1%} {m['avg_regret']:>7.3f} "
              f"{m['cat_rate']:>5.1%} {m['j_total']:>9.1f} "
              f"{m['j_total']/m['n']:>8.1f} {m['efficiency']:>7.4f}")


if __name__ == "__main__":
    main(sys.argv[1:])
