#!/usr/bin/env python3
"""Condition- and item-level summaries, clustered bootstrap CIs, and figures.

Usage
    python src/analyze.py --run-dir runs/hostllm_r3 --out-dir outputs

Primary outcome
    correct_to_incorrect, conditioned on an initially correct answer:
        P(turn-2 answer incorrect | turn-1 answer correct)

Inference
    All intervals are item-clustered percentile bootstrap intervals (10,000
    resamples, seed 20260923). Items, not trials, are the resampling unit,
    because each item contributes repeated trials and those trials are not
    independent. Condition contrasts resample items once and evaluate both
    conditions on the same resample, so the interval is for the paired
    difference and correctly reflects the shared item variance.
"""

import argparse
import json
from pathlib import Path

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd

SEED = 20260923
N_BOOT = 10000
CONDITION_ORDER = [
    "neutral_reconsideration",
    "majority_false",
    "unanimous_experts_false",
    "fictional_claim_transparent",
]
PRIMARY = "correct_to_incorrect"


def rate_with_ci(df, mask_col, outcome_col, n_boot=N_BOOT, seed=SEED):
    """Clustered bootstrap rate of outcome_col among rows where mask_col is True."""
    d = df[df[mask_col]]
    if d.empty:
        return dict(n=0, k=0, rate=np.nan, lo=np.nan, hi=np.nan)
    item_codes = d["item_id"].astype("category").cat.codes.to_numpy()
    outcomes = d[outcome_col].to_numpy(dtype=float)
    n_items = item_codes.max() + 1

    by_item_k = np.bincount(item_codes, weights=outcomes, minlength=n_items)
    by_item_n = np.bincount(item_codes, minlength=n_items).astype(float)

    rng = np.random.default_rng(seed)
    k, n = outcomes.sum(), len(outcomes)
    rate = k / n
    draws = rng.integers(0, n_items, size=(n_boot, n_items))
    num = by_item_k[draws].sum(axis=1)
    den = by_item_n[draws].sum(axis=1)
    valid = den > 0
    boot = num[valid] / den[valid]
    lo, hi = np.percentile(boot, [2.5, 97.5])
    return dict(n=int(n), k=int(k), rate=float(rate), lo=float(lo), hi=float(hi))


def paired_diff_ci(df, mask_col, outcome_col, cond_a, cond_b, n_boot=N_BOOT, seed=SEED):
    """Clustered bootstrap CI for rate(cond_a) - rate(cond_b), resampling items."""
    d = df[df[mask_col]]
    items = np.sort(df["item_id"].unique())
    index = {it: i for i, it in enumerate(items)}
    n_items = len(items)

    k = {cond_a: np.zeros(n_items), cond_b: np.zeros(n_items)}
    n = {cond_a: np.zeros(n_items), cond_b: np.zeros(n_items)}
    for cond, dfc in d.groupby("condition"):
        if cond not in k:
            continue
        idx = dfc["item_id"].map(index).to_numpy()
        np.add.at(k[cond], idx, dfc[outcome_col].to_numpy(dtype=float))
        np.add.at(n[cond], idx, 1.0)

    def ratio(kk, nn):
        return kk.sum() / nn.sum() if nn.sum() else np.nan

    point = ratio(k[cond_a], n[cond_a]) - ratio(k[cond_b], n[cond_b])

    rng = np.random.default_rng(seed)
    draws = rng.integers(0, n_items, size=(n_boot, n_items))
    ka, na = k[cond_a][draws].sum(axis=1), n[cond_a][draws].sum(axis=1)
    kb, nb = k[cond_b][draws].sum(axis=1), n[cond_b][draws].sum(axis=1)
    valid = (na > 0) & (nb > 0)
    boot = ka[valid] / na[valid] - kb[valid] / nb[valid]
    lo, hi = np.percentile(boot, [2.5, 97.5])
    # two-sided bootstrap p-value for a difference of zero
    p = 2 * min((boot <= 0).mean(), (boot >= 0).mean())
    return dict(
        contrast=f"{cond_a} minus {cond_b}",
        point=float(point),
        lo=float(lo),
        hi=float(hi),
        p_value=float(min(p, 1.0)),
    )


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--run-dir", required=True)
    ap.add_argument("--out-dir", default="outputs")
    ap.add_argument("--label", default=None, help="label used in figure titles and the run banner")
    args = ap.parse_args()

    run_dir = Path(args.run_dir)
    out_dir = Path(args.out_dir)
    tables_dir = out_dir / "tables"
    figures_dir = out_dir / "figures"
    tables_dir.mkdir(parents=True, exist_ok=True)
    figures_dir.mkdir(parents=True, exist_ok=True)

    df = pd.read_csv(run_dir / "scored_trials.csv")
    manifest_path = run_dir / "run_manifest.json"
    manifest = json.loads(manifest_path.read_text()) if manifest_path.exists() else {}
    is_simulated = bool(manifest.get("is_simulated", False))
    label = args.label or run_dir.name

    banner = "SIMULATED PIPELINE CHECK - NOT EXPERIMENTAL RESULTS" if is_simulated else "EXPERIMENTAL RUN"
    if is_simulated:
        print("!" * 78)
        print(f"!! {banner}")
        print("!! These numbers come from a deterministic mock responder. They are")
        print("!! not measurements of any model and must not be reported as findings.")
        print("!" * 78)

    conditions = [c for c in CONDITION_ORDER if c in set(df["condition"])]
    conditions += [c for c in df["condition"].unique() if c not in conditions]

    # ---------------------------------------------------- condition summary ---
    rows = []
    for cond in conditions:
        d = df[df["condition"] == cond]
        d_el = d[d["eligible_baseline"]]
        prim = rate_with_ci(d, "eligible_primary", PRIMARY)
        supplied = rate_with_ci(d, "eligible_primary", "correct_to_supplied_false")
        other_wrong = rate_with_ci(d, "eligible_primary", "correct_to_other_wrong")
        changed = rate_with_ci(d, "eligible_primary", "answer_changed")
        repair = rate_with_ci(d, "eligible_baseline", "incorrect_to_correct")
        rows.append(
            {
                "condition": cond,
                "n_trials": len(d),
                "n_eligible_baseline": len(d_el),
                "baseline_accuracy": d_el["baseline_correct"].mean() if len(d_el) else np.nan,
                "n_baseline_correct": int(d["eligible_primary"].sum()),
                "correct_to_incorrect_rate": prim["rate"],
                "correct_to_incorrect_lo": prim["lo"],
                "correct_to_incorrect_hi": prim["hi"],
                "correct_to_incorrect_n": prim["n"],
                "correct_to_incorrect_k": prim["k"],
                "correct_to_supplied_false_rate": supplied["rate"],
                "correct_to_supplied_false_lo": supplied["lo"],
                "correct_to_supplied_false_hi": supplied["hi"],
                "correct_to_other_wrong_rate": other_wrong["rate"],
                "answer_changed_rate": changed["rate"],
                "answer_changed_lo": changed["lo"],
                "answer_changed_hi": changed["hi"],
                "incorrect_to_correct_rate": repair["rate"],
                "final_accuracy": (
                    d[d["eligible_primary"]]["final_correct"].mean() if d["eligible_primary"].any() else np.nan
                ),
                "turn2_unparseable_rate": (
                    d[d["eligible_baseline"]]["turn2_unparseable"].mean() if len(d_el) else np.nan
                ),
            }
        )
    summary = pd.DataFrame(rows)
    summary.to_csv(tables_dir / f"condition_summary_{label}.csv", index=False)
    summary.to_csv(tables_dir / "condition_summary.csv", index=False)

    # --------------------------------------------------------- contrasts -------
    contrast_rows = []
    for cond in conditions:
        if cond == "neutral_reconsideration":
            continue
        for outcome in [PRIMARY, "correct_to_supplied_false", "answer_changed"]:
            contrast_rows.append(
                {
                    "outcome": outcome,
                    **paired_diff_ci(df, "eligible_primary" if outcome != "answer_changed" else "eligible_primary",
                                     outcome, cond, "neutral_reconsideration"),
                }
            )
    if "unanimous_experts_false" in conditions and "majority_false" in conditions:
        contrast_rows.append(
            {
                "outcome": PRIMARY,
                **paired_diff_ci(df, "eligible_primary", PRIMARY, "unanimous_experts_false", "majority_false"),
            }
        )
    if "fictional_claim_transparent" in conditions and "unanimous_experts_false" in conditions:
        contrast_rows.append(
            {
                "outcome": PRIMARY,
                **paired_diff_ci(df, "eligible_primary", PRIMARY, "fictional_claim_transparent", "unanimous_experts_false"),
            }
        )
    contrasts = pd.DataFrame(contrast_rows)
    contrasts.to_csv(tables_dir / f"contrasts_{label}.csv", index=False)
    contrasts.to_csv(tables_dir / "contrasts.csv", index=False)

    # -------------------------------------------------------- item level ------
    item_rows = []
    for (item, cond), d in df.groupby(["item_id", "condition"]):
        elig = d[d["eligible_primary"]]
        item_rows.append(
            {
                "item_id": item,
                "domain": d["domain"].iloc[0],
                "condition": cond,
                "n_primary": len(elig),
                "n_correct_to_incorrect": int(elig[PRIMARY].sum()),
                "rate_correct_to_incorrect": elig[PRIMARY].mean() if len(elig) else np.nan,
                "baseline_correct_any": bool(d[d["eligible_baseline"]]["baseline_correct"].any()),
            }
        )
    item_level = pd.DataFrame(item_rows)
    item_level.to_csv(tables_dir / f"item_level_{label}.csv", index=False)
    item_level.to_csv(tables_dir / "item_level.csv", index=False)

    # --------------------------------------------------- outcome categories ---
    cat_rows = []
    for cond in conditions:
        d = df[(df["condition"] == cond) & df["eligible_primary"]]
        cat_rows.append(
            {
                "condition": cond,
                "n": len(d),
                "stayed_correct": int(d["final_correct"].sum()),
                "to_supplied_false": int(d["correct_to_supplied_false"].sum()),
                "to_other_wrong": int(d["correct_to_other_wrong"].sum()),
            }
        )
    pd.DataFrame(cat_rows).to_csv(tables_dir / f"outcome_distribution_{label}.csv", index=False)

    # ------------------------------------------------------------ figures -----
    plt.rcParams.update({"font.size": 10, "axes.spines.top": False, "axes.spines.right": False})

    def bar_with_ci(ax, values, los, his, title, ylabel, color):
        x = np.arange(len(values))
        err = np.array([[v - lo for v, lo in zip(values, los)], [hi - v for v, hi in zip(values, his)]])
        ax.bar(x, values, color=color, width=0.6)
        ax.errorbar(x, values, yerr=err, fmt="none", ecolor="#222222", capsize=4, linewidth=1.2)
        ax.set_xticks(x)
        ax.set_xticklabels([c.replace("_", "\n") for c in conditions], fontsize=8)
        ax.set_ylabel(ylabel)
        ax.set_title(title, fontsize=10)
        ax.set_ylim(0, 1)
        for xi, v in zip(x, values):
            ax.text(xi, min(v + 0.03, 0.96), f"{v:.2f}", ha="center", fontsize=8)

    fig, axes = plt.subplots(1, 2, figsize=(11, 4.2))
    bar_with_ci(
        axes[0],
        summary["correct_to_incorrect_rate"],
        summary["correct_to_incorrect_lo"],
        summary["correct_to_incorrect_hi"],
        "Correct to incorrect after follow-up",
        "P(incorrect | initially correct)",
        "#c0392b",
    )
    bar_with_ci(
        axes[1],
        summary["correct_to_supplied_false_rate"],
        summary["correct_to_supplied_false_lo"],
        summary["correct_to_supplied_false_hi"],
        "Correct to the supplied false answer",
        "P(supplied false | initially correct)",
        "#8e44ad",
    )
    suptitle = f"{banner}\nPrimary outcomes ({label})" if is_simulated else f"Primary outcomes ({label})"
    fig.suptitle(suptitle, fontsize=11, color="#c0392b" if is_simulated else "#111111")
    fig.tight_layout()
    fig.savefig(figures_dir / f"fig1_primary_outcomes_{label}.png", dpi=200)
    plt.close(fig)

    # outcome composition
    comp = pd.DataFrame(cat_rows).set_index("condition").loc[conditions]
    fig, ax = plt.subplots(figsize=(7.5, 4))
    bottom = np.zeros(len(comp))
    for col, color, lab in [
        ("stayed_correct", "#2e86c1", "stayed correct"),
        ("to_supplied_false", "#c0392b", "switched to supplied false answer"),
        ("to_other_wrong", "#f39c12", "switched to another wrong answer"),
    ]:
        vals = comp[col].to_numpy(dtype=float)
        ax.bar(np.arange(len(comp)), vals, bottom=bottom, color=color, label=lab, width=0.6)
        bottom += vals
    ax.set_xticks(np.arange(len(comp)))
    ax.set_xticklabels([c.replace("_", "\n") for c in comp.index], fontsize=8)
    ax.set_ylabel("trials (initially correct)")
    ax.set_title("Where initially correct answers go", fontsize=10)
    ax.legend(fontsize=8)
    fig.tight_layout()
    fig.savefig(figures_dir / f"fig2_outcome_composition_{label}.png", dpi=200)
    plt.close(fig)

    # per-item distribution
    fig, ax = plt.subplots(figsize=(7.5, 4))
    data = [
        item_level[(item_level["condition"] == cond)]["rate_correct_to_incorrect"].dropna().to_numpy()
        for cond in conditions
    ]
    ax.boxplot(data, tick_labels=[c.replace("_", "\n") for c in conditions], showfliers=True)
    ax.set_ylabel("per-item P(correct to incorrect)")
    ax.set_title("Item-level spread of the primary outcome", fontsize=10)
    ax.tick_params(axis="x", labelsize=8)
    fig.tight_layout()
    fig.savefig(figures_dir / f"fig3_item_level_{label}.png", dpi=200)
    plt.close(fig)

    # baseline accuracy check (should not differ by condition)
    fig, ax = plt.subplots(figsize=(7.5, 4))
    ax.bar(np.arange(len(conditions)), summary["baseline_accuracy"], color="#16a085", width=0.6)
    ax.set_xticks(np.arange(len(conditions)))
    ax.set_xticklabels([c.replace("_", "\n") for c in conditions], fontsize=8)
    ax.set_ylabel("turn-1 accuracy")
    ax.set_ylim(0, 1)
    ax.set_title("Manipulation check: turn-1 accuracy before any pressure", fontsize=10)
    fig.tight_layout()
    fig.savefig(figures_dir / f"fig4_baseline_accuracy_{label}.png", dpi=200)
    plt.close(fig)

    # ------------------------------------------------------------- console ----
    with pd.option_context("display.width", 200, "display.max_columns", 50):
        print("\n== condition summary ==")
        print(summary[["condition", "n_trials", "baseline_accuracy", "correct_to_incorrect_rate",
                       "correct_to_incorrect_lo", "correct_to_incorrect_hi",
                       "correct_to_supplied_false_rate", "answer_changed_rate"]].to_string(index=False))
        print("\n== contrasts (item-clustered bootstrap) ==")
        print(contrasts.to_string(index=False))

    sidecar = {
        "label": label,
        "run_dir": str(run_dir),
        "is_simulated": is_simulated,
        "n_trials": int(len(df)),
        "bootstrap_resamples": N_BOOT,
        "bootstrap_unit": "item_id (clustered)",
        "seed": SEED,
        "primary_outcome": PRIMARY,
        "generated_at": pd.Timestamp.now(tz="UTC").isoformat(),
    }
    (out_dir / f"analysis_manifest_{label}.json").write_text(json.dumps(sidecar, indent=2) + "\n")
    print(f"\nwrote tables to {tables_dir} and figures to {figures_dir}")


if __name__ == "__main__":
    main()
