"""
Clean EEG Pipeline for TCH Dataset
Handles despiking, continuous bandpass filtering, session z-score standardization,
epoch extraction, and pre-stimulus baseline subtraction without data leakage.
"""

from __future__ import annotations
import csv
import json
from pathlib import Path
import numpy as np
from scipy.signal import butter, sosfiltfilt, resample_poly

CHANNELS = ("Fp1", "Fp2", "Fz", "C3", "C4", "Pz", "O1", "O2")
LABELS = {-1: 0, 1: 1}  # left=0, right=1

EXCLUDED_TRIALS = {
    "s1_5": {21},
    "s2_1": {5, 12, 13, 19, 25},
    "s2_2": {17, 21, 37, 40},
    "s2_3": {5, 6, 8, 10, 11, 12, 20, 27, 32},
    "s2_4": {12, 19, 21, 27, 37, 39},
}
SKIPPED_SESSIONS = {"s3_1"}


def session_id(path: Path) -> str:
    return "_".join(path.stem.split("_")[:2])


def despike_continuous(eeg: np.ndarray, max_diff: float = 200.0) -> np.ndarray:
    """Detect and remove sudden gradient spikes in continuous EEG."""
    cleaned = eeg.copy()
    n_ch, n_samples = cleaned.shape
    for ch in range(n_ch):
        diff = np.abs(np.diff(cleaned[ch]))
        spike_indices = np.where(diff > max_diff)[0]
        for idx in spike_indices:
            if 0 < idx < n_samples - 1:
                cleaned[ch, idx + 1] = (cleaned[ch, idx] + cleaned[ch, min(idx + 2, n_samples - 1)]) / 2.0
    return cleaned


def process_continuous_session(
    eeg: np.ndarray,
    source_rate: int = 1000,
    target_rate: int = 200,
    band: tuple[float, float] = (4.0, 32.0),
    max_diff: float = 200.0,
) -> np.ndarray:
    """
    1. Despike
    2. Bandpass filter continuous signal
    3. Session-level Z-score standardization per channel
    4. Resample to target rate (200 Hz)
    """
    # 1. Despike
    eeg_clean = despike_continuous(eeg.astype(np.float64), max_diff=max_diff)

    # 2. Continuous IIR bandpass filter (zero-phase)
    sos = butter(4, band, btype="bandpass", fs=source_rate, output="sos")
    filtered = sosfiltfilt(sos, eeg_clean, axis=1)

    # 3. Session-level z-score standardization per channel
    for ch in range(filtered.shape[0]):
        std = np.std(filtered[ch])
        mean = np.mean(filtered[ch])
        if std > 1e-6:
            filtered[ch] = (filtered[ch] - mean) / std

    # 4. Resample to target rate (200 Hz)
    resampled = resample_poly(filtered, up=target_rate, down=source_rate, axis=1).astype(np.float32)
    return resampled


def extract_event_epochs(
    path: Path,
    continuous_eeg: np.ndarray,
    source_rate: int = 1000,
    target_rate: int = 200,
    pre_samples: int = 200,
    post_samples: int = 400,
):
    """
    Extract epochs aligned to marker timestamps and apply pre-stimulus baseline subtraction.
    """
    record = np.load(path, allow_pickle=True)
    timestamps = np.asarray(record["eeg_ts"], dtype=np.float64)
    codes = record["marker_code"].tolist()
    marker_timestamps = np.asarray(record["marker_ts"], dtype=np.float64)

    events = []
    target_number = 0
    for code, timestamp in zip(codes, marker_timestamps):
        if int(code) not in LABELS:
            continue
        target_number += 1
        source_index = int(np.argmin(np.abs(timestamps - timestamp)))
        target_index = round(source_index * target_rate / source_rate)
        events.append((target_number, int(code), target_index, float(timestamp)))

    session = session_id(path)
    excluded = EXCLUDED_TRIALS.get(session, set())
    kept_epochs, kept_labels, kept_numbers = [], [], []

    for number, code, index, timestamp in events:
        start, stop = index - pre_samples, index + post_samples
        if number in excluded or start < 0 or stop > continuous_eeg.shape[1]:
            continue

        epoch = continuous_eeg[:, start:stop].copy()
        # Pre-stimulus baseline subtraction: [-1, 0]s -> first pre_samples (samples 0:200)
        baseline_mean = np.mean(epoch[:, :pre_samples], axis=1, keepdims=True)
        epoch_corrected = epoch - baseline_mean

        kept_epochs.append(epoch_corrected)
        kept_labels.append(LABELS[code])
        kept_numbers.append(number)

    x = np.stack(kept_epochs).astype(np.float32)
    y = np.asarray(kept_labels, dtype=np.int64)
    numbers = np.asarray(kept_numbers, dtype=np.int64)
    return x, y, numbers


def prepare_cleaned_cache(source_dir: Path, output_dir: Path) -> None:
    source_rate, target_rate = 1000, 200
    pre_samples, post_samples = 200, 400
    output_dir.mkdir(parents=True, exist_ok=True)
    cleaned_dir = output_dir / "cleaned_sessions"
    cleaned_dir.mkdir(exist_ok=True)

    manifest = []
    for path in sorted(source_dir.glob("*.npz")):
        session = session_id(path)
        if session in SKIPPED_SESSIONS:
            continue

        record = np.load(path, allow_pickle=True)
        eeg_raw = np.asarray(record["eeg"], dtype=np.float32)

        continuous_cleaned = process_continuous_session(
            eeg_raw, source_rate=source_rate, target_rate=target_rate, band=(4.0, 32.0), max_diff=200.0
        )

        x, y, numbers = extract_event_epochs(
            path, continuous_cleaned, source_rate=source_rate, target_rate=target_rate, pre_samples=pre_samples, post_samples=post_samples
        )

        np.save(cleaned_dir / f"{session}_x.npy", x, allow_pickle=False)
        np.save(cleaned_dir / f"{session}_y.npy", y, allow_pickle=False)
        np.save(cleaned_dir / f"{session}_trial_numbers.npy", numbers, allow_pickle=False)

        manifest.append({
            "session_id": session,
            "trials": len(y),
            "left": int((y == 0).sum()),
            "right": int((y == 1).sum()),
            "shape": str(tuple(x.shape)),
        })
        print(f"Processed {session}: x.shape={x.shape}, left={(y==0).sum()}, right={(y==1).sum()}")

    (output_dir / "cleaned_config.json").write_text(
        json.dumps({
            "channels": CHANNELS,
            "source_rate": source_rate,
            "target_rate": target_rate,
            "filter": "continuous 4th-order Butterworth zero-phase 4-32 Hz + despike + session zscore",
            "epoch": "event[-1, 2] s with [-1, 0]s baseline subtraction",
            "task_window": "event[0, 2] s (samples 200:600)",
        }, ensure_ascii=False, indent=2),
        encoding="utf-8",
    )


if __name__ == "__main__":
    package = Path(__file__).resolve().parents[1]
    workspace = package.parent
    prepare_cleaned_cache(workspace / "中榮data", package / "output_cache")
