"""Clinical Data Lab: independent Python analysis of unchanged CDISC Pilot XPT.
Run: python analyze.py INPUT_DIRECTORY OUTPUT_DIRECTORY
No random sampling, imputation, fitting service, or patient upload.
"""
import sys, json, math, platform, importlib.metadata as metadata
from pathlib import Path
import numpy as np
import pandas as pd
from scipy.stats import t

def read(name):
    d = pd.read_sas(INPUT / (name + '.xpt'), format='xport', encoding='utf-8')
    for c in d.select_dtypes(include=['object','str']).columns:
        d[c] = d[c].fillna('').str.strip()
    # Legacy IBM XPORT numeric zero can decode as 16**-65 in pandas.
    # Normalize only nonzero subnormal XPORT placeholders, never missing values.
    for c in d.select_dtypes(include='number').columns:
        d.loc[d[c].abs().between(0, 1e-70, inclusive='neither'), c] = 0.0
    return d

def summary(values):
    x = np.asarray(pd.Series(values).dropna(),dtype=float)
    n = len(x)
    mean = float(x.mean()) if n else None
    sd = float(x.std(ddof=1)) if n > 1 else None
    margin = float(t.ppf(.975,n-1)*sd/math.sqrt(n)) if n>1 else None
    return dict(n=n,estimate=mean,sd=sd,lower=mean-margin if margin is not None else None,upper=mean+margin if margin is not None else None)

def welch(a,b):
    a,b=np.asarray(a.dropna()),np.asarray(b.dropna())
    if min(len(a),len(b))<2: return dict(estimate=None,lower=None,upper=None)
    va,vb=a.var(ddof=1)/len(a),b.var(ddof=1)/len(b)
    se=math.sqrt(va+vb); df=(va+vb)**2/(va**2/(len(a)-1)+vb**2/(len(b)-1))
    delta=float(a.mean()-b.mean()); margin=float(t.ppf(.975,df)*se)
    return dict(estimate=delta,lower=delta-margin,upper=delta+margin)

def km(times,status):
    times,status=np.asarray(times),np.asarray(status)
    s=1.; greenwood=0.; rows=[]
    for tm in sorted(set(times)):
        risk=int((times>=tm).sum()); events=int(((times==tm)&(status==1)).sum()); censored=int(((times==tm)&(status==0)).sum())
        s*=1-events/risk
        if events and risk>events: greenwood+=events/(risk*(risk-events))
        if s==1: lo=hi=1.
        elif s==0: lo=hi=None
        else:
            se=math.sqrt(greenwood)/abs(math.log(s)); z=1.959963984540054
            lo=math.exp(-math.exp(math.log(-math.log(s))+z*se)); hi=math.exp(-math.exp(math.log(-math.log(s))-z*se))
        rows.append(dict(time=float(tm),risk=risk,events=events,censored=censored,estimate=s,lower=lo,upper=hi))
    return rows

def run():
    adsl,ae,q,lab,tte=[read(n) for n in ['adsl','adae','adqsadas','adlbc','adtte']]
    assert adsl.USUBJID.is_unique
    assert set(tte.CNSR)=={0,1} and tte.USUBJID.is_unique
    assert (adsl.DISCONFL.eq('Y') == adsl.DCDECOD.ne('COMPLETED')).all()
    q=q[(q.PARAMCD=='ACTOT')&(q.ANL01FL=='Y')&(q.DTYPE=='')]
    assert not q.duplicated(['USUBJID','AVISIT']).any()
    # End of Treatment is an explicit analysis visit; do not reuse ANL01FL as a visit selector.
    lab=lab[(lab.PARAMCD=='ALT')&(lab.AVISIT=='End of Treatment')]
    assert not lab.duplicated(['USUBJID']).any()
    out={k:[] for k in ['baseline','disposition','longitudinal','adverse-events','laboratory','subgroups','time-to-event']}
    arms=['Placebo','Xanomeline Low Dose','Xanomeline High Dose']
    terms=sorted(set(zip(ae.AEBODSYS,ae.AEDECOD)))
    for stratum in ['All','F','M']:
        cohort=adsl if stratum=='All' else adsl[adsl.SEX==stratum]
        for arm in arms:
            subjects=cohort[(cohort.TRT01A==arm)&(cohort.SAFFL=='Y')]
            ids=set(subjects.USUBJID); N=len(ids)
            def add(case,measure,**kw):
                out[case].append(dict(stratum=stratum,arm=arm,measure=measure,**kw))
            for col,label in [('AGE','Age (years)'),('WEIGHTBL','Weight (kg)'),('BMIBL','BMI (kg/m^2)')]:
                stats=summary(subjects[col]); add('baseline',label,denominator=N,missing=N-stats['n'],**stats)
            for col,levels in [('SEX',['F','M']),('RACE',sorted(adsl.RACE.unique()))]:
                for level in levels:
                    n=int((subjects[col]==level).sum()); add('baseline',f'{col}: {level}',n=n,denominator=N,missing=int(subjects[col].eq('').sum()),estimate=100*n/N)
            add('disposition','Analysis cohort',n=N,denominator=N,estimate=100.)
            for reason in sorted(adsl.DCDECOD.unique()):
                n=int((subjects.DCDECOD==reason).sum()); add('disposition',reason,n=n,denominator=N,estimate=100*n/N)
            qi=q[q.USUBJID.isin(ids)]
            for visit in ['Baseline','Week 8','Week 16','Week 24']:
                for col,label in [('AVAL','ADAS-Cog(11) score'),('CHG','Change from baseline')]:
                    stats=summary(qi.loc[qi.AVISIT==visit,col]); add('longitudinal',label,visit=visit,time={'Baseline':0,'Week 8':8,'Week 16':16,'Week 24':24}[visit],denominator=N,missing=N-stats['n'],**stats)
            ai=ae[ae.USUBJID.isin(ids)&(ae.TRTEMFL=='Y')]
            for soc,pt in [('', 'Any treatment-emergent AE')]+[(s,'') for s in sorted(ae.AEBODSYS.unique())]+terms:
                subset=ai if not soc else ai[(ai.AEBODSYS==soc)&((ai.AEDECOD==pt) if pt else True)]
                n=int(subset.USUBJID.nunique()); add('adverse-events',pt or soc,soc=soc,level='Any' if not soc else ('PT' if pt else 'SOC'),n=n,events=len(subset),denominator=N,estimate=100*n/N)
            li=lab[lab.USUBJID.isin(ids)]
            paired=li[li.BNRIND.isin(['L','N','H'])&li.ANRIND.isin(['L','N','H'])]
            for base in ['L','N','H']:
                for follow in ['L','N','H']:
                    n=int(((paired.BNRIND==base)&(paired.ANRIND==follow)).sum()); add('laboratory',f'{base} -> {follow}',n=n,denominator=len(paired),missing=N-len(paired),estimate=100*n/len(paired) if len(paired) else None)
            ti=tte[tte.USUBJID.isin(ids)&(tte.PARAMCD=='TTDE')]
            assert ti.AVAL.notna().all() and len(ti)==N
            add('time-to-event','Event-free probability',time=0.,risk=N,events=0,censored=0,estimate=1.,lower=1.,upper=1.,denominator=N)
            for row in km(ti.AVAL,1-ti.CNSR): add('time-to-event','Event-free probability',denominator=N,**row)
        for label,mask in [('Overall',pd.Series(True,index=cohort.index)),('Age <65',cohort.AGE<65),('Age >=65',cohort.AGE>=65)]:
            group=cohort[mask & cohort.ITTFL.eq('Y')]
            end=q[(q.AVISIT=='Week 24')&q.USUBJID.isin(group.USUBJID)]
            placebo=end[end.USUBJID.isin(group[group.TRT01P=='Placebo'].USUBJID)].CHG.dropna()
            for arm in arms[1:]:
                active=end[end.USUBJID.isin(group[group.TRT01P==arm].USUBJID)].CHG.dropna()
                out['subgroups'].append(dict(stratum=stratum,arm=arm,measure=label,n=len(active),control_n=len(placebo),denominator=len(group[group.TRT01P==arm]),missing=len(group[group.TRT01P==arm])-len(active),**welch(active,placebo)))
    for case,rows in out.items():
        (OUTPUT/(case+'.json')).write_text(json.dumps(rows,ensure_ascii=False,allow_nan=False,indent=2)+'\n',encoding='utf8')
    fixture=summary([1,2,3,None]); assert fixture['n']==3 and fixture['estimate']==2 and fixture['sd']==1
    f=km([1,2,2,3],[1,1,0,1]); assert abs(f[1]['estimate']-.5)<1e-12 and f[1]['risk']==3
    (OUTPUT/'environment.json').write_text(json.dumps(dict(python=platform.python_version(),packages={p:metadata.version(p) for p in ['numpy','pandas','scipy']},fixtures='pass',seed='not applicable: deterministic'),indent=2),encoding='utf8')
    print(json.dumps({k:len(v) for k,v in out.items()}))

if __name__=='__main__':
    INPUT=Path(sys.argv[1]); OUTPUT=Path(sys.argv[2]); OUTPUT.mkdir(parents=True,exist_ok=True); run()
