"""Pre-registered analysis of the sit-to-stand features (see plan/analysis-plan.md).

Usage: python code/analyze.py <features.csv> <GE-71_Data_Summary_Table.csv> <out.json>
"""
import json
import sys

import numpy as np
import pandas as pd
from scipy import stats

SEED = 20261007
N_BOOT = 10000


def hl_paired(d):
    """Hodges-Lehmann estimate for paired differences (median of Walsh averages)."""
    d = np.asarray(d)
    w = (d[:, None] + d[None, :]) / 2
    return float(np.median(w[np.triu_indices(len(d))]))


def boot_ci(fn, *arrays, rng):
    vals = []
    for _ in range(N_BOOT):
        vals.append(fn(*[a[rng.integers(0, len(a), len(a))] for a in arrays]))
    lo, hi = np.percentile(vals, [2.5, 97.5])
    return round(float(lo), 4), round(float(hi), 4)


def auc(pos, neg):
    """Probability that a value from pos exceeds one from neg (ties count half)."""
    pos, neg = np.asarray(pos), np.asarray(neg)
    gt = (pos[:, None] > neg[None, :]).mean()
    eq = (pos[:, None] == neg[None, :]).mean()
    return float(gt + 0.5 * eq)


def paired_test(d, rng):
    d = np.asarray(d, float)
    d = d[np.isfinite(d)]
    p = stats.wilcoxon(d, alternative="greater").pvalue
    return {"n": int(len(d)), "median_diff": round(float(np.median(d)), 4),
            "hodges_lehmann": round(hl_paired(d), 4),
            "hl_ci95": boot_ci(hl_paired, d, rng=rng),
            "dz": round(float(d.mean() / d.std(ddof=1)), 3),
            "n_positive": int((d > 0).sum()),
            "p_one_sided": round(float(p), 5)}


def group_test(dm, ctl, rng, alternative="greater"):
    dm, ctl = np.asarray(dm, float), np.asarray(ctl, float)
    dm, ctl = dm[np.isfinite(dm)], ctl[np.isfinite(ctl)]
    p = stats.mannwhitneyu(dm, ctl, alternative=alternative).pvalue
    return {"n_dm": int(len(dm)), "n_control": int(len(ctl)),
            "median_dm": round(float(np.median(dm)), 4), "median_control": round(float(np.median(ctl)), 4),
            "auc_dm_gt_control": round(auc(dm, ctl), 3),
            "auc_ci95": boot_ci(auc, dm, ctl, rng=rng),
            "p": round(float(p), 5), "alternative": alternative}


def loo_auc(X, y):
    """Leave-one-out AUC of a ridge-penalised logistic regression on standardised features."""
    from numpy.linalg import solve
    n = len(y)
    scores = np.empty(n)
    for i in range(n):
        tr = np.arange(n) != i
        mu, sd = X[tr].mean(0), X[tr].std(0, ddof=1)
        Z = (X[tr] - mu) / sd
        Z1 = np.column_stack([np.ones(tr.sum()), Z])
        w = np.zeros(Z1.shape[1])
        for _ in range(50):  # Newton iterations, L2 penalty 1 on slopes
            p = 1 / (1 + np.exp(-Z1 @ w))
            g = Z1.T @ (p - y[tr]) + np.r_[0, w[1:]]
            H = (Z1 * (p * (1 - p))[:, None]).T @ Z1 + np.diag(np.r_[0, np.ones(len(w) - 1)])
            w -= solve(H, g)
        zi = (X[i] - mu) / sd
        scores[i] = w[0] + zi @ w[1:]
    return round(auc(scores[y == 1], scores[y == 0]), 3)


def main():
    feats, table, out = sys.argv[1], sys.argv[2], sys.argv[3]
    rng = np.random.default_rng(SEED)
    f = pd.read_csv(feats)
    t = pd.read_csv(table, encoding="latin-1")
    t = t.assign(subject=t["SUBJECT NUMBER"].astype(str).str.upper().str.strip())
    t = t[["subject", "group2", "Group", "Age", "Dizziness"]]
    f = f.merge(t, on="subject", how="left")
    res = {"bouts_by_status": {str(k): int(v) for k, v in f["status"].value_counts().items()}}
    ok = f[(f["status"] == "ok") & f["group2"].isin(["Control", "DM"])].copy()
    for b in (1, 2):
        for v in ("sbp_ar1", "cbfv_ar1"):
            ok.loc[ok.bout == b, f"d_{v}"] = ok["stand_" + v] - ok["sit_" + v]
        ok.loc[ok.bout == b, "d_log_sbp_var"] = np.log(ok["stand_sbp_var"]) - np.log(ok["sit_sbp_var"])
        ok.loc[ok.bout == b, "d_sbp_mean"] = ok["stand_sbp_mean"] - ok["sit_sbp_mean"]
    b1 = ok[ok.bout == 1]
    b2 = ok[ok.bout == 2]
    dm, ctl = b1[b1.group2 == "DM"], b1[b1.group2 == "Control"]
    res["sample"] = {"subjects_bout1": int(len(b1)), "dm": int(len(dm)), "control": int(len(ctl)),
                     "subjects_bout2": int(len(b2))}
    # H1 (primary): standing raises SBP lag-1 autocorrelation (eyes open)
    res["H1_sbp_ar1_stand_minus_sit"] = paired_test(b1["d_sbp_ar1"], rng)
    # H2 (primary): the rise is larger with diabetes
    res["H2_sbp_ar1_rise_dm_vs_control"] = group_test(dm["d_sbp_ar1"], ctl["d_sbp_ar1"], rng)
    # Holm adjustment across the two primary tests
    ps = [res["H1_sbp_ar1_stand_minus_sit"]["p_one_sided"], res["H2_sbp_ar1_rise_dm_vs_control"]["p"]]
    order = np.argsort(ps)
    adj = [0.0, 0.0]
    running = 0.0
    for rank, i in enumerate(order):
        running = max(running, min(1.0, ps[i] * (2 - rank)))
        adj[i] = round(running, 5)
    res["holm_adjusted_p"] = {"H1": adj[0], "H2": adj[1]}
    # Secondary
    res["S1_log_sbp_var_stand_minus_sit"] = paired_test(b1["d_log_sbp_var"], rng)
    res["S2_log_sbp_var_rise_dm_vs_control"] = group_test(dm["d_log_sbp_var"], ctl["d_log_sbp_var"], rng)
    res["S3_cbfv_ar1_stand_minus_sit"] = paired_test(b1["d_cbfv_ar1"], rng)
    res["S4_cbfv_ar1_rise_dm_vs_control"] = group_test(dm["d_cbfv_ar1"], ctl["d_cbfv_ar1"], rng)
    res["S5_eyes_closed_sbp_ar1_stand_minus_sit"] = paired_test(b2["d_sbp_ar1"], rng)
    dm2, ctl2 = b2[b2.group2 == "DM"], b2[b2.group2 == "Control"]
    res["S6_eyes_closed_sbp_ar1_rise_dm_vs_control"] = group_test(dm2["d_sbp_ar1"], ctl2["d_sbp_ar1"], rng)
    # H3: sway-SBP coupling while standing exceeds phase-randomised surrogates
    z = b1["coh_sway_sbp_z"].to_numpy(float)
    z = z[np.isfinite(z)]
    res["H3_sway_sbp_coherence_z"] = {"n": int(len(z)), "median_z": round(float(np.median(z)), 3),
                                      "n_z_above_1_645": int((z > 1.645).sum()),
                                      "p_one_sided": round(float(stats.wilcoxon(z, alternative="greater").pvalue), 5)}
    res["S7_coherence_dm_vs_control"] = group_test(dm["coh_sway_sbp"], ctl["coh_sway_sbp"], rng, "two-sided")
    # Sensitivity: group effect on the H1 measure adjusted for age (rank regression)
    r = b1[["d_sbp_ar1", "group2", "Age"]].dropna()
    yv = stats.rankdata(r["d_sbp_ar1"])
    X = np.column_stack([np.ones(len(r)), (r["group2"] == "DM").astype(float), stats.rankdata(r["Age"])])
    beta, *_ = np.linalg.lstsq(X, yv, rcond=None)
    resid = yv - X @ beta
    s2 = resid @ resid / (len(yv) - 3)
    se = np.sqrt(np.diag(s2 * np.linalg.inv(X.T @ X)))
    res["S8_age_adjusted_rank_regression"] = {"n": int(len(r)), "dm_coef_ranks": round(float(beta[1]), 3),
                                              "t": round(float(beta[1] / se[1]), 3),
                                              "p_two_sided": round(float(2 * stats.t.sf(abs(beta[1] / se[1]), len(yv) - 3)), 5)}
    # E1 (exploratory): joint versus single-system separation of DM from controls
    cols = {"bp_ar1_rise": "d_sbp_ar1", "bp_fall": "d_sbp_mean", "sway_ap": "sway_ap_sd_g", "cbfv_ar1_rise": "d_cbfv_ar1"}
    e = b1[list(cols.values()) + ["group2"]].dropna()
    y = (e["group2"] == "DM").astype(int).to_numpy()
    res["E1_loo_auc"] = {"n": int(len(e))}
    for name, c in cols.items():
        res["E1_loo_auc"][name] = loo_auc(e[[c]].to_numpy(float), y)
    res["E1_loo_auc"]["joint"] = loo_auc(e[list(cols.values())].to_numpy(float), y)
    # Descriptive only: dizziness and the orthostatic-hypotension label are too rare to test
    res["descriptive_counts"] = {"dizziness_yes": int((b1["Dizziness"].astype(str).str.lower() == "yes").sum()),
                                 "group_DMOH": int((b1["Group"] == "DMOH").sum())}
    with open(out, "w") as fh:
        json.dump(res, fh, indent=2, sort_keys=True)
        fh.write("\n")
    print(json.dumps(res, indent=1))


if __name__ == "__main__":
    main()
