"""Design-weighted partial identification of SCD sensitivity; public NHANES only."""
import json, hashlib, itertools
from pathlib import Path
import numpy as np
import pandas as pd
from scipy.stats import t as student_t
ROOT=Path(__file__).resolve().parents[1]

def load():
    for x in json.loads((ROOT/'data/external.json').read_text()):
        b=(ROOT/x['path']).read_bytes()
        assert len(b)==x['bytes'] and 'sha256:'+hashlib.sha256(b).hexdigest()==x['sha256']
    dfs=[]; stratum_sets=[]
    for suffix in ['G','H']:
        frames={c:pd.read_sas(ROOT/f'data/nhanes/{c}_{suffix}.xpt',format='xport') for c in ['DEMO','MCQ','CFQ']}
        for f in frames.values(): assert f.SEQN.is_unique
        d=frames['DEMO'][['SEQN','RIDAGEYR','RIDSTATR','WTMEC2YR','SDMVSTRA','SDMVPSU']].copy()
        d=d.merge(frames['MCQ'][['SEQN','MCQ084']],on='SEQN',how='left',validate='one_to_one')
        c=frames['CFQ']; keep=['SEQN','CFDDS','CFDCSR','CFDAST','CFDCST1','CFDCST2','CFDCST3','CFASTAT','CFDDRNC','CFDDPP']
        d=d.merge(c[keep],on='SEQN',how='left',validate='one_to_one')
        d=d[(d.RIDSTATR==2)&(d.WTMEC2YR>0)].copy(); d['cycle']=suffix
        stratum_sets.append(set(d.SDMVSTRA));dfs.append(d)
    assert not stratum_sets[0]&stratum_sets[1], 'pooled strata must be distinct'
    d=pd.concat(dfs,ignore_index=True);assert d.SEQN.is_unique
    assert d[['WTMEC2YR','SDMVSTRA','SDMVPSU']].notna().all().all()
    assert d.CFDDS.dropna().between(0,133).all()
    d['weight']=d.WTMEC2YR/2
    return d

def ratio(d,num,den,w=None):
    w=d.weight.to_numpy() if w is None else np.asarray(w)
    num=np.asarray(num,dtype=float);den=np.asarray(den,dtype=float)
    D=float(w@den);r=float(w@num/D)
    z=w*(num-r*den)/D
    g=pd.DataFrame({'h':d.SDMVSTRA,'p':d.SDMVPSU,'z':z}).groupby(['h','p']).z.sum()
    var=0;df=0
    for _,v in g.groupby(level=0):
        m=len(v);assert m>1
        var+=m/(m-1)*float(((v-v.mean())**2).sum());df+=m-1
    se=float(np.sqrt(var));crit=float(student_t.ppf(.975,df))
    return {'estimate':r,'se':se,'ci_low':max(0,r-crit*se),'ci_high':min(1,r+crit*se),'df':df}

def indicators(d,cut,domain,allow_missing_scd=False):
    yes=(d.MCQ084==1).to_numpy();no=(d.MCQ084==2).to_numpy(); unk=~(yes|no)
    obs=d.CFDDS.notna().to_numpy();low=(d.CFDDS<=cut).to_numpy();miss=~obs
    a=domain&low&yes; b=domain&low&no; u1=domain&miss&yes;u0=domain&miss&no
    kl=domain&low&unk if allow_missing_scd else np.zeros(len(d),bool)
    km=domain&miss&unk if allow_missing_scd else np.zeros(len(d),bool)
    return {'observed':(a,a|b),'lower':(a,a|b|kl|u0|km),'upper':(a|u1|kl|km,a|b|kl|u1|km)}

def summarize(d,cut,domain,allow=False):
    inds=indicators(d,cut,domain,allow)
    return {name:ratio(d,*pair) for name,pair in inds.items()}

def main():
    d=load();old=(d.RIDAGEYR>=60).to_numpy();known=d.MCQ084.isin([1,2]).to_numpy();domain=old&known
    obs=d.CFDDS.notna().to_numpy();yes=(d.MCQ084==1).to_numpy()
    anytest=d[['CFDDS','CFDCSR','CFDAST']].notna().any(axis=1).to_numpy()|d[['CFDCST1','CFDCST2','CFDCST3']].notna().all(axis=1).to_numpy()
    primary=summarize(d,40,domain)
    p=primary['upper'];test_t=(p['estimate']-.5)/p['se'];pvalue=float(2*student_t.sf(abs(test_t),p['df']))
    out={'n_examined_older':int(old.sum()),'n_known_scd':int(domain.sum()),'n_missing_scd':int((old&~known).sum()),'n_dsst_observed':int((domain&obs).sum()),'n_dsst_missing':int((domain&~obs).sum()),'cutoff':40,'df':p['df'],'upper_test_t':test_t,'upper_test_p':pvalue,'upper_below_half':p['estimate']<.5,'upper_ci_below_half':p['ci_high']<.5}
    for k,v in primary.items():
        for stat,val in v.items():out[k+'_'+stat]=val
    out['missing_weight_share']=ratio(d,domain&~obs,domain)['estimate']
    secondary={f'cutoff_{c}':summarize(d,c,domain) for c in [30,40,50]}
    for s in ['G','H']:
        dc=d.loc[d.cycle==s].copy()
        secondary['cycle_'+s]=summarize(dc,40,((dc.RIDAGEYR>=60)&dc.MCQ084.isin([1,2])).to_numpy())
    secondary['any_test']=summarize(d,40,domain&anytest)
    secondary['missing_scd_allowed']=summarize(d,40,old,True)
    secondary['unweighted']={name:float(np.sum(n)/np.sum(de)) for name,(n,de) in indicators(d,40,domain).items()}
    counts=[]
    for scd,scmask in [('yes',d.MCQ084==1),('no',d.MCQ084==2),('unknown',~d.MCQ084.isin([1,2]))]:
        for status,mask in [('low',d.CFDDS<=40),('high',d.CFDDS>40),('missing',d.CFDDS.isna())]:
            m=old&scmask.to_numpy()&mask.to_numpy()
            counts.append({'scd':scd,'dsst':status,'n':int(m.sum()),'weight':float(d.loc[m,'weight'].sum())})
    prevalence={k:ratio(d,mask&yes,mask) for k,mask in [('observed',domain&obs),('missing',domain&~obs)]}
    reasons={var:{str(k):int(v) for k,v in d.loc[old&~obs,var].fillna(-1).value_counts().sort_index().items()} for var in ['CFASTAT','CFDDRNC','CFDDPP']}
    for k in ['observed','lower','upper']:
        for stat in ['estimate','se','ci_low','ci_high']:out[k+'_'+stat+'_pct']=round(100*out[k+'_'+stat],2)
    out['missing_weight_pct']=round(100*out['missing_weight_share'],2)
    out['primary_gap_below_half_pp']=round(100*(.5-out['upper_estimate']),2)
    r2={}
    for name,vals in secondary.items():
        r2[name]={k+'_pct':round(100*(v['estimate'] if isinstance(v,dict) else v),2) for k,v in vals.items()}
    for name,vals in prevalence.items():r2['scd_prevalence_'+name]={'estimate_pct':round(100*vals['estimate'],2),'ci_low_pct':round(100*vals['ci_low'],2),'ci_high_pct':round(100*vals['ci_high'],2)}
    results=ROOT/'results';results.mkdir(exist_ok=True)
    (results/'R2.json').write_text(json.dumps(r2,indent=2)+'\n')
    for name,data in [('R1',out),('sensitivity',secondary),('missingness_cells',counts),('noncompletion',{'reasons':reasons,'scd_prevalence':prevalence})]:
        (results/(name+'.json')).write_text(json.dumps(data,indent=2,allow_nan=False)+'\n')
    print(json.dumps(out,indent=2))
if __name__=='__main__':main()
