import io, json, zipfile, hashlib, platform
from pathlib import Path
from collections import Counter
import numpy as np
import pandas as pd
import scipy
from scipy.stats import false_discovery_control

SEED=20261009
OUT=Path("optotagging-results")
OUT.mkdir(exist_ok=True)
ZIP_PATH=Path("inputs/optotagging-data-c64e8857f006.zip")
z=zipfile.ZipFile(ZIP_PATH)
manifest=json.loads(z.read("SHA256SUMS.json"))
assert all(hashlib.sha256(z.read(n)).hexdigest()==h for n,h in manifest.items())
DATA={}; checks={"manifest_members_verified":len(manifest),"source_zip_sha256":hashlib.sha256(ZIP_PATH.read_bytes()).hexdigest(),"sessions":{}}
for group in ("sst","wt"):
    m=json.loads(z.read(group+"-metadata.json"))
    u=pd.read_csv(z.open(group+"-units.csv"))
    e=pd.read_csv(z.open(group+"-electrodes.csv"))
    s=pd.read_csv(z.open(group+"-stimuli.csv"))
    assert u.id.is_unique and e.id.is_unique and s.id.is_unique
    a=u.merge(e[["id","probe_id","location","valid_data"]],left_on="peak_channel_id",right_on="id",suffixes=("","_electrode"),validate="many_to_one")
    rules={"quality":a.quality.eq("good"),"amplitude_cutoff":np.isfinite(a.amplitude_cutoff)&a.amplitude_cutoff.lt(.1),"presence_ratio":np.isfinite(a.presence_ratio)&a.presence_ratio.gt(.95),"isi_violations":np.isfinite(a.isi_violations)&a.isi_violations.lt(.5),"peak_electrode":a.valid_data.astype(str).str.lower().eq("true")}
    a["eligible"]=pd.DataFrame(rules).all(axis=1)
    a["exclusion_reasons"]=[";".join(k for k,v in rules.items() if not v.iloc[i]) for i in range(len(a))]
    a.to_csv(OUT/(group+"-unit-audit.csv"),index=False)
    trials=s[s.condition.eq("a single square pulse")&s.stimulus_name.eq("pulse")&np.isclose(s.duration,.01,atol=.0001,rtol=0)].sort_values("start_time").copy()
    assert np.allclose(s.stop_time-s.start_time,s.duration,atol=1e-8,rtol=0)
    for row in trials.itertuples():
        assert not ((s.start_time<row.start_time+.1)&(s.stop_time>row.start_time-.1)&s.id.ne(row.id)).any()
    with np.load(io.BytesIO(z.read(group+"-spikes.npz")),allow_pickle=False) as n:
        arr={k:n[k].copy() for k in n.files}
    assert np.array_equal(arr["unit_ids"],u.id.to_numpy())
    expected=np.array([[t-.1,t+.1] for t in trials.start_time])
    assert np.array_equal(arr["windows"],expected)
    ends=arr["spike_times_index"]; times=arr["spike_times"]
    assert len(ends)==len(u) and ends[-1]==len(times) and np.all(np.diff(np.r_[0,ends])>=0)
    assert np.isfinite(times).all()
    trains={}; previous=0
    for uid,end in zip(arr["unit_ids"],ends):
        sp=times[previous:end]; previous=end
        assert np.all(np.diff(sp)>=0)
        wi=np.searchsorted(expected[:,0],sp,side="right")-1
        assert np.all(wi>=0) and np.all(sp<expected[wi,1])
        trains[int(uid)]=sp
    coverage=[]
    for probe in a.loc[a.eligible,"probe_id"].unique():
        cov=json.loads(z.read(f"{group}-probe-{int(probe)}-coverage.json"))
        ts=[v for k,v in cov["datasets"].items() if k.endswith("/timestamps")][0]
        assert expected[:,0].min()>=ts["first_timestamp"] and expected[:,1].max()<=ts["last_timestamp"]
        coverage.append({"probe_id":int(probe),"first_timestamp":ts["first_timestamp"],"last_timestamp":ts["last_timestamp"]})
    trials["coverage_status"]="assumed_from_acquisition_metadata"
    trials.to_csv(OUT/(group+"-trials.csv"),index=False)
    DATA[group]={"meta":m,"units":a,"trials":trials,"trains":trains,"coverage":coverage}
    checks["sessions"][group]={"units":len(u),"eligible_units":int(a.eligible.sum()),"spikes_in_extract":len(times),"trials":len(trials),"coverage":coverage,"all_structural_checks_pass":True}
(OUT/"input-validation.json").write_text(json.dumps(checks,indent=2))
(OUT/"environment.json").write_text(json.dumps({"python":platform.python_version(),"numpy":np.__version__,"pandas":pd.__version__,"scipy":scipy.__version__,"seed":SEED},indent=2))

def signflip(d, seed):
    d=np.asarray(d,dtype=int); d=d[d!=0]; target=int(d.sum())
    if len(d)==0:return 1.0
    if len(d)<=20:
        dist=Counter({0:1})
        for v in np.abs(d):
            nxt=Counter()
            for k,n in dist.items():nxt[k+int(v)]+=n;nxt[k-int(v)]+=n
            dist=nxt
        return sum(n for k,n in dist.items() if k>=target)/(2**len(d))
    rng=np.random.default_rng(seed); exceed=0; B=99999
    for start in range(0,B,5000):
        signs=rng.integers(0,2,size=(min(5000,B-start),len(d)),dtype=np.int8)*2-1
        exceed+=int(np.sum(signs@d>=target))
    return (exceed+1)/(B+1)

def counts(sp,t,lo,hi):
    return np.searchsorted(sp,t+hi,side="left")-np.searchsorted(sp,t+lo,side="left")

records=[]; trial_records=[]; bin_cache={}
for group,D in DATA.items():
    primary=D["trials"].level.max()
    for unit in D["units"].loc[D["units"].eligible].itertuples():
        sp=D["trains"][int(unit.id)]
        for level,tr in D["trials"].groupby("level",sort=True):
            t=tr.start_time.to_numpy(); n=len(t)
            pre=counts(sp,t,-.008,-.002); post=counts(sp,t,.002,.008)
            pre3=counts(sp,t,-.007,-.002);post3=counts(sp,t,.003,.008)
            first=np.full(n,np.nan)
            for j,time in enumerate(t):
                inds=np.searchsorted(sp,[time+.002,time+.008])
                if inds[1]>inds[0]:first[j]=(sp[inds[0]]-time)*1000
                trial_records.append({"group":group,"session_id":D["meta"]["session_id"],"unit_id":int(unit.id),"level":level,"trial_id":int(tr.id.iloc[j]),"start_time_s":time,"baseline_count":int(pre[j]),"response_count":int(post[j]),"baseline3_count":int(pre3[j]),"response3_count":int(post3[j]),"first_spike_ms":first[j],"coverage_status":"assumed_from_acquisition_metadata"})
            eligible=group=="sst" and n>=20
            seed=SEED+int(unit.id)+int(round(level*100))
            rec={"group":group,"session_id":D["meta"]["session_id"],"unit_id":int(unit.id),"probe_id":int(unit.probe_id),"region":unit.location,"level":level,"primary":level==primary,"n_trials":n,
            "baseline_hz":pre.mean()/.006,"response_hz":post.mean()/.006,"difference_hz":(post-pre).mean()/.006,
            "fold_change":post.sum()/pre.sum() if pre.sum() else np.nan,
            "baseline_trial_fraction":np.mean(pre>0),"response_trial_fraction":np.mean(post>0),
            "baseline3_hz":pre3.mean()/.005,"response3_hz":post3.mean()/.005,"difference3_hz":(post3-pre3).mean()/.005,
            "p":signflip(post-pre,seed) if eligible else np.nan,"p3":signflip(post3-pre3,seed) if eligible else np.nan,
            "first_spike_median_ms":float(np.nanmedian(first)) if np.isfinite(first).any() else np.nan,
            "first_spike_q25_ms":float(np.nanpercentile(first,25)) if np.isfinite(first).any() else np.nan,
            "first_spike_q75_ms":float(np.nanpercentile(first,75)) if np.isfinite(first).any() else np.nan,
            "inference_status":"conditional_on_coverage" if eligible else "descriptive_only"}
            records.append(rec)
R=pd.DataFrame(records); T=pd.DataFrame(trial_records)
for primary in (True,False):
    ix=(R.group=="sst")&(R.primary==primary)&R.p.notna()
    R.loc[ix,"q"]=false_discovery_control(R.loc[ix,"p"].to_numpy())
    R.loc[ix,"q3"]=false_discovery_control(R.loc[ix,"p3"].to_numpy())
R["responder"]=pd.Series(pd.NA,index=R.index,dtype="boolean")
R["responder3"]=pd.Series(pd.NA,index=R.index,dtype="boolean")
ix=R.group=="sst"
R.loc[ix,"responder"]=(R.loc[ix,"q"]<=.05)&(R.loc[ix,"response_hz"]>=2*R.loc[ix,"baseline_hz"])&(R.loc[ix,"difference_hz"]>=10)
R.loc[ix,"responder3"]=(R.loc[ix,"q3"]<=.05)&(R.loc[ix,"response3_hz"]>=2*R.loc[ix,"baseline3_hz"])&(R.loc[ix,"difference3_hz"]>=10)
R["onset_ms"]=np.nan;R["first_spike_ci_low_ms"]=np.nan;R["first_spike_ci_high_ms"]=np.nan
for idx,row in R.loc[R.primary & R.responder.fillna(False)].iterrows():
    D=DATA[row.group];sp=D["trains"][int(row.unit_id)];t=D["trials"].loc[D["trials"].level==row.level,"start_time"].to_numpy()
    diffs=np.stack([counts(sp,t,(2+b)/1000,(3+b)/1000)-counts(sp,t,(-8+b)/1000,(-7+b)/1000) for b in range(6)],axis=1)
    rng=np.random.default_rng(SEED+int(row.unit_id)); null=[]
    for start in range(0,99999,5000):
        signs=rng.integers(0,2,size=(min(5000,99999-start),len(t)),dtype=np.int8)*2-1
        null.append(np.maximum(0,(signs@diffs).max(axis=1)))
    threshold=np.quantile(np.concatenate(null),.99)
    above=diffs.sum(axis=0)>threshold
    consecutive=np.flatnonzero(above[:-1]&above[1:])
    if len(consecutive): R.loc[idx,"onset_ms"]=2+consecutive[0]
    R.loc[idx,"onset_null_max_count_threshold"]=threshold
    first=T.loc[(T.group==row.group)&(T.unit_id==row.unit_id)&(T.level==row.level),"first_spike_ms"].to_numpy()
    samples=first[rng.integers(0,len(first),size=(2000,len(first)))]
    valid=np.any(np.isfinite(samples),axis=1)
    boot=np.nanmedian(samples[valid],axis=1)
    R.loc[idx,["first_spike_ci_low_ms","first_spike_ci_high_ms"]]=np.quantile(boot,[.025,.975])
R.to_csv(OUT/"unit-results.csv",index=False);T.to_csv(OUT/"trial-counts-and-latencies.csv",index=False)
summary=[]
for (g,l),rr in R.groupby(["group","level"]):
    responders=rr.responder.fillna(False)
    summary.append({"group":g,"level":l,"units":len(rr),"trials_per_unit":int(rr.n_trials.iloc[0]),"baseline_mean_hz":rr.baseline_hz.mean(),"response_mean_hz":rr.response_hz.mean(),"median_difference_hz":rr.difference_hz.median(),"responders":int(responders.sum()) if g=="sst" else None,"strict_responders":int(rr.responder3.fillna(False).sum()) if g=="sst" else None,"resolved_onsets":int(rr.onset_ms.notna().sum()),"onset_values_ms":rr.onset_ms.dropna().tolist()})
for rec in summary:
    rec["onset_evaluated"]=bool(rec["group"]=="sst" and rec["level"]==DATA["sst"]["trials"].level.max())
    if not rec["onset_evaluated"]:rec["resolved_onsets"]=None;rec["onset_values_ms"]=None
(OUT/"summary.json").write_text(json.dumps(summary,indent=2))
print(json.dumps(summary,indent=2))
