#!/usr/bin/env python3
"""
Bayesian beta regression models for game-level outcomes.

Outcomes: consensus_strength, evidence_alignment, SDI (supplementary)
Predictors: PCE (numeric), pressure (binary), discussion length (centered + quadratic),
            and key interactions (PCE × pressure, PCE × length).

Uses Bambi/PyMC for Bayesian estimation with weakly informative priors.
"""
from __future__ import annotations

import os
import warnings

import arviz as az
import bambi as bmb
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd

warnings.filterwarnings("ignore", category=FutureWarning)
warnings.filterwarnings("ignore", category=UserWarning, module="pytensor")

BASE = "./MSM"  # Update this path to your local working directory
TABLE_DIR = os.path.join(BASE, "tables")
FIG_DIR = os.path.join(BASE, "figures")
os.makedirs(TABLE_DIR, exist_ok=True)
os.makedirs(FIG_DIR, exist_ok=True)


# ── 1. Load and prepare data ──────────────────────────────────────────

def load_game_data() -> pd.DataFrame:
    df = pd.read_csv(os.path.join(BASE, "processed_data", "game_level.csv"))

    # Numeric coding of predictors
    pce_map = {"0%": 0.0, "25%": 0.25, "50%": 0.50, "75%": 0.75}
    df["pce_numeric"] = df["pce_level"].map(pce_map)

    df["pressure_ec"] = (df["pressure_condition"] == "evidence_centric").astype(int)

    length_map = {"short": 2, "medium": 4, "long": 6}
    df["length_numeric"] = df["discussion_length"].map(length_map)
    df["length_centered"] = df["length_numeric"] - 4  # center at medium
    df["length_sq"] = df["length_centered"] ** 2

    # Interaction terms
    df["pce_x_pressure"] = df["pce_numeric"] * df["pressure_ec"]
    df["pce_x_length"] = df["pce_numeric"] * df["length_centered"]

    # Smithson & Verkuilen (2006) transformation for beta regression
    N = len(df)
    for col in ["consensus_strength", "evidence_alignment", "semantic_decoupling_index"]:
        df[f"{col}_beta"] = (df[col] * (N - 1) + 0.5) / N

    return df


# ── 2. Fit Bayesian beta regression ──────────────────────────────────

def fit_beta_model(
    df: pd.DataFrame,
    outcome: str,
    label: str,
    n_draws: int = 2000,
    n_tune: int = 2000,
    n_chains: int = 4,
    target_accept: float = 0.95,
) -> dict:
    """Fit a Bayesian beta regression with Bambi."""
    formula = f"{outcome} ~ pce_numeric + pressure_ec + length_centered + length_sq + pce_x_pressure + pce_x_length"

    print(f"\n{'='*60}")
    print(f"Fitting model: {label}")
    print(f"Formula: {formula}")
    print(f"N = {len(df)}")
    print(f"{'='*60}")

    model = bmb.Model(
        formula,
        data=df,
        family="beta",
        # Weakly informative priors (Bambi defaults are already weakly informative)
    )

    idata = model.fit(
        draws=n_draws,
        tune=n_tune,
        chains=n_chains,
        target_accept=target_accept,
        random_seed=42,
        progressbar=True,
    )

    # Diagnostics
    summary = az.summary(idata, hdi_prob=0.95)
    print(f"\n--- Posterior Summary ({label}) ---")
    print(summary.to_string())

    # Convergence check
    rhat_ok = (summary["r_hat"] <= 1.05).all()
    ess_ok = (summary["ess_bulk"] >= 400).all()
    print(f"\nR-hat all <= 1.05: {rhat_ok}")
    print(f"ESS_bulk all >= 400: {ess_ok}")

    return {
        "label": label,
        "outcome": outcome,
        "model": model,
        "idata": idata,
        "summary": summary,
        "rhat_ok": rhat_ok,
        "ess_ok": ess_ok,
    }


# ── 3. Extract results and make tables ───────────────────────────────

def make_coefficient_table(results: list[dict]) -> pd.DataFrame:
    """Combine posterior summaries across models into one publication table."""
    rows = []
    for res in results:
        s = res["summary"]
        for param in s.index:
            row = {
                "model": res["label"],
                "parameter": param,
                "mean": round(s.loc[param, "mean"], 4),
                "sd": round(s.loc[param, "sd"], 4),
                "hdi_2.5%": round(s.loc[param, "hdi_2.5%"], 4),
                "hdi_97.5%": round(s.loc[param, "hdi_97.5%"], 4),
                "ess_bulk": int(s.loc[param, "ess_bulk"]),
                "ess_tail": int(s.loc[param, "ess_tail"]),
                "r_hat": round(s.loc[param, "r_hat"], 4),
            }
            rows.append(row)
    return pd.DataFrame(rows)


def make_diagnostics_table(results: list[dict]) -> pd.DataFrame:
    rows = []
    for res in results:
        s = res["summary"]
        rows.append({
            "model": res["label"],
            "outcome": res["outcome"],
            "n_obs": 48,
            "n_params": len(s),
            "all_rhat_ok": res["rhat_ok"],
            "all_ess_ok": res["ess_ok"],
            "min_ess_bulk": int(s["ess_bulk"].min()),
            "max_rhat": round(s["r_hat"].max(), 4),
        })
    return pd.DataFrame(rows)


# ── 4. Figures ────────────────────────────────────────────────────────

def plot_coefficients(results: list[dict], save_path: str):
    """Forest plot of posterior intervals for all game-level models."""
    fig, axes = plt.subplots(1, len(results), figsize=(5 * len(results), 6), sharey=False)
    if len(results) == 1:
        axes = [axes]

    predictor_labels = {
        "Intercept": "Intercept",
        "pce_numeric": "PCE (linear)",
        "pressure_ec": "Pressure (EC vs HP)",
        "length_centered": "Disc. Length (linear)",
        "length_sq": "Disc. Length (quadratic)",
        "pce_x_pressure": "PCE × Pressure",
        "pce_x_length": "PCE × Length",
    }

    for ax, res in zip(axes, results):
        s = res["summary"]
        # Filter to fixed effects only (exclude phi / kappa)
        params = [p for p in s.index if p in predictor_labels or p == "Intercept"]
        if not params:
            params = [p for p in s.index if "kappa" not in p.lower() and p != res["outcome"] + "_kappa"]

        # Filter summary to these params
        s_plot = s.loc[[p for p in params if p in s.index]]

        y_pos = np.arange(len(s_plot))
        means = s_plot["mean"].values
        lo = s_plot["hdi_2.5%"].values
        hi = s_plot["hdi_97.5%"].values

        labels = [predictor_labels.get(p, p) for p in s_plot.index]

        ax.hlines(y_pos, lo, hi, color="#2196F3", linewidth=2, alpha=0.8)
        ax.scatter(means, y_pos, color="#1565C0", s=50, zorder=5)
        ax.axvline(0, color="gray", linestyle="--", linewidth=0.8, alpha=0.6)
        ax.set_yticks(y_pos)
        ax.set_yticklabels(labels, fontsize=9)
        ax.set_xlabel("Posterior Mean (logit scale)", fontsize=10)
        ax.set_title(res["label"], fontsize=11, fontweight="bold")
        ax.grid(axis="x", alpha=0.3)

    fig.suptitle("Bayesian Beta Regression: Posterior Coefficients (95% HDI)", fontsize=13, fontweight="bold", y=1.02)
    fig.tight_layout()
    for ext in ["png", "pdf"]:
        fig.savefig(f"{save_path}.{ext}", dpi=300, bbox_inches="tight")
    plt.close(fig)
    print(f"Saved: {save_path}.png/.pdf")


def plot_interaction_predictions(df: pd.DataFrame, results: list[dict], save_path: str):
    """Plot predicted outcomes for key interactions: PCE × Pressure and PCE × Length."""
    fig, axes = plt.subplots(2, len(results), figsize=(5 * len(results), 9))
    if len(results) == 1:
        axes = axes.reshape(-1, 1)

    pce_vals = np.array([0.0, 0.25, 0.50, 0.75])
    pce_labels_str = ["0%", "25%", "50%", "75%"]

    for col_idx, res in enumerate(results):
        idata = res["idata"]
        posterior = idata.posterior

        # Extract coefficient samples (shape: chains × draws)
        def get_samples(name):
            return posterior[name].values.flatten()

        intercept = get_samples("Intercept")
        b_pce = get_samples("pce_numeric")
        b_press = get_samples("pressure_ec")
        b_len = get_samples("length_centered")
        b_len_sq = get_samples("length_sq")
        b_pce_press = get_samples("pce_x_pressure")
        b_pce_len = get_samples("pce_x_length")

        # --- Panel 1: PCE × Pressure ---
        ax = axes[0, col_idx]
        for press_val, press_label, color in [(0, "High Pressure", "#E53935"), (1, "Evidence Centric", "#1E88E5")]:
            means_pred = []
            lo_pred = []
            hi_pred = []
            for pce in pce_vals:
                eta = intercept + b_pce * pce + b_press * press_val + b_pce_press * pce * press_val
                # At length_centered=0 (medium), length_sq=0
                mu = 1.0 / (1.0 + np.exp(-eta))  # inverse logit
                means_pred.append(np.mean(mu))
                lo_pred.append(np.percentile(mu, 2.5))
                hi_pred.append(np.percentile(mu, 97.5))
            ax.plot(pce_vals, means_pred, "o-", label=press_label, color=color, linewidth=2, markersize=6)
            ax.fill_between(pce_vals, lo_pred, hi_pred, alpha=0.15, color=color)

        ax.set_xlabel("PCE Level", fontsize=10)
        ax.set_ylabel(f"Predicted {res['label']}", fontsize=10)
        ax.set_title(f"PCE × Pressure", fontsize=11)
        ax.set_xticks(pce_vals)
        ax.set_xticklabels(pce_labels_str)
        ax.legend(fontsize=8)
        ax.grid(alpha=0.3)

        # --- Panel 2: PCE × Length ---
        ax = axes[1, col_idx]
        length_vals = np.array([-2, 0, 2])  # centered
        length_labels_str = ["Short (2)", "Medium (4)", "Long (6)"]
        for length_val, length_label, color in zip(length_vals, length_labels_str, ["#FF9800", "#4CAF50", "#9C27B0"]):
            means_pred = []
            lo_pred = []
            hi_pred = []
            for pce in pce_vals:
                eta = intercept + b_pce * pce + b_len * length_val + b_len_sq * (length_val**2) + b_pce_len * pce * length_val
                # At pressure_ec=0 (HP reference)
                mu = 1.0 / (1.0 + np.exp(-eta))
                means_pred.append(np.mean(mu))
                lo_pred.append(np.percentile(mu, 2.5))
                hi_pred.append(np.percentile(mu, 97.5))
            ax.plot(pce_vals, means_pred, "o-", label=length_label, color=color, linewidth=2, markersize=6)
            ax.fill_between(pce_vals, lo_pred, hi_pred, alpha=0.15, color=color)

        ax.set_xlabel("PCE Level", fontsize=10)
        ax.set_ylabel(f"Predicted {res['label']}", fontsize=10)
        ax.set_title(f"PCE × Discussion Length", fontsize=11)
        ax.set_xticks(pce_vals)
        ax.set_xticklabels(pce_labels_str)
        ax.legend(fontsize=8)
        ax.grid(alpha=0.3)

    fig.suptitle("Model-Predicted Outcomes: Interaction Effects (95% CI)", fontsize=13, fontweight="bold", y=1.02)
    fig.tight_layout()
    for ext in ["png", "pdf"]:
        fig.savefig(f"{save_path}.{ext}", dpi=300, bbox_inches="tight")
    plt.close(fig)
    print(f"Saved: {save_path}.png/.pdf")


# ── 5. Main ──────────────────────────────────────────────────────────

def main():
    print("Loading data...")
    df = load_game_data()

    # Fit models
    results = []

    # Model 1: Consensus Strength
    res_cs = fit_beta_model(df, "consensus_strength_beta", "Consensus Strength")
    results.append(res_cs)

    # Model 2: Evidence Alignment
    res_ea = fit_beta_model(df, "evidence_alignment_beta", "Evidence Alignment")
    results.append(res_ea)

    # Model 3 (supplementary): SDI
    res_sdi = fit_beta_model(df, "semantic_decoupling_index_beta", "SDI (Supplementary)")
    results.append(res_sdi)

    # ── Tables ──
    coef_table = make_coefficient_table(results)
    coef_table.to_csv(os.path.join(TABLE_DIR, "table_bayesian_game_level_models.csv"), index=False)
    print(f"\nSaved coefficient table: {TABLE_DIR}/table_bayesian_game_level_models.csv")

    # Also save as markdown
    with open(os.path.join(TABLE_DIR, "table_bayesian_game_level_models.md"), "w") as f:
        f.write("# Bayesian Beta Regression: Game-Level Models\n\n")
        f.write(coef_table.to_markdown(index=False))
        f.write("\n\nNote: Coefficients on logit scale. HDI = highest density interval (95%).\n")
    print(f"Saved markdown table: {TABLE_DIR}/table_bayesian_game_level_models.md")

    diag_table = make_diagnostics_table(results)
    diag_table.to_csv(os.path.join(TABLE_DIR, "table_model_diagnostics.csv"), index=False)
    print(f"Saved diagnostics table: {TABLE_DIR}/table_model_diagnostics.csv")

    # ── Figures ──
    plot_coefficients(results, os.path.join(FIG_DIR, "fig_bayes_game_coefficients"))
    plot_interaction_predictions(df, results, os.path.join(FIG_DIR, "fig_bayes_interaction_predictions"))

    # ── Save InferenceData for later use ──
    for res in results:
        safe_name = res["outcome"].replace("_beta", "")
        az.to_netcdf(res["idata"], os.path.join(BASE, "scripts", f"idata_{safe_name}.nc"))

    print("\n" + "=" * 60)
    print("GAME-LEVEL MODELS COMPLETE")
    print("=" * 60)

    # Print key results summary
    for res in results:
        s = res["summary"]
        fixed = [p for p in s.index if "kappa" not in p.lower()]
        print(f"\n--- {res['label']} ---")
        for p in fixed:
            mean_val = s.loc[p, "mean"]
            lo = s.loc[p, "hdi_2.5%"]
            hi = s.loc[p, "hdi_97.5%"]
            sig = "*" if (lo > 0 and hi > 0) or (lo < 0 and hi < 0) else ""
            print(f"  {p:25s}: {mean_val:+.3f} [{lo:+.3f}, {hi:+.3f}] {sig}")


if __name__ == "__main__":
    main()
