#!/usr/bin/env python3
"""Preregistered TWFE DiD: pharmacist direct-authority naloxone x fentanyl dominance."""
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):
    """OLS with state-clustered HC1-like SE via statsmodels."""
    model = sm.OLS(y, X)
    fit = model.fit(cov_type="cluster", cov_kwds={"groups": groups})
    return fit


def twfe_design(df, extra_cols):
    """Build y, X with state and year dummies; drop one state and one year."""
    d = df.copy()
    states = sorted(d["state"].unique())
    years = sorted(d["year"].unique())
    # drop first state, first year as reference
    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 = pd.read_csv(DATA / "analysis_panel.csv")
    panel = panel.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"]

    # Primary interaction model
    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"])
    # 95% CI
    z = 1.959963984540054
    ci1 = (b1 - z * se1, b1 + z * se1)
    ci2 = (b2 - z * se2, b2 + z * se2)
    att_dom = b1 + b2
    # SE of sum
    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 relative years for law_direct (among states with start year)
    es = panel.copy()
    es = es[es["direct_auth_start_year"].notna()].copy()
    es["rel"] = es["year"] - es["direct_auth_start_year"]
    # bin <=-5 and >=5
    es["rel_bin"] = es["rel"].clip(-5, 5)
    # dummies excluding -1
    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)
    # also include never-treated as zeros on all R_
    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)
    # joint test leads -4,-3,-2
    lead_names = [f"R_{r}" for r in [-4, -3, -2] if f"R_{r}" in fit_es.params.index]
    if lead_names:
        hyp = " = ".join(lead_names) + " = 0"
        # use f_test
        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")

    # Success rules (prereg)
    beta1_neg_sig = (b1 < 0) and (ci1[1] < 0)
    beta2_pos_sig = (b2 > 0) and (ci2[0] > 0)
    # ATT_dominant 95% CI excludes a reduction as large as -5: lower bound > -5
    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)
    # refute if no pre-dom protection OR no attenuation as defined
    refute = bool(
        (b1 >= 0 or ci1[0] <= 0 and ci1[1] >= 0 and not beta1_neg_sig)
        and True
    )
    # clearer refute:
    # 1) no protective pre-dominance (not beta1_neg_sig)
    # OR 2) attenuation fails: not (beta2_pos_sig and att_dom_excludes_large_reduction)
    no_pre_protection = not beta1_neg_sig
    no_attenuation = not (beta2_pos_sig and att_dom_excludes_large_reduction)
    # For negative_result we need clear refute of the attenuation claim
    # Prereg: Refute if beta1 fails OR attenuation definition fails OR interaction contradicts
    if support:
        decision = "support"
    elif no_pre_protection or no_attenuation:
        # if pretrends fail but otherwise mixed -> inconclusive
        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 (reported only)
    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)

    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()),
        "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),
        "outcome": "crude_opioid_od_deaths_per_100k",
        "exposure": "optic_nal_Rx_prescriptive_auth",
    }
    (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()
