"""Independent small-table extremum checks and ratio-variance cross-checks."""
import itertools,json,math
from pathlib import Path
import numpy as np,pandas as pd
from analyze import indicators,ratio,load

def brute(df):
    choices=[]
    for row in df.itertuples():
        s=[int(row.MCQ084==1)] if row.MCQ084 in [1,2] else [0,1]
        l=[int(row.CFDDS<=40)] if not math.isnan(row.CFDDS) else [0,1]
        choices.append(list(itertools.product(s,l)))
    vals=[]
    for labels in itertools.product(*choices):
        den=sum(w*l for w,(s,l) in zip(df.weight,labels))
        if den: vals.append(sum(w*s*l for w,(s,l) in zip(df.weight,labels))/den)
    return min(vals),max(vals),len(vals)

def main():
    base=pd.DataFrame({'MCQ084':[1.,2.,1.,2.,np.nan,np.nan],'CFDDS':[20.,20.,np.nan,np.nan,20.,np.nan],'weight':[1.,3.,2.,4.,5.,6.],'SDMVSTRA':[1,1,1,1,2,2],'SDMVPSU':[1,2,1,2,1,2]})
    tested=0
    for allow in [False,True]:
        d=base if allow else base.iloc[:4].copy()
        vals=brute(d); inds=indicators(d,40,np.ones(len(d),bool),allow)
        bounds=[float(np.dot(d.weight,inds[k][0])/np.dot(d.weight,inds[k][1])) for k in ['lower','upper']]
        assert np.allclose(bounds,vals[:2],rtol=0,atol=1e-14);tested+=vals[2]
    # Hand example: A=1,B=3,U1=2,U0=4 -> [1/8,3/6].
    assert brute(base.iloc[:4])[:2]==(.125,.5)
    d=load();domain=((d.RIDAGEYR>=60)&d.MCQ084.isin([1,2])).to_numpy()
    realchecks={}
    for endpoint,(n,de) in indicators(d,40,domain).items():
        # Loop-based alternative to vectorized ratio and pandas aggregation.
        W=list(map(float,d.weight));N=sum(w for w,a in zip(W,n) if a);D=sum(w for w,b in zip(W,de) if b);r=N/D
        totals={}
        for h,p,w,a,b in zip(d.SDMVSTRA,d.SDMVPSU,W,n,de):
            totals[(h,p)]=totals.get((h,p),0)+w*(int(a)-r*int(b))/D
        hs=sorted(set(h for h,p in totals));var=0
        for h in hs:
            ps=sorted(p for hh,p in totals if h==hh);assert len(ps)>1
            var+=sum((totals[(h,p)]-totals[(h,q)])**2 for p,q in itertools.combinations(ps,2))/(len(ps)-1)
        got=ratio(d,n,de)
        assert abs(r-got['estimate'])<1e-12 and abs(math.sqrt(var)-got['se'])<1e-12
        realchecks[endpoint]={'estimate':r,'se':math.sqrt(var)}
    out={'synthetic_assignments_checked':tested,'hand_example_passed':True,'independent_real_ratios_and_psu_variance_match':True,'independent_estimates':realchecks}
    p=Path('results');p.mkdir(exist_ok=True);(p/'validation.json').write_text(json.dumps(out,indent=2)+'\n');print(json.dumps(out))
if __name__=='__main__':main()
