"""
SGB - scoring and reporting.

Unit of analysis is the TASK, not the call. Repeats of the same task on the same prompt
are not independent observations, so repeats are averaged inside a task first and the
paired test runs across tasks.

  python analyze.py ../results/haiku_v2.jsonl
"""

import json
import math
import sys
from collections import defaultdict

ARMS = ["schema", "prose", "prose_full", "vercy", "vercy_pad", "vercy_fed"]
ARM_LABEL = {
    "schema": "A schema only",
    "prose": "B prose, incomplete",
    "prose_full": "E prose, complete",
    "vercy": "C versioned records",
    "vercy_pad": "F records + padding",
    "vercy_fed": "D records + federation",
}


def wilson(k, n, z=1.96):
    if n == 0:
        return (0.0, 0.0, 0.0)
    p = k / n
    d = 1 + z * z / n
    c = (p + z * z / (2 * n)) / d
    h = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d
    return (p, max(0.0, c - h), min(1.0, c + h))


def cluster_ci(values):
    """Mean over tasks with a normal CI on the task-level mean (cluster = task)."""
    n = len(values)
    if n == 0:
        return (0.0, 0.0, 0.0)
    m = sum(values) / n
    if n < 2:
        return (m, m, m)
    var = sum((v - m) ** 2 for v in values) / (n - 1)
    se = math.sqrt(var / n)
    return (m, max(0.0, m - 1.96 * se), min(1.0, m + 1.96 * se))


def fmt_ci(t):
    m, lo, hi = t
    return "%5.1f%% [%.0f-%.0f]" % (100 * m, 100 * lo, 100 * hi)


def sign_test(diffs):
    """Exact two-sided sign test on non-zero paired differences."""
    pos = sum(1 for d in diffs if d > 0)
    neg = sum(1 for d in diffs if d < 0)
    n = pos + neg
    if n == 0:
        return pos, neg, 1.0
    k = min(pos, neg)
    tail = sum(math.comb(n, i) for i in range(0, k + 1)) / (2 ** n)
    return pos, neg, min(1.0, 2 * tail)


def load(paths):
    rows = []
    for p in paths:
        with open(p, encoding="utf-8") as fh:
            for line in fh:
                line = line.strip()
                if line:
                    rows.append(json.loads(line))
    return rows


def task_means(rows, arm, metric, gate):
    """Mean of `metric` per task for one arm, over tasks where `gate` is set."""
    acc = defaultdict(list)
    for r in rows:
        if r["arm"] == arm and r[gate]:
            acc[r["task"]].append(r[metric])
    return {t: sum(v) / len(v) for t, v in acc.items()}


def report(rows, model):
    rows = [r for r in rows if r["model"] == model]
    if not rows:
        return
    arms = [a for a in ARMS if any(r["arm"] == a for r in rows)]
    print("=" * 100)
    print("MODEL: %s   calls: %d   arms: %d" % (model, len(rows), len(arms)))
    print("=" * 100)

    print("\nTask-level means (repeats averaged within a task, CI clustered by task)")
    print("M1 answer accuracy | M2 governing definition cited | M5 correct abstention")
    print("M7 answer exactly equal to the prespecified competing definition's answer")
    print("-" * 100)
    print("%-24s %-19s %-19s %-19s %-19s" % ("arm", "M1 accuracy", "M2 definition",
                                             "M5 abstention", "M7 competing-answer"))
    store = {}
    for a in arms:
        m1 = task_means(rows, a, "answer_correct", "scored_answer")
        m2 = task_means(rows, a, "def_correct", "scored_def")
        m5 = task_means(rows, a, "abstained", "scored_abst")
        m7 = task_means(rows, a, "used_alt_definition", "scored_answer")
        store[a] = m1
        print("%-24s %-19s %-19s %-19s %-19s" % (
            ARM_LABEL[a], fmt_ci(cluster_ci(list(m1.values()))),
            fmt_ci(cluster_ci(list(m2.values()))),
            fmt_ci(cluster_ci(list(m5.values()))),
            fmt_ci(cluster_ci(list(m7.values())))))

    print("\nM1 by task family (task-level means)")
    print("-" * 100)
    fams = ["COMP", "TEMP", "XORG", "ABST"]
    print("%-24s %-17s %-17s %-17s %-17s" % ("arm", *fams))
    for a in arms:
        cells = []
        for f in fams:
            sub = [r for r in rows if r["arm"] == a and r["family"] == f]
            metric, gate = ("abstained", "scored_abst") if f == "ABST" else \
                           ("answer_correct", "scored_answer")
            acc = defaultdict(list)
            for r in sub:
                if r[gate]:
                    acc[r["task"]].append(r[metric])
            vals = [sum(v) / len(v) for v in acc.values()]
            cells.append("%5.1f%% (n=%d)" % (100 * sum(vals) / len(vals) if vals else 0,
                                             len(vals)))
        print("%-24s %-17s %-17s %-17s %-17s" % (ARM_LABEL[a], *cells))

    print("\ncost")
    print("-" * 100)
    for a in arms:
        sub = [r for r in rows if r["arm"] == a]
        n = len(sub)
        print("%-24s prompt %6d chars   latency %5.1f s   parse ok %5.1f%%   false abstention %4.1f%%"
              % (ARM_LABEL[a], sum(r["prompt_chars"] for r in sub) // n,
                 sum(r["latency_s"] for r in sub) / n,
                 100 * sum(r["parsed"] for r in sub) / n,
                 100 * sum(r["false_abstention"] for r in sub) /
                 max(1, sum(r["scored_answer"] for r in sub))))

    print("\nPaired comparison across tasks (exact sign test on task accuracies)")
    print("-" * 100)
    pairs = [("schema", "vercy"), ("prose", "vercy"), ("prose_full", "vercy"),
             ("vercy", "vercy_pad"), ("vercy", "vercy_fed"), ("vercy_pad", "vercy_fed"),
             ("schema", "prose_full"), ("prose", "prose_full")]
    for a, b in pairs:
        if a not in store or b not in store:
            continue
        keys = sorted(set(store[a]) & set(store[b]))
        diffs = [store[b][k] - store[a][k] for k in keys]
        pos, neg, p = sign_test(diffs)
        ma = sum(store[a][k] for k in keys) / len(keys)
        mb = sum(store[b][k] for k in keys) / len(keys)
        print("%-11s %5.1f%%  ->  %-11s %5.1f%%   tasks=%2d  better:%2d worse:%2d  p=%.4f"
              % (a, 100 * ma, b, 100 * mb, len(keys), pos, neg, p))

    print("\nScope discriminators (federation records must NOT apply)")
    print("-" * 100)
    for tid in ("XORG-11", "XORG-12"):
        cells = []
        for a in arms:
            sub = [r for r in rows if r["arm"] == a and r["task"] == tid]
            if sub:
                cells.append("%s %d/%d" % (a, sum(r["answer_correct"] for r in sub), len(sub)))
        print("%-9s %s" % (tid, "   ".join(cells)))


def main():
    paths = sys.argv[1:]
    if not paths:
        print(__doc__)
        return
    rows = load(paths)
    for model in sorted({r["model"] for r in rows}):
        report(rows, model)


if __name__ == "__main__":
    main()
