#!/usr/bin/env python3
"""Preregistered TWFE DiD: pharmacist direct-authority naloxone x fentanyl dominance (WONDER AA rates)."""
from __future__ import annotations

import json
import math
from pathlib import Path

import numpy as np
import pandas as pd
import statsmodels.api as sm

ROOT = Path(__file__).resolve().parents[1]
DATA = ROOT / "data"
OUT = ROOT / "results"
OUT.mkdir(parents=True, exist_ok=True)

ATTENUATION_EXCLUSION = 5.0  # per 100k


def clustered_ols(y, X, groups):
    model = sm.OLS(y, X)
    fit = model.fit(cov_type="cluster", cov_kwds={"groups": groups})
    return fit


def twfe_design(df, extra_cols):
    d = df.copy()
    states = sorted(d["state"].unique())
    years = sorted(d["year"].unique())
    for s in states[1:]:
        d[f"S_{s}"] = (d["state"] == s).astype(float)
    for y in years[1:]:
        d[f"Y_{y}"] = (d["year"] == y).astype(float)
    fe = [c for c in d.columns if c.startswith("S_") or c.startswith("Y_")]
    cols = extra_cols + fe
    X = sm.add_constant(d[cols].astype(float), has_constant="add")
    y = d["y_opioid_rate"].astype(float)
    return y, X, d["state"].values


def main():
    panel_all = pd.read_csv(DATA / "analysis_panel.csv")
    # Locked primary sample: observed AA rate AND observed T40.4 share
    panel = panel_all[panel_all["in_primary"].astype(bool)].dropna(
        subset=["y_opioid_rate", "law_direct", "synthetic_dominant"]
    ).copy()
    panel["law"] = panel["law_direct"].astype(float)
    panel["dom"] = panel["synthetic_dominant"].astype(float)
    panel["law_x_dom"] = panel["law"] * panel["dom"]

    y, X, g = twfe_design(panel, ["law", "dom", "law_x_dom"])
    fit = clustered_ols(y, X, g)
    b1 = float(fit.params["law"])
    b0 = float(fit.params["dom"])
    b2 = float(fit.params["law_x_dom"])
    se1 = float(fit.bse["law"])
    se2 = float(fit.bse["law_x_dom"])
    z = 1.959963984540054
    ci1 = (b1 - z * se1, b1 + z * se1)
    ci2 = (b2 - z * se2, b2 + z * se2)
    att_dom = b1 + b2
    cov = float(fit.cov_params().loc["law", "law_x_dom"])
    se_dom = math.sqrt(max(se1**2 + se2**2 + 2 * cov, 0.0))
    ci_dom = (att_dom - z * se_dom, att_dom + z * se_dom)

    # Event study
    es = panel[panel["direct_auth_start_year"].notna()].copy()
    es["rel"] = es["year"] - es["direct_auth_start_year"]
    es["rel_bin"] = es["rel"].clip(-5, 5)
    rel_levels = [i for i in range(-5, 6) if i != -1]
    for r in rel_levels:
        es[f"R_{r}"] = (es["rel_bin"] == r).astype(float)
    never = panel[panel["direct_auth_start_year"].isna()].copy()
    for r in rel_levels:
        never[f"R_{r}"] = 0.0
    never["rel_bin"] = np.nan
    es_all = pd.concat([es, never], ignore_index=True)
    y_es, X_es, g_es = twfe_design(es_all, [f"R_{r}" for r in rel_levels])
    fit_es = clustered_ols(y_es, X_es, g_es)
    lead_names = [f"R_{r}" for r in [-4, -3, -2] if f"R_{r}" in fit_es.params.index]
    if lead_names:
        R = np.zeros((len(lead_names), len(fit_es.params)))
        for i, name in enumerate(lead_names):
            R[i, list(fit_es.params.index).index(name)] = 1.0
        ft = fit_es.f_test(R)
        pretrend_p = float(np.asarray(ft.pvalue).reshape(-1)[0])
    else:
        pretrend_p = float("nan")

    beta1_neg_sig = (b1 < 0) and (ci1[1] < 0)
    beta2_pos_sig = (b2 > 0) and (ci2[0] > 0)
    att_dom_excludes_large_reduction = ci_dom[0] > -ATTENUATION_EXCLUSION
    pretrend_ok = (not math.isnan(pretrend_p)) and (pretrend_p >= 0.05)

    support = bool(beta1_neg_sig and beta2_pos_sig and att_dom_excludes_large_reduction and pretrend_ok)
    no_pre_protection = not beta1_neg_sig
    no_attenuation = not (beta2_pos_sig and att_dom_excludes_large_reduction)
    if support:
        decision = "support"
    elif no_pre_protection or no_attenuation:
        if (not pretrend_ok) and beta1_neg_sig and (beta2_pos_sig or att_dom_excludes_large_reduction):
            decision = "inconclusive"
        else:
            decision = "refute"
    else:
        decision = "inconclusive"

    # Secondary standing-order model
    panel2 = panel.dropna(subset=["law_standing"]).copy()
    panel2["law_s"] = panel2["law_standing"].astype(float)
    panel2["law_s_x_dom"] = panel2["law_s"] * panel2["dom"]
    y2, X2, g2 = twfe_design(panel2, ["law_s", "dom", "law_s_x_dom"])
    fit2 = clustered_ols(y2, X2, g2)

    # Secondary: population-weighted primary model
    # Use WLS with population weights via sm.WLS
    y_w, X_w, g_w = twfe_design(panel, ["law", "dom", "law_x_dom"])
    w = panel["population"].astype(float).values
    fit_w = sm.WLS(y_w, X_w, weights=w).fit(cov_type="cluster", cov_kwds={"groups": g_w})

    # Sensitivity: code unavailable T40.4 share as SyntheticDominant=0 among AA-observed rows
    sens = panel_all[panel_all["aa_ok"].astype(bool)].copy()
    sens["law"] = sens["law_direct"].astype(float)
    sens["dom"] = sens["synthetic_dominant"].fillna(0.0).astype(float)
    # for share-unavailable, force 0 per prereg sensitivity
    sens.loc[sens["share_ok"] == False, "dom"] = 0.0
    sens["law_x_dom"] = sens["law"] * sens["dom"]
    sens = sens.dropna(subset=["y_opioid_rate"]).copy()
    y_s, X_s, g_s = twfe_design(sens, ["law", "dom", "law_x_dom"])
    fit_s = clustered_ols(y_s, X_s, g_s)

    event = {
        str(r): {
            "coef": float(fit_es.params.get(f"R_{r}", 0.0)) if r != -1 else 0.0,
            "se": float(fit_es.bse.get(f"R_{r}", float("nan"))) if r != -1 else 0.0,
        }
        for r in range(-5, 6)
    }

    R1 = {
        "n_obs": int(len(panel)),
        "n_states": int(panel["state"].nunique()),
        "n_dropped": int((~panel_all["in_primary"].astype(bool)).sum()),
        "year_min": int(panel["year"].min()),
        "year_max": int(panel["year"].max()),
        "beta1_law": round(b1, 6),
        "beta1_se": round(se1, 6),
        "beta1_ci95_low": round(ci1[0], 6),
        "beta1_ci95_high": round(ci1[1], 6),
        "beta2_interaction": round(b2, 6),
        "beta2_se": round(se2, 6),
        "beta2_ci95_low": round(ci2[0], 6),
        "beta2_ci95_high": round(ci2[1], 6),
        "beta0_synthetic_dominant": round(b0, 6),
        "att_non_dominant": round(b1, 6),
        "att_dominant": round(att_dom, 6),
        "att_dominant_se": round(se_dom, 6),
        "att_dominant_ci95_low": round(ci_dom[0], 6),
        "att_dominant_ci95_high": round(ci_dom[1], 6),
        "pretrend_joint_p": None if math.isnan(pretrend_p) else round(pretrend_p, 6),
        "attenuation_exclusion_per_100k": ATTENUATION_EXCLUSION,
        "beta1_negative_sig": beta1_neg_sig,
        "beta2_positive_sig": beta2_pos_sig,
        "att_dominant_ci_excludes_reduction_ge_5": att_dom_excludes_large_reduction,
        "pretrend_ok": pretrend_ok,
        "decision": decision,
        "support_attenuation_claim": support,
        "secondary_standing_beta1": round(float(fit2.params["law_s"]), 6),
        "secondary_standing_beta2": round(float(fit2.params["law_s_x_dom"]), 6),
        "secondary_popwt_beta1": round(float(fit_w.params["law"]), 6),
        "secondary_popwt_beta2": round(float(fit_w.params["law_x_dom"]), 6),
        "sensitivity_dom0_beta1": round(float(fit_s.params["law"]), 6),
        "sensitivity_dom0_beta2": round(float(fit_s.params["law_x_dom"]), 6),
        "sensitivity_dom0_n": int(len(sens)),
        "outcome": "wonder_age_adjusted_opioid_od_deaths_per_100k",
        "exposure": "optic_nal_Rx_prescriptive_auth",
        "prereg": "prereg:aad4dbc7aff96bc09b242d7ed7cff2628fb50f0f1a47ac7a128f8509bafab8a9",
    }
    (OUT / "R1.json").write_text(json.dumps(R1, indent=2, sort_keys=True) + "\n")
    (OUT / "event_study.json").write_text(json.dumps(event, indent=2, sort_keys=True) + "\n")
    panel.to_csv(OUT / "analysis_used.csv", index=False)
    print(json.dumps(R1, indent=2))


if __name__ == "__main__":
    main()
