"""Create a separate configurable ASR cache for the TCH recordings."""
from __future__ import annotations

import argparse
import csv
import json
from pathlib import Path

import mne
import numpy as np
from asrpy import ASR
from scipy.signal import butter, resample_poly, sosfiltfilt

from prepare_tch_data import CHANNELS, EXCLUDED_TRIALS, LABELS, SKIPPED_SESSIONS, _event_rows, session_id


def prepare(source: Path, output: Path, cutoff: float = 10.0):
    source_rate, target_rate = 1000, 200
    pre_samples, post_samples = 200, 400
    output.mkdir(parents=True, exist_ok=True)
    sessions = output / "sessions"
    sessions.mkdir(exist_ok=True)
    sos = butter(4, (1.0, 40.0), btype="bandpass", fs=source_rate, output="sos")
    info = mne.create_info(list(CHANNELS), source_rate, "eeg")
    manifest, quality, trials = [], [], []
    for path in sorted(source.glob("*.npz")):
        session = session_id(path)
        if session in SKIPPED_SESSIONS:
            manifest.append({"session_id":session,"included":False,"reason":"record marked unusable"})
            continue
        raw_values, events = _event_rows(path, source_rate, target_rate, pre_samples, post_samples)
        filtered = sosfiltfilt(sos, raw_values.astype(np.float64), axis=1)
        # TCH recorder values are stored in nV; MNE/ASR operates in volts.
        raw = mne.io.RawArray(filtered * 1e-9, info.copy(), verbose=False)
        asr = ASR(sfreq=source_rate, cutoff=cutoff, win_len=.5, win_overlap=.66, max_bad_chans=.1)
        _, clean_mask = asr.fit(raw, return_clean_window=True)
        cleaned = asr.transform(raw).get_data() * 1e9
        continuous = resample_poly(cleaned, up=target_rate, down=source_rate, axis=1).astype(np.float32)
        kept, labels, numbers = [], [], []
        excluded = EXCLUDED_TRIALS.get(session, set())
        for number, code, index, timestamp in events:
            start, stop = index-pre_samples, index+post_samples
            reason = "excluded_by_record" if number in excluded else ("boundary" if start < 0 or stop > continuous.shape[1] else "")
            if not reason:
                kept.append(continuous[:,start:stop]); labels.append(LABELS[code]); numbers.append(number)
            trials.append({"session_id":session,"target_trial_number":number,"event_code":code,"label":LABELS[code],"kept":not bool(reason),"exclude_reason":reason,"marker_timestamp":timestamp})
        x=np.stack(kept).astype(np.float32); y=np.asarray(labels,dtype=np.int64)
        np.save(sessions/f"{session}_x.npy",x,allow_pickle=False); np.save(sessions/f"{session}_y.npy",y,allow_pickle=False); np.save(sessions/f"{session}_trial_numbers.npy",np.asarray(numbers,dtype=np.int64),allow_pickle=False)
        correction = cleaned-filtered
        quality.append({"session_id":session,"raw_centered_sd_nV":float(np.std(raw_values-raw_values.mean(axis=1,keepdims=True))),"filtered_1_40_sd_nV":float(np.std(filtered)),"asr_clean_sd_nV":float(np.std(cleaned)),"asr_correction_rms_nV":float(np.sqrt(np.mean(correction**2)),),"asr_cutoff":cutoff,"asr_calibration_fraction":float(clean_mask.mean()),"asr_calibration_samples":int(clean_mask.sum())})
        manifest.append({"session_id":session,"included":True,"source_file":path.name,"trials_raw":len(events),"trials_kept":len(y),"left":int((y==0).sum()),"right":int((y==1).sum()),"shape":str(tuple(x.shape))})
        print(f"[ASR cutoff={cutoff:g}] {session}: kept={len(y)} calibration={clean_mask.mean():.3f}",flush=True)
    for name, rows in (("session_manifest.csv",manifest),("trial_manifest.csv",trials),("quality_summary.csv",quality)):
        fields=sorted({key for row in rows for key in row})
        with (output/name).open("w",encoding="utf-8-sig",newline="") as handle:
            writer=csv.DictWriter(handle,fieldnames=fields); writer.writeheader(); writer.writerows(rows)
    (output/"config.json").write_text(json.dumps({"pipeline":f"1-40 Hz fourth-order zero-phase SOS IIR -> ASR cutoff {cutoff:g} on continuous filtered signal -> 1000 to 200 Hz -> event[-1,2] s","asr":{"cutoff":cutoff,"win_len":.5,"win_overlap":.66,"max_bad_chans":.1},"units":"recorder values interpreted as nV for MNE/ASR conversion"},ensure_ascii=False,indent=2),encoding="utf-8")


if __name__=="__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--cutoff", type=float, default=10.0)
    parser.add_argument("--output", type=Path)
    args = parser.parse_args()
    root=Path(__file__).resolve().parent.parent
    output = args.output or Path(__file__).resolve().parent / f"data_cache_asr{args.cutoff:g}"
    prepare(root/"中榮data", output, args.cutoff)
