import csv,json,sys,math
from pathlib import Path
import numpy as np
from scipy import signal,stats
import wfdb
sys.path.insert(0,'/bundle/code')
import features
p=Path('/bundle');labels={r['SUBJECT NUMBER'].strip().upper():r['group2'] for r in csv.DictReader((p/'data/GE-71_Data_Summary_Table.csv').open(encoding='latin-1'))};rows=[r for r in csv.DictReader((p/'results/features.csv').open()) if r['bout']=='1' and r['status']=='ok' and labels.get(r['subject']) in ['DM','Control']]
observed=[];matched_time=[];matched_beats=[];pvals=[];emp_z=[];nullz=[];rng=np.random.default_rng(48101)
for r in rows:
 rec=wfdb.rdrecord(str(p/'data/sit-to-stand'/('s'+r['subject'][1:]+'DC')));s0=float(r['stand_onset_s']);s1=float(r['stand_end_s']);a=s0+30;b=s1-5;abp=rec.p_signal[:,rec.sig_name.index('abp')]
 sit=features.beats(abp,rec.fs,s0-250,s0-10);stand=features.beats(abp,rec.fs,a,b);sitv=sit[1][sit[2]];stv=stand[1][stand[2]];tm=stand[0][stand[2]]
 observed.append(features.ews(stv)[0]-features.ews(sitv)[0]);match=features.beats(abp,rec.fs,s0-10-(b-a),s0-10);mv=match[1][match[2]];matched_time.append(features.ews(stv)[0]-features.ews(mv)[0]);n=min(len(stv),len(sitv));matched_beats.append(features.ews(stv[:n])[0]-features.ews(sitv[-n:])[0])
 start,end=int(a*rec.fs),int(b*rec.fs);fz=np.mean(rec.p_signal[start:end,rec.sig_name.index('fz')]);ap=rec.p_signal[start:end,rec.sig_name.index('fy')]/fz
 if np.isnan(ap).mean()>.1 or np.nanstd(ap)==0:continue
 ap4=signal.sosfiltfilt(signal.butter(4,1,fs=rec.fs,output='sos'),np.nan_to_num(ap,nan=np.nanmedian(ap)))[::int(rec.fs/4)];grid=a+np.arange(len(ap4))/4;keep=(grid>=tm[0])&(grid<=tm[-1]);x=signal.detrend(ap4[keep]);y=signal.detrend(np.interp(grid[keep],tm,stv))
 if len(x)<512:continue
 f,c=signal.coherence(x,y,fs=4,nperseg=256);band=(f>=.05)&(f<=.15);coh=c[band].mean();Y=np.fft.rfft(y);sur=[]
 for k in range(1000):
  ph=np.exp(1j*rng.uniform(0,2*np.pi,len(Y)));ph[0]=1;ys=np.fft.irfft(np.abs(Y)*ph,n=len(y));_,cs=signal.coherence(x,ys,fs=4,nperseg=256);sur.append(cs[band].mean())
 sur=np.array(sur);pvals.append(float((1+(sur>=coh).sum())/(len(sur)+1)));emp_z.append(float((coh-sur.mean())/sur.std(ddof=1)));nullz.append((sur-sur.mean())/sur.std(ddof=1))
def summarize(v):
 v=np.asarray(v);v=v[np.isfinite(v)];n=len(v);walsh=[(v[i]+v[j])/2 for i in range(n) for j in range(i,n)];return {'n':n,'HL':float(np.median(walsh)),'median':float(np.median(v)),'positive':int((v>0).sum()),'wilcoxon_p':float(stats.wilcoxon(v,alternative='greater').pvalue),'sign_test_p':float(stats.binomtest(int((v>0).sum()),n,.5,alternative='greater').pvalue)}
# Calibrate the median across participants directly against independently sampled surrogate values.
medobs=float(np.median(emp_z));mednull=np.median(np.array([rng.choice(z,size=20000) for z in nullz]),axis=0)
out={'baseline':summarize(observed),'matched_time':summarize(matched_time),'matched_beats':summarize(matched_beats),'surrogates':{'n':len(pvals),'per_subject':1000,'nominal_p_below_05':sum(x<.05 for x in pvals),'minimum_empirical_p':min(pvals),'fisher_combined_p':float(stats.combine_pvalues(pvals).pvalue),'median_z':medobs,'calibrated_median_p':float((1+(mednull>=medobs).sum())/(len(mednull)+1)),'pvalues':pvals}}
print(json.dumps(out,indent=2))
