"""SeeSway browser preview 0.2.0. New definitions, not legacy-equivalence certified."""
import csv,io,json,math
import numpy as np
VERSION='0.2.0'
def analyse(payload):
    text=payload.get('csv','').lstrip('\ufeff')
    if len(text)>5_000_000: raise ValueError('File exceeds the 5 MB analysis limit.')
    fs=float(payload.get('sample_rate',100)); unit=payload.get('unit','mm'); method=payload.get('filter','none')
    if not math.isfinite(fs) or fs<=0 or fs>10000: raise ValueError('Sample rate must be greater than 0 and no more than 10000 Hz.')
    if unit not in ('mm','cm','m'): raise ValueError('Select mm, cm or m.')
    delimiter=payload.get('delimiter',',')
    if delimiter not in (',',';','\t'): raise ValueError('Choose comma, semicolon or tab.')
    columns=payload.get('columns',[0,1]);header_mode=payload.get('header','auto')
    if not isinstance(columns,list) or len(columns)!=2 or any(not isinstance(c,int) or isinstance(c,bool) or c<0 or c>99 for c in columns) or columns[0]==columns[1]: raise ValueError('Choose two different column numbers from 1 to 100.')
    if header_mode not in ('auto','yes','no'): raise ValueError('Choose the header setting.')
    rows=[]; header=None; expected=None
    for line,row in enumerate(csv.reader(io.StringIO(text),delimiter=delimiter),1):
        if not row or all(not c.strip() for c in row):continue
        if expected is None:expected=len(row)
        if len(row)!=expected or len(row)<=max(columns):raise ValueError(f'Row {line}: expected a consistent column count, with at least ML and AP.')
        if not rows and header is None:
            names=[row[c].strip().lower() for c in columns]
            if header_mode=='yes' or (header_mode=='auto' and names[0] in ('ml','medial-lateral','ml_mm') and names[1] in ('ap','anterior-posterior','ap_mm')):
                header=row;continue
        try: values=[float(v.strip().replace(',','.')) if delimiter!=',' else float(v.strip()) for v in [row[c] for c in columns]]
        except ValueError:raise ValueError(f'Row {line}: ML and AP must be numbers. Optional header: ML,AP.')
        if not all(math.isfinite(v) for v in values):raise ValueError(f'Row {line}: missing or non-finite values are not accepted.')
        rows.append(values)
        if len(rows)>100000:raise ValueError('Maximum 100000 samples per analysis.')
    if len(rows)<16:raise ValueError('At least 16 valid samples are required.')
    x=np.asarray(rows,dtype=float)*{'mm':1.,'cm':10.,'m':1000.}[unit]
    y=x.copy();settings={'method':method,'sample_rate_hz':fs,'input_unit':unit,'output_unit':'mm','duration_definition':'(N - 1) / sampling rate','columns_zero_based':columns,'column_names':[header[c] for c in columns] if header else None,'header_mode':header_mode,'unused_columns_zero_based':[c for c in range(expected) if c not in columns],'extra_columns':'Unused columns ignored; no vertical-force measures'}
    if method=='moving':
        raw_width=float(payload.get('window',5))
        if not math.isfinite(raw_width) or not raw_width.is_integer():raise ValueError('Moving-average window must be a whole number.')
        width=int(raw_width)
        if width<3 or width>101 or width%2!=1 or width>=len(x):raise ValueError('Moving-average window must be odd, 3–101, and shorter than the trace.')
        k=width//2;y=np.column_stack([np.convolve(np.pad(x[:,i],(k,k),mode='reflect'),np.ones(width)/width,mode='valid') for i in range(2)])
        settings.update(window_samples=width,padding='reflect, endpoints excluded')
    elif method=='butterworth':
        from scipy.signal import butter,sosfiltfilt
        cutoff=float(payload.get('cutoff',10));order=int(payload.get('order',6))
        if not math.isfinite(cutoff) or cutoff<=0 or cutoff>=fs/2:raise ValueError('Cutoff must be positive and below half the sample rate (Nyquist).')
        if order<1 or order>10:raise ValueError('Butterworth order must be from 1 to 10.')
        sos=butter(order,cutoff,fs=fs,output='sos')
        try:y=sosfiltfilt(sos,x,axis=0,padtype='odd')
        except ValueError:raise ValueError('Trace is too short for this filter order. Use more samples or a lower order.')
        settings.update(cutoff_hz=cutoff,design_order=order,passes='forward/backward',padding='SciPy odd extension; default SOS pad length')
    elif method=='wavelet':
        import pywt
        level=3
        margin=(pywt.Wavelet('sym8').dec_len-1)*(2**level-1)
        if len(x)<2*margin+1: raise ValueError('Symlet-8 filtering requires at least 211 samples.')
        extra=(-(len(x)+2*margin))%(2**level)
        padded=np.pad(x,((margin,margin+extra),(0,0)),mode='reflect')
        coeffs=pywt.swt(padded,'sym8',level=level,axis=0)
        approximation=[(a,np.zeros_like(d)) for a,d in coeffs]
        y=pywt.iswt(approximation,'sym8',axis=0)[margin:margin+len(x)]
        settings.update(wavelet='sym8',wavelet_level=level,transform='stationary undecimated; inverse reconstruction with all detail coefficients zero',nominal_upper_band_hz=fs/16,padding='reflect, endpoints excluded; 105 samples each end plus right padding to multiple of 8',padding_right_extra=extra)
    elif method!='none':raise ValueError('Unknown filter.')
    if not np.isfinite(y).all():raise ValueError('Analysis produced non-finite values; check input scale and settings.')
    duration=(len(y)-1)/fs;d=np.diff(y,axis=0);centre=y-y.mean(axis=0);path=float(np.linalg.norm(d,axis=1).sum())
    metrics={'samples':len(y),'duration_s':duration,'resultant_path_mm':path,'mean_resultant_speed_mm_s':path/duration}
    for i,axis in enumerate(('ML','AP')):
        metrics.update({f'{axis}_path_mm':float(np.abs(d[:,i]).sum()),f'{axis}_range_mm':float(np.ptp(y[:,i])),f'{axis}_rms_about_mean_mm':float(np.sqrt(np.mean(centre[:,i]**2))),f'{axis}_rms_about_zero_mm':float(np.sqrt(np.mean(y[:,i]**2))),f'{axis}_sd_sample_mm':float(np.std(y[:,i],ddof=1))})
    raw_path=float(np.linalg.norm(np.diff(x,axis=0),axis=1).sum())
    baseline_metrics={'resultant_path_mm':raw_path,'mean_resultant_speed_mm_s':raw_path/duration}
    bands=wavelet_bands(x,fs) if payload.get('wavelet_bands',True) else None
    dfa=None
    if payload.get('dfa',False):
        if payload.get('dfa') is not True: raise ValueError('DFA selection must be a boolean.')
        dfa={'ML':dfa_analysis(x[:,0],fs,payload), 'AP':dfa_analysis(x[:,1],fs,payload)}
        settings['dfa']={'input':'unfiltered displacement in mm', 'order':1, 'integration':'cumulative sum after subtracting mean', 'segments':'non-overlapping, forward and backward; duplicate segments omitted when exactly divisible', 'scale_spacing':'up to 20 unique rounded logarithmic integer scales', 'fit':'unweighted least-squares log10 F versus log10 window size', 'min_samples':dfa['ML']['scales_samples'][0], 'max_samples':dfa['ML']['scales_samples'][-1]}
    step=max(1,math.ceil(len(y)/2000));idx=list(range(0,len(y),step))
    if idx[-1]!=len(y)-1:idx.append(len(y)-1)
    return {'version':VERSION,'status':'research_preview_not_legacy_equivalent','settings':settings,'metrics':metrics,'baseline_metrics':baseline_metrics,'dfa':dfa,'wavelet_bands':bands,'plot':{'time':[i/fs for i in idx],'raw':x[idx].tolist(),'filtered':y[idx].tolist(),'display_decimation':step},'filtered':y.tolist(),'warnings':['New implementation: equivalence to 2018 SeeSway has not been established.','Plots may be decimated for display; metrics use every sample.','Only ML/AP displacement is analysed. Extra columns are ignored.']}


def dfa_analysis(signal,fs,payload):
    """DFA-1, explicit new method; not an implementation specified in the ACLR paper."""
    n=len(signal)
    if n<128: raise ValueError('DFA requires at least 128 samples.')
    lo=float(payload.get('dfa_min',16));hi=float(payload.get('dfa_max',n//4))
    if not all(math.isfinite(v) and v.is_integer() for v in (lo,hi)): raise ValueError('DFA window limits must be whole sample counts.')
    lo,hi=int(lo),int(hi)
    if lo<4 or hi>n//4 or hi<lo*2: raise ValueError('DFA windows must start at 4 or more samples, span at least a factor of two, and end at N/4 or less.')
    scales=np.unique(np.rint(np.geomspace(lo,hi,20)).astype(int))
    if len(scales)<6: raise ValueError('DFA needs at least six distinct window sizes.')
    profile=np.cumsum(signal-np.mean(signal));fluct=[]
    for size in scales:
        count=n//size
        windows=profile[:count*size].reshape(count,size)
        if n%size: windows=np.concatenate([windows,profile[n-count*size:].reshape(count,size)],axis=0)
        t=np.arange(size,dtype=float);t-=t.mean()
        centred=windows-windows.mean(axis=1,keepdims=True)
        slopes=centred@t/(t@t)
        residual=centred-slopes[:,None]*t
        fluct.append(float(np.sqrt(np.mean(residual**2))))
    result={'scales_samples':scales.tolist(),'scales_seconds':(scales/fs).tolist(),'fluctuation_mm':fluct,'alpha':None,'r_squared':None,'status':'undefined: constant or numerically degenerate signal'}
    tolerance=np.finfo(float).eps*max(1,float(np.max(np.abs(profile))))*100
    if any(v<=tolerance for v in fluct): return result
    lx=np.log10(scales);ly=np.log10(fluct)
    alpha,intercept=np.polyfit(lx,ly,1);pred=alpha*lx+intercept
    total=float(np.sum((ly-ly.mean())**2));r2=None if total<=np.finfo(float).eps else float(1-np.sum((ly-pred)**2)/total)
    result.update(alpha=float(alpha),intercept=float(intercept),r_squared=r2,status='computed; inspect scaling curve and selected range')
    return result


def wavelet_bands(x,fs):
    """Nine-level decimated Symlet-8 MRA; grouped as Clark et al. 2014."""
    import pywt,warnings
    if len(x)<512: return {'status':'At least 512 samples are required for the four wavelet bands.','bands':[]}
    # Paper first applies level-3 undecimated approximation smoothing.
    margin=105;extra=(-(len(x)+2*margin))%8
    padded=np.pad(x,((margin,margin+extra),(0,0)),mode='reflect')
    swt=pywt.swt(padded,'sym8',level=3,axis=0)
    smooth=pywt.iswt([(a,np.zeros_like(d)) for a,d in swt],'sym8',axis=0)[margin:margin+len(x)]
    with warnings.catch_warnings():
        warnings.simplefilter('ignore',UserWarning)
        coeffs=pywt.wavedec(smooth,'sym8',mode='symmetric',level=9,axis=0)
    definitions=[('Moderate', [5,6],fs/64,fs/16),('Low',[3,4],fs/256,fs/64),('Very low',[1,2],fs/1024,fs/256),('Ultralow',[0],0,fs/1024)]
    bands=[]
    for name,indices,lo,hi in definitions:
        selected=[c if i in indices else np.zeros_like(c) for i,c in enumerate(coeffs)]
        trace=pywt.waverec(selected,'sym8',mode='symmetric',axis=0)[:len(x)]
        path=np.abs(np.diff(trace,axis=0)).sum(axis=0)
        bands.append({'name':name,'lower_hz':lo,'upper_hz':hi,'trace_mm':trace.tolist(),'ML_path_speed_mm_s':float(path[0]/((len(x)-1)/fs)),'AP_path_speed_mm_s':float(path[1]/((len(x)-1)/fs))})
    return {'status':'computed','wavelet':'sym8','level':9,'input':'raw displacement, then level-3 undecimated Symlet-8 approximation (independent of main filter selection)','transform':'decimated DWT multiresolution reconstruction; D4+D5, D6+D7, D8+D9, A9','padding':'symmetric DWT; reflected padding for initial stationary filter','boundary_warning':'Coarse bands may be dominated by boundary extension in short recordings. Frequency limits are nominal, not sharp cutoffs. No physiological diagnosis can be inferred from a band. LabVIEW equivalence is unverified.','bands':bands}
