#!/usr/bin/env python3
"""Reanalysis of integrity-bench per-question data: disentangle bias vs
discrimination (AUROC) vs banding artifact. Data: data/cached/*.json,
rows = [qid, level, correct(0/1), confidence(0-100), cost, time]."""
import json, glob, os

def auroc(pos, neg):
    """AUROC via Mann-Whitney U: P(conf_correct > conf_wrong) + 0.5*ties."""
    if not pos or not neg:
        return float('nan')
    wins = ties = 0
    sneg = sorted(neg)
    import bisect
    for p in pos:
        lo = bisect.bisect_left(sneg, p)
        hi = bisect.bisect_right(sneg, p)
        wins += lo
        ties += hi - lo
    return (wins + 0.5 * ties) / (len(pos) * len(neg))

def band_rows(rows, mode):
    """rows: list of (level, correct, conf). Pick 3-consecutive-level band
    with accuracy closest to 50%. mode: 'all' | 'global' | 'perdomain'
    (perdomain handled by caller passing single-domain rows with 'global')."""
    if mode == 'all':
        return rows
    from collections import defaultdict
    by_level = defaultdict(list)
    for lv, c, cf in rows:
        by_level[lv].append((lv, c, cf))
    levels = sorted(by_level)
    best, best_d = None, 9e9
    for i in range(len(levels) - 2):
        window = levels[i:i+3]
        # require consecutive
        if window[2] - window[0] != 2:
            continue
        sel = [r for lv in window for r in by_level[lv]]
        acc = sum(r[1] for r in sel) / len(sel)
        if abs(acc - 0.5) < best_d:
            best_d, best = abs(acc - 0.5), sel
    return best if best else rows

def metrics(rows):
    n = len(rows)
    a = sum(r[1] for r in rows) / n
    cbar = sum(r[2] for r in rows) / n / 100.0
    brier = sum((r[2]/100.0 - r[1])**2 for r in rows) / n
    integrity = 100 - 400 * brier
    floor = 100 - 400 * a * (1 - a)
    pos = [r[2] for r in rows if r[1] == 1]
    neg = [r[2] for r in rows if r[1] == 0]
    return dict(n=n, acc=a, meanconf=cbar, bias=cbar - a,
                integrity=integrity, floor=floor, gap=integrity - floor,
                auroc=auroc(pos, neg))

results = {}
for path in sorted(glob.glob('data/cached/*.json')):
    d = json.load(open(path))
    model = d['model']
    allrows, banded = [], []
    for dom, rows in d['domains'].items():
        drows = [(r[1], int(r[2]), float(r[3])) for r in rows]
        allrows += drows
        banded += band_rows(drows, 'global')  # per-domain 3-level band
    results[model] = {
        'all': metrics(allrows),
        'band': metrics(banded),
    }

hdr = f"{'model':22s} {'n':>4s} {'acc':>5s} {'conf':>5s} {'bias':>6s} {'Integ':>7s} {'floor':>6s} {'gap':>7s} {'AUROC':>6s} | {'AUROC-all':>9s}"
print("Per-domain 3-level band closest to 50% (reproduction attempt); AUROC-all = full 8-level scout range")
print(hdr)
for m, r in sorted(results.items(), key=lambda kv: -kv[1]['band']['gap']):
    b, al = r['band'], r['all']
    print(f"{m:22s} {b['n']:4d} {b['acc']:5.1%} {b['meanconf']:5.1%} {b['bias']:+6.1%} "
          f"{b['integrity']:7.1f} {b['floor']:6.1f} {b['gap']:+7.1f} {b['auroc']:6.3f} | {al['auroc']:9.3f}")
