#!/usr/bin/env python3
"""Independent re-implementation of NOHARM benchmark statistical claims.

Written from scratch against the raw released data files only:
  data/donoharm-case-performance.csv
  data/severe-mode-counts.csv
  data/human-study.json

No code from repro_stats.py or analysis.ipynb was read or imported.

Usage:
  python3 verify_claims.py [--draws 20000] [--flips 200000] [--seed 7]
"""

import argparse
import csv
import itertools
import json
import math
import os
from collections import defaultdict

import numpy as np
from scipy import stats

DATA = "/Users/haohu/work/groundtruth/research/noharm/repos/noharm/data"
RAG = ["nia-1.0", "doxgpt-condensed", "openevidence", "glass-5.6-max"]


# ---------------------------------------------------------------- loading
def load_severe():
    """dict model -> {case_id: (stratum, n, sev_full)} for prompt=default."""
    out = defaultdict(dict)
    with open(os.path.join(DATA, "severe-mode-counts.csv")) as f:
        for r in csv.DictReader(f):
            if r["prompt"] != "default":
                continue
            out[r["model"]][r["case_id"]] = (
                r["stratum"], int(r["n"]), int(r["sev_full"]))
    return out


def load_perf():
    """dict model -> {case_id: (stratum, case_split, F1)} for prompt=default."""
    out = defaultdict(dict)
    with open(os.path.join(DATA, "donoharm-case-performance.csv")) as f:
        for r in csv.DictReader(f):
            if r["prompt"] != "default":
                continue
            out[r["model"]][r["case_id"]] = (
                r["stratum"], r["case_split"], float(r["F1_weighted"]))
    return out


# ---------------------------------------------------------- bootstrap core
def stratified_paired_bootstrap(values_by_model, strata, n_draws, rng):
    """values_by_model: dict model -> np.array aligned on cases.
    strata: list of stratum labels aligned on the same case order.
    Resample case indices with replacement WITHIN each stratum; the same
    resampled index set is applied to every model (paired / clustered on
    case). Returns dict model -> (n_draws,) array of bootstrap means."""
    strata = np.asarray(strata)
    idx_by_stratum = [np.flatnonzero(strata == s) for s in np.unique(strata)]
    n_cases = len(strata)
    # Pre-draw all resampled indices: (n_draws, n_cases)
    draws = np.empty((n_draws, n_cases), dtype=np.int64)
    col = 0
    for idx in idx_by_stratum:
        k = len(idx)
        pick = rng.integers(0, k, size=(n_draws, k))
        draws[:, col:col + k] = idx[pick]
        col += k
    return {m: v[draws].mean(axis=1) for m, v in values_by_model.items()}


def boot_p_one_sided(boot_diff):
    """Achieved significance level for H0: diff <= 0 vs observed diff > 0,
    i.e. share of bootstrap draws in which the difference is <= 0
    (the 'leader' loses its lead)."""
    return float(np.mean(boot_diff <= 0))


def holm(pvals):
    """Holm step-down adjustment. pvals: dict name -> p. Returns dict."""
    items = sorted(pvals.items(), key=lambda kv: kv[1])
    m = len(items)
    adj, running = {}, 0.0
    for i, (name, p) in enumerate(items):
        running = max(running, min(1.0, (m - i) * p))
        adj[name] = running
    return adj


def sign_flip_test(diffs, n_flips, rng):
    """Paired sign-flip (permutation) test of H0: mean(diff)=0.
    Exact enumeration over the nonzero diffs when feasible, else Monte
    Carlo. Returns (p_two_sided, p_one_sided_greater, exact?)."""
    d = np.asarray(diffs, dtype=float)
    nz = d[d != 0]
    obs = d.mean()
    n = len(nz)
    if n == 0:
        return 1.0, 1.0, True
    if n <= 20:  # exact: 2^n enumerations
        signs = np.array(list(itertools.product([1.0, -1.0], repeat=n)))
        null_means = (signs * nz).sum(axis=1) / len(d)
        p2 = float(np.mean(np.abs(null_means) >= abs(obs) - 1e-12))
        p1 = float(np.mean(null_means >= obs - 1e-12))
        return p2, p1, True
    signs = rng.choice([1.0, -1.0], size=(n_flips, n))
    null_means = (signs * nz).sum(axis=1) / len(d)
    # add-one correction for Monte Carlo permutation p-values
    p2 = float((np.sum(np.abs(null_means) >= abs(obs) - 1e-12) + 1)
               / (n_flips + 1))
    p1 = float((np.sum(null_means >= obs - 1e-12) + 1) / (n_flips + 1))
    return p2, p1, False


# ------------------------------------------------------------------ claims
def claim1(sev):
    print("=" * 78)
    print("CLAIM 1 - mean per-case severe-harm fraction (sev_full/n),"
          " default prompt, 100 cases")
    print("=" * 78)
    res = {}
    for m in RAG:
        cases = sev[m]
        assert len(cases) == 100, (m, len(cases))
        rates = np.array([sf / n for (_s, n, sf) in cases.values()])
        res[m] = rates.mean()
        # sensitivity: binary "any severe run on this case" indicator
        binary = np.array([1.0 if sf > 0 else 0.0
                           for (_s, n, sf) in cases.values()])
        # sensitivity: pooled (total severe runs / total runs)
        tot_sf = sum(sf for (_s, n, sf) in cases.values())
        tot_n = sum(n for (_s, n, sf) in cases.values())
        print(f"  {m:20s} mean(sev_full/n) = {rates.mean()*100:6.2f}%   "
              f"[binary any-severe = {binary.mean()*100:5.1f}%,"
              f" pooled = {tot_sf/tot_n*100:5.2f}%]")
    return res


def get_aligned(sev, models):
    """case-aligned rate arrays + strata for the given models."""
    case_ids = sorted(sev[models[0]].keys())
    strata = [sev[models[0]][c][0] for c in case_ids]
    vals = {m: np.array([sev[m][c][2] / sev[m][c][1] for c in case_ids])
            for m in models}
    return case_ids, strata, vals


def claim2(sev, n_draws, n_flips, rng):
    print("=" * 78)
    print(f"CLAIM 2 - severe-harm separability among RAG systems "
          f"(stratified paired cluster bootstrap, {n_draws} draws)")
    print("=" * 78)
    _cids, strata, vals = get_aligned(sev, RAG)
    means = {m: v.mean() for m, v in vals.items()}
    leader = min(means, key=means.get)
    print(f"  leader (lowest severe rate): {leader} "
          f"({means[leader]*100:.2f}%)")
    boots = stratified_paired_bootstrap(vals, strata, n_draws, rng)
    raw = {}
    for m in RAG:
        if m == leader:
            continue
        raw[m] = boot_p_one_sided(boots[m] - boots[leader])
    adj = holm(raw)
    for m in sorted(raw, key=raw.get):
        d = means[m] - means[leader]
        print(f"  {m:20s} obs diff = +{d*100:.2f}pp   "
              f"one-sided boot p = {raw[m]:.4f}   Holm-adj = {adj[m]:.4f}")

    # unstratified sensitivity
    boots_u = stratified_paired_bootstrap(
        vals, ["all"] * len(strata), n_draws, rng)
    raw_u = {m: boot_p_one_sided(boots_u[m] - boots_u[leader])
             for m in RAG if m != leader}
    adj_u = holm(raw_u)
    print("  [sensitivity: unstratified bootstrap] " +
          ", ".join(f"{m}: p={raw_u[m]:.4f} (Holm {adj_u[m]:.4f})"
                    for m in sorted(raw_u, key=raw_u.get)))

    # paired sign-flip tests, nia vs each other tool
    print(f"  sign-flip permutation tests on per-case rate differences "
          f"(vs {leader}):")
    sf_raw = {}
    for m in RAG:
        if m == leader:
            continue
        d = vals[m] - vals[leader]
        p2, p1, exact = sign_flip_test(d, n_flips, rng)
        sf_raw[m] = p1
        nz = int(np.sum(d != 0))
        print(f"    {m:20s} mean diff = {d.mean()*100:+.2f}pp  "
              f"nonzero cases = {nz:3d}  "
              f"p(two-sided) = {p2:.4f}  p(one-sided) = {p1:.4f}  "
              f"({'exact' if exact else 'MC'})")
    sf_adj = holm(sf_raw)
    print("    Holm-adjusted one-sided: " +
          ", ".join(f"{m}: {sf_adj[m]:.4f}"
                    for m in sorted(sf_adj, key=sf_adj.get)))
    # Wilcoxon signed-rank as an extra sanity check
    for m in RAG:
        if m == leader:
            continue
        d = vals[m] - vals[leader]
        try:
            w = stats.wilcoxon(d, alternative="greater",
                               zero_method="wilcox")
            print(f"    wilcoxon {m:20s} p = {w.pvalue:.4f}")
        except ValueError as e:
            print(f"    wilcoxon {m:20s} n/a ({e})")


def claim3_4(perf, n_draws, n_flips, rng, split=None, label="all 100 cases"):
    print("=" * 78)
    print(f"CLAIM {'3' if split is None else '4'} - F1_weighted, {label} "
          f"(stratified paired cluster bootstrap, {n_draws} draws)")
    print("=" * 78)
    case_ids = sorted(c for c, (_s, sp, _f) in perf[RAG[0]].items()
                      if split is None or sp == split)
    strata = [perf[RAG[0]][c][0] for c in case_ids]
    vals = {m: np.array([perf[m][c][2] for c in case_ids]) for m in RAG}
    print(f"  n cases = {len(case_ids)}")
    means = {m: v.mean() for m, v in vals.items()}
    for m in sorted(means, key=means.get, reverse=True):
        print(f"  {m:20s} mean F1 = {means[m]:.4f}")
    leader = max(means, key=means.get)
    print(f"  leader (highest F1): {leader}")
    boots = stratified_paired_bootstrap(vals, strata, n_draws, rng)
    raw = {}
    for m in RAG:
        if m == leader:
            continue
        # one-sided ASL that the leader's advantage is <= 0
        raw[m] = boot_p_one_sided(boots[leader] - boots[m])
    adj = holm(raw)
    for m in sorted(raw, key=raw.get):
        d = means[leader] - means[m]
        print(f"  {leader} vs {m:20s} obs diff = +{d:.4f}   "
              f"one-sided boot p = {raw[m]:.5f}   Holm-adj = {adj[m]:.5f}")
    # sign-flip permutation sanity checks
    print("  sign-flip permutation tests (leader minus other, per case):")
    sf_raw = {}
    for m in RAG:
        if m == leader:
            continue
        d = vals[leader] - vals[m]
        p2, p1, exact = sign_flip_test(d, n_flips, rng)
        sf_raw[m] = p1
        print(f"    vs {m:20s} mean diff = {d.mean():+.4f}  "
              f"p(two-sided) = {p2:.5f}  p(one-sided) = {p1:.5f}  "
              f"({'exact' if exact else 'MC'})")
    sf_adj = holm(sf_raw)
    print("    Holm-adjusted one-sided: " +
          ", ".join(f"{m}: {sf_adj[m]:.5f}"
                    for m in sorted(sf_adj, key=sf_adj.get)))
    # head-to-head nia vs doxgpt regardless of leader (for claim 4 wording)
    if split == "held_out":
        d = vals["doxgpt-condensed"] - vals["nia-1.0"]
        p2, p1, exact = sign_flip_test(d, n_flips, rng)
        bdiff = boots["doxgpt-condensed"] - boots["nia-1.0"]
        print(f"  head-to-head doxgpt-condensed minus nia-1.0: "
              f"mean diff = {d.mean():+.5f}, "
              f"boot one-sided p = {boot_p_one_sided(bdiff):.4f}, "
              f"sign-flip two-sided p = {p2:.4f}")


def claim5():
    print("=" * 78)
    print("CLAIM 5 - human-study resource_breakdown")
    print("=" * 78)
    with open(os.path.join(DATA, "human-study.json")) as f:
        h = json.load(f)
    rb = h["resource_breakdown"]
    n = rb["n"]
    print(f"  n = {n}")
    rows = sorted(rb["resources"], key=lambda r: r["count"], reverse=True)
    for rank, r in enumerate(rows, 1):
        my_pct = 100.0 * r["count"] / n
        print(f"  #{rank} {r['name'][:45]:45s} count={r['count']:3d}  "
              f"file pct={r['pct']:5.1f}  recomputed={my_pct:5.2f}")
    return rb


def claim6():
    print("=" * 78)
    print("CLAIM 6 - uncertainty on OpenEvidence 45 vs External AI 40 "
          "(n=202, multi-select, joint counts unknown)")
    print("=" * 78)
    n, a, b = 202, 45, 40
    # Same respondents answered both items -> paired binary data.
    # Let k = # who used BOTH. Discordant counts: b10 = 45-k (OE only),
    # b01 = 40-k (ExtAI only). McNemar exact: b10 ~ Bin(b10+b01, 1/2).
    print(f"  paired-data view: overlap k ranges "
          f"{max(0, a + b - n)}..{min(a, b)}")
    best = None
    for k in range(max(0, a + b - n), min(a, b) + 1):
        b10, b01 = a - k, b - k
        m = b10 + b01
        p1 = stats.binomtest(b10, m, 0.5, alternative="greater").pvalue
        p2 = stats.binomtest(b10, m, 0.5, alternative="two-sided").pvalue
        if best is None or p2 < best[3]:
            best = (k, m, p1, p2)
        if k in (max(0, a + b - n), min(a, b)) or k % 10 == 0:
            print(f"    k={k:2d}: discordant={m:2d} "
                  f"(b10={b10}, b01={b01})  exact McNemar "
                  f"one-sided p={p1:.4f}  two-sided p={p2:.4f}")
    k, m, p1, p2 = best
    print(f"  most favorable overlap for significance: k={k} "
          f"(complete nesting), discordant={m}, one-sided p={p1:.4f}, "
          f"two-sided p={p2:.4f}")
    # naive independent-samples comparison (wrong model, upper bound
    # on evidence if treated as two independent proportions)
    cont = np.array([[a, n - a], [b, n - b]])
    chi = stats.chi2_contingency(cont, correction=False)
    print(f"  naive two-independent-proportions chi2 p = {chi.pvalue:.4f} "
          f"(45/202=22.3% vs 40/202=19.8%)")
    # normal-approx CI on the difference in the best case
    print(f"  conclusion: min possible two-sided p = {p2:.4f} > 0.05 -> "
          f"no overlap structure makes the gap significant two-sided; "
          f"only complete nesting (k=40) reaches p={p1:.4f} one-sided.")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--draws", type=int, default=20000)
    ap.add_argument("--flips", type=int, default=200000)
    ap.add_argument("--seed", type=int, default=7)
    args = ap.parse_args()
    rng = np.random.default_rng(args.seed)

    sev = load_severe()
    perf = load_perf()

    claim1(sev)
    claim2(sev, args.draws, args.flips, rng)
    claim3_4(perf, args.draws, args.flips, rng,
             split=None, label="all 100 cases")
    claim3_4(perf, args.draws, args.flips, rng,
             split="held_out", label="held_out subset")
    claim5()
    claim6()


if __name__ == "__main__":
    main()
