import json, zipfile, hashlib
from pathlib import Path
import numpy as np
import pandas as pd

def verify(zip_path, outdir):
    out = Path(outdir)
    out.mkdir(exist_ok=True)
    summary = {"source_zip_sha256": hashlib.sha256(Path(zip_path).read_bytes()).hexdigest(), "software": {"numpy": np.__version__, "pandas": pd.__version__}, "sessions": {}}
    with zipfile.ZipFile(zip_path) as z:
        summary["member_sha256"] = {n: hashlib.sha256(z.read(n)).hexdigest() for n in z.namelist() if not n.endswith("/")}
        for g in ("sst", "wt"):
            m = json.loads(z.read(g+"-metadata.json"))
            s = pd.read_csv(z.open(g+"-stimuli.csv"))
            u = pd.read_csv(z.open(g+"-units.csv"))
            e = pd.read_csv(z.open(g+"-electrodes.csv"))
            schema = json.loads(z.read(g+"-remote-schema.json"))
            assert s.id.is_unique and u.id.is_unique and e.id.is_unique
            assert len(u) == schema["units/id"]["shape"][0]
            assert len(s) == schema["processing/optotagging/optogenetic_stimulation/id"]["shape"][0]
            assert s.session_id.eq(int(m["session_id"])).all()
            assert u.session_id.eq(int(m["session_id"])).all()
            assert np.isfinite(s[["start_time","stop_time","duration","level"]]).all().all()
            assert s.duration.gt(0).all()
            assert np.allclose(s.stop_time-s.start_time,s.duration,atol=1e-8,rtol=0)
            a = u.merge(e[["id","probe_id","location","valid_data"]],left_on="peak_channel_id",right_on="id",suffixes=("","_electrode"),how="left",validate="many_to_one")
            a["pass_quality"] = a.quality.eq("good")
            a["pass_amplitude_cutoff"] = np.isfinite(a.amplitude_cutoff)&a.amplitude_cutoff.lt(.1)
            a["pass_presence_ratio"] = np.isfinite(a.presence_ratio)&a.presence_ratio.gt(.95)
            a["pass_isi_violations"] = np.isfinite(a.isi_violations)&a.isi_violations.lt(.5)
            a["pass_original_qc"] = a[["pass_quality","pass_amplitude_cutoff","pass_presence_ratio","pass_isi_violations"]].all(axis=1)
            a["pass_peak_electrode"] = a.valid_data.astype(str).str.lower().eq("true") & a.probe_id.notna()
            a["candidate"] = a.pass_original_qc & a.pass_peak_electrode
            a["coverage_status"] = "unverified"
            flags = ["pass_quality","pass_amplitude_cutoff","pass_presence_ratio","pass_isi_violations","pass_peak_electrode"]
            a["exclusion_reasons"] = a[flags].apply(lambda row: ";".join(k for k in flags if not row[k]),axis=1)
            a.to_csv(out/(g+"-unit-eligibility.csv"),index=False)
            selected = s.loc[s.condition.eq("a single square pulse") & s.stimulus_name.eq("pulse") & np.isclose(s.duration,.01,atol=.0001,rtol=0)].copy()
            selected["window_start_s"] = selected.start_time-.1
            selected["window_stop_s"] = selected.start_time+.1
            selected["other_event_overlap"] = [bool(((s.start_time<t+.1)&(s.stop_time>t-.1)&s.id.ne(i)).any()) for i,t in zip(selected.id, selected.start_time)]
            selected["primary_level"] = selected.level.eq(selected.level.max())
            selected["coverage_status"] = "unverified"
            selected.to_csv(out/(g+"-extraction-windows.csv"),index=False)
            counts = s.assign(duration_ms=(s.duration*1000).round(3)).groupby(["duration_ms","condition","stimulus_name","level"],dropna=False).size().reset_index(name="trials")
            counts.to_csv(out/(g+"-stimulus-counts.csv"),index=False)
            summary["sessions"][g] = {
              "session_id":m["session_id"],"recording_start":m["session_start_time"],
              "timestamps_reference_time":m["timestamps_reference_time"],"subject":m["subject"],
              "dandi_version":m["dandi_version"],"nwb_version":m["nwb_version"],
              "asset_id":m["source"]["identifier"],"asset_path":m["source"]["path"],
              "source_digests_as_reported":m["source"]["digest"],
              "n_units":len(u),"n_stimuli":len(s),"n_electrodes":len(e),
              "original_qc_pass":int(a.pass_original_qc.sum()),"candidate_units":int(a.candidate.sum()),
              "qc_pass_but_invalid_peak":int((a.pass_original_qc & ~a.pass_peak_electrode).sum()),
              "unmapped_peak_channels":int(a.probe_id.isna().sum()),
              "missing_qc":u[["quality","amplitude_cutoff","presence_ratio","isi_violations"]].isna().sum().to_dict(),
              "ten_ms_counts":selected.groupby("level").size().to_dict(),
              "primary_level":float(selected.level.max()),
              "other_event_overlap_rejections":int(selected.other_event_overlap.sum()),
              "extraction_envelope_s":[float(selected.window_start_s.min()),float(selected.window_stop_s.max())],
              "selected_window_total_s":float(len(selected)*.2),
              "minimum_all_event_onset_spacing_s":float(np.diff(np.sort(s.start_time)).min()),
              "maximum_duration_error_s":float(abs(s.stop_time-s.start_time-s.duration).max()),
              "interval_tables":m["interval_tables"],
              "obs_intervals_in_unit_columns":"obs_intervals" in m["units_attrs"]["colnames"],
              "coverage_status":"unverified",
              "layout":m["dataset_layout"]
            }
    (out/"verification.json").write_text(json.dumps(summary,indent=2))
    return summary

# Usage: import this module and call verify(path_to_zip, output_directory).
