# Independent recomputation of the SCD-sensitivity bounds, written from the paper's Methods only.
import pandas as pd, numpy as np
from scipy import stats

def read(name):
    return pd.read_sas(f"/data/{name}.xpt", format="xport")

rows = []
for s in "GH":
    demo = read(f"DEMO_{s}")[["SEQN", "RIDAGEYR", "RIDSTATR", "WTMEC2YR", "SDMVSTRA", "SDMVPSU"]]
    mcq = read(f"MCQ_{s}")[["SEQN", "MCQ084"]]
    cfq = read(f"CFQ_{s}")[["SEQN", "CFDDS"]]
    d = demo.merge(mcq, on="SEQN", how="left").merge(cfq, on="SEQN", how="left")
    d = d[(d.RIDSTATR == 2) & (d.WTMEC2YR > 0)].copy()
    d["cyc"] = s
    rows.append(d)
d = pd.concat(rows, ignore_index=True)
d["w"] = d.WTMEC2YR / 2
old = d.RIDAGEYR >= 60
known = d.MCQ084.isin([1, 2])
dom = old & known
yes = d.MCQ084 == 1
no = d.MCQ084 == 2
miss = d.CFDDS.isna()

def lin(num, den, frame=d):
    w = frame.w.values
    N = (w * num).sum(); D = (w * den).sum(); r = N / D
    z = w * (num.astype(float) - r * den.astype(float)) / D
    t = pd.DataFrame({"h": frame.SDMVSTRA.values, "p": frame.SDMVPSU.values, "z": z}).groupby(["h", "p"]).z.sum().reset_index()
    v = 0.0; npsu = 0; nstr = 0
    for h, g in t.groupby("h"):
        m = len(g); nstr += 1; npsu += m
        v += m / (m - 1) * ((g.z - g.z.mean()) ** 2).sum()
    se = np.sqrt(v); df = npsu - nstr
    q = stats.t.ppf(0.975, df)
    return r, se, r - q * se, r + q * se, df

for cut in (30, 40, 50):
    low = (~miss) & (d.CFDDS <= cut)
    A = (dom & low & yes).values; B = (dom & low & no).values
    U1 = (dom & miss & yes).values; U0 = (dom & miss & no).values
    print(f"cut {cut}")
    print("  observed", lin(A, A | B))
    print("  lower   ", lin(A, A | B | U0))
    print("  upper   ", lin(A | U1, A | B | U1))
print("n older examined", int(old.sum()), "known", int(dom.sum()), "missing dsst", int((dom & miss).sum()))
# How many with MCQ084 coded 7/9 (refused / don't know) among older examined
print("MCQ084 codes among older examined:", d.loc[old, "MCQ084"].value_counts(dropna=False).to_dict())
print("CFDDS range observed:", d.loc[dom & ~miss, "CFDDS"].min(), d.loc[dom & ~miss, "CFDDS"].max())
# Unweighted 25th percentile of CFDDS among older examined, and weighted
x = d.loc[old & ~miss, ["CFDDS", "w"]].sort_values("CFDDS")
cw = x.w.cumsum() / x.w.sum()
print("unweighted p25:", x.CFDDS.quantile(0.25), "weighted p25:", x.CFDDS[cw >= 0.25].iloc[0])
print("share at or below 40, weighted:", x.w[x.CFDDS <= 40].sum() / x.w.sum(), " below 40:", x.w[x.CFDDS < 40].sum() / x.w.sum())
# Sensitivity with strict < 40 cutoff
low = (~miss) & (d.CFDDS < 40)
A = (dom & low & yes).values; B = (dom & low & no).values; U1 = (dom & miss & yes).values
print("upper with strict <40:", lin(A | U1, A | B | U1))
