"""PDB analysis. The unit of analysis is the item, not the call.

Replicates of one item under one arm are not independent observations, so they are
averaged inside the item before any test is applied. Paired comparisons across arms use
the exact two-sided sign test over item means; intervals are clustered by item.
"""
import json
import math
import os
import sys
from collections import defaultdict

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import views as V
import items as IT

R = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),
                                  "..", "results"))

ARM_LABEL = {
    "notes": "A notes", "notes_profile": "B notes + profile",
    "flat": "C flat records", "dimension": "D dimension",
    "dim_no_prov": "E dimension - provenance", "dim_no_time": "F dimension - validity",
}


def load(name):
    p = os.path.join(R, name)
    if not os.path.exists(p):
        return []
    rows = [json.loads(l) for l in open(p, encoding="utf-8")]
    if name == "grid.jsonl":
        import run as RUN
        by_id = {i["id"]: i for i in IT.I}
        for r in rows:
            r.update(RUN.score(by_id[r["item"]], r["arm"], r.get("obj")))
    return rows


def item_means(rows, arm, metric, gate, family=None):
    acc = defaultdict(list)
    for r in rows:
        if r["arm"] != arm or not r.get(gate):
            continue
        if family and r["family"] != family:
            continue
        acc[r["item"]].append(r[metric])
    return {k: sum(v) / len(v) for k, v in acc.items()}


def mean(d):
    return sum(d.values()) / len(d) if d else None


def cluster_ci(d):
    """Normal-approximation interval over item means, clustered by item."""
    n = len(d)
    if n < 2:
        return (None, None)
    vals = list(d.values())
    m = sum(vals) / n
    var = sum((v - m) ** 2 for v in vals) / (n - 1)
    se = math.sqrt(var / n)
    return (max(0.0, m - 1.96 * se), min(1.0, m + 1.96 * se))


def sign_test(a, b):
    keys = sorted(set(a) & set(b))
    pos = sum(1 for k in keys if b[k] > a[k])
    neg = sum(1 for k in keys if b[k] < a[k])
    n = pos + neg
    if n == 0:
        return 1.0, pos, neg
    k = min(pos, neg)
    p = min(1.0, 2 * sum(math.comb(n, i) for i in range(k + 1)) / 2 ** n)
    return round(p, 5), pos, neg


def pct(x):
    return "-" if x is None else f"{100 * x:5.1f}%"


def bar(title):
    print("\n" + "=" * 96)
    print(title)
    print("=" * 96)


def main():
    grid = load("grid.jsonl")
    if not grid:
        raise SystemExit("no grid.jsonl yet")

    bar(f"PDB grid: {len(grid)} calls, {len(set(r['item'] for r in grid))} items, "
        f"{len(set(r['arm'] for r in grid))} arms")
    print(f"parse success: {100*sum(r['parsed'] for r in grid)/len(grid):.1f}%")

    print(f"\n{'arm':26} {'M1 answer':>10} {'95% CI':>14} {'M2 source':>10} "
          f"{'M4 abstain':>11} {'M6 stale':>9} {'M7 overclaim':>13} {'M8 fabricated':>14}")
    print("-" * 112)
    summary = {}
    for arm in V.ARMS:
        m1 = item_means(grid, arm, "answer_correct", "scored_answer")
        lo, hi = cluster_ci(m1)
        m2 = item_means(grid, arm, "source_correct", "scored_source")
        m4 = item_means(grid, arm, "abstained", "scored_abst")
        m6 = item_means(grid, arm, "stale_correct", "scored_stale")
        m7 = item_means(grid, arm, "overclaim", "scored_answer")
        m8 = item_means(grid, arm, "fabricated", "scored_answer")
        summary[arm] = {"m1": mean(m1), "ci": (lo, hi), "m2": mean(m2), "m4": mean(m4),
                        "m6": mean(m6), "m7": mean(m7), "m8": mean(m8),
                        "items": len(m1)}
        ci = f"[{100*lo:.0f}-{100*hi:.0f}]" if lo is not None else "-"
        print(f"{ARM_LABEL[arm]:26} {pct(mean(m1)):>10} {ci:>14} {pct(mean(m2)):>10} "
              f"{pct(mean(m4)):>11} {pct(mean(m6)):>9} {pct(mean(m7)):>13} "
              f"{pct(mean(m8)):>14}")

    bar("M1 by family (item means)")
    fams = IT.FAMILIES
    print(f"{'arm':26} " + " ".join(f"{f:>9}" for f in fams))
    print("-" * 112)
    byfam = {}
    for arm in V.ARMS:
        row = []
        byfam[arm] = {}
        for f in fams:
            metric, gate = ("abstained", "scored_abst") if f == "GAP" else \
                           ("answer_correct", "scored_answer")
            d = item_means(grid, arm, metric, gate, family=f)
            byfam[arm][f] = mean(d)
            row.append(pct(mean(d)))
        print(f"{ARM_LABEL[arm]:26} " + " ".join(f"{v:>9}" for v in row))

    bar("Policy naming on the conflict family, and false stale alarms")
    print(f"{'arm':26} {'M3 policy named':>16} {'M5 false abstention':>21} "
          f"{'false stale alarm':>19}")
    print("-" * 112)
    extra = {}
    for arm in V.ARMS:
        m3 = item_means(grid, arm, "policy_correct", "scored_policy")
        m5 = item_means(grid, arm, "false_abstention", "scored_answer")
        fs = item_means(grid, arm, "false_stale", "scored_false_stale")
        extra[arm] = {"m3": mean(m3), "m5": mean(m5), "false_stale": mean(fs)}
        print(f"{ARM_LABEL[arm]:26} {pct(mean(m3)):>16} {pct(mean(m5)):>21} "
              f"{pct(mean(fs)):>19}")

    bar("Paired comparisons on M1 (exact two-sided sign test over item means)")
    pairs = [("notes", "dimension"), ("notes_profile", "dimension"),
             ("flat", "dimension"), ("dim_no_prov", "dimension"),
             ("dim_no_time", "dimension"), ("notes", "notes_profile"),
             ("notes", "flat"), ("flat", "dim_no_time"), ("flat", "dim_no_prov")]
    tests = {}
    for a, b in pairs:
        da, db = (item_means(grid, a, "answer_correct", "scored_answer"),
                  item_means(grid, b, "answer_correct", "scored_answer"))
        p, pos, neg = sign_test(da, db)
        tests[f"{a}_vs_{b}"] = p
        print(f"{ARM_LABEL[a]:26} -> {ARM_LABEL[b]:26} "
              f"{pct(mean(da))} -> {pct(mean(db))}  better:{pos:2} worse:{neg:2}  p={p}")

    bar("Targeted ablation contrasts (the family each field class is supposed to serve)")
    targeted = [("HIST", "dim_no_time", "dimension", "validity intervals"),
                ("CONF", "dim_no_prov", "dimension", "provenance and policy"),
                ("CURR", "dim_no_time", "dimension", "validity intervals"),
                ("CURR", "dim_no_prov", "dimension", "provenance and policy")]
    abl = {}
    for fam, a, b, what in targeted:
        metric, gate = ("answer_correct", "scored_answer")
        da = item_means(grid, a, metric, gate, family=fam)
        db = item_means(grid, b, metric, gate, family=fam)
        p, pos, neg = sign_test(da, db)
        abl[f"{fam}:{a}_vs_{b}"] = {"a": mean(da), "b": mean(db), "p": p}
        print(f"{fam:6} removing {what:24} {pct(mean(da))} vs {pct(mean(db))}  p={p}")

    # ---------------------------------------------------------------- tracks
    disc = load("disclosure.jsonl")
    disc_out = {}
    if disc:
        bar("Disclosure track: 36 decisions per arm per replicate")
        print(f"{'arm':20} {'precision':>10} {'recall':>9} {'accuracy':>10} "
              f"{'critical errors':>16}")
        print("-" * 112)
        for arm in ("no_policy", "prose_policy", "structured_policy"):
            s = [r for r in disc if r["arm"] == arm]
            tp = sum(r["tp"] for r in s); fp = sum(r["fp"] for r in s)
            tn = sum(r["tn"] for r in s); fn = sum(r["fn"] for r in s)
            crit = sum(r["critical"] for r in s)
            tot = tp + fp + tn + fn
            disc_out[arm] = {
                "precision": tp / (tp + fp) if tp + fp else None,
                "recall": tp / (tp + fn) if tp + fn else None,
                "accuracy": (tp + tn) / tot if tot else None,
                "critical": crit, "decisions": tot}
            print(f"{arm:20} {pct(disc_out[arm]['precision']):>10} "
                  f"{pct(disc_out[arm]['recall']):>9} {pct(disc_out[arm]['accuracy']):>10} "
                  f"{str(crit) + ' of ' + str(tot):>16}")
        print("\ncritical errors by requester")
        for rec in IT.RECIPIENTS:
            line = f"  {rec:14}"
            for arm in ("no_policy", "prose_policy", "structured_policy"):
                c = sum(r["critical"] for r in disc
                        if r["arm"] == arm and r["recipient"] == rec)
                line += f"  {arm}: {c}"
            print(line)

    wb = load("writeback.jsonl")
    wb_out = {}
    if wb:
        bar("Write-back track: record the change without destroying history")
        print(f"{'arm':20} {'new value':>10} {'valid_from':>11} {'closed id':>10} "
              f"{'valid_to':>9} {'history kept':>13} {'well formed':>12}")
        print("-" * 112)
        for arm in ("notes", "flat", "dimension"):
            s = [r for r in wb if r["arm"] == arm]
            n = len(s) or 1
            wb_out[arm] = {k: sum(r[k] for r in s) / n for k in
                           ("value_ok", "from_ok", "close_id_ok", "close_at_ok",
                            "history_ok", "well_formed")}
            wb_out[arm]["chains"] = len(s)
            w = wb_out[arm]
            print(f"{arm:20} {pct(w['value_ok']):>10} {pct(w['from_ok']):>11} "
                  f"{pct(w['close_id_ok']):>10} {pct(w['close_at_ok']):>9} "
                  f"{pct(w['history_ok']):>13} {pct(w['well_formed']):>12}")

    bar("Context size")
    for arm in V.ARMS:
        s = [r for r in grid if r["arm"] == arm]
        if s:
            print(f"{ARM_LABEL[arm]:26} {s[0]['context_chars']:6} chars   "
                  f"median latency {sorted(r['latency_s'] for r in s)[len(s)//2]:.1f}s")

    out = {"summary": summary, "by_family": byfam, "extra": extra, "tests": tests,
           "ablation": abl, "disclosure": disc_out, "writeback": wb_out,
           "calls": {"grid": len(grid), "disclosure": len(disc), "writeback": len(wb)}}
    dest = os.path.join(R, "analysis.json")
    with open(dest, "w", encoding="utf-8") as fh:
        json.dump(out, fh, indent=2, ensure_ascii=False, default=str)
    print("\nwritten:", dest)


if __name__ == "__main__":
    main()
