"""
Build per-unit GLM prediction PSTH NPZ files.

The GLM fit files are external upstream inputs, expected as one compressed
``<session_id>.pbz2`` file per session. This script preserves the original
condition definitions and output schema while replacing hard-coded paths with
CLI arguments.
"""

import argparse
import bz2
from pathlib import Path
from typing import Optional, Sequence

import _pickle as cPickle
import numpy as np
import pandas as pd
from scipy.interpolate import interp1d


def resample_df_to_times(df, time_column, val_column, new_times):
    timestamps = df[time_column].values
    vals = df[val_column].values
    interpolator = interp1d(timestamps, vals, kind="linear", bounds_error=False)
    new_values = interpolator(new_times)
    return new_values, new_times


def get_condition_filters(session_stims: pd.DataFrame) -> Sequence[pd.Series]:
    hits = (
        session_stims["hit"]
        & session_stims["is_change"]
    )
    misses = (
        session_stims["miss"]
        & session_stims["is_change"]
    )
    nonchange_licks = (
        (~session_stims["is_change"])
        & (~session_stims["omitted"])
        & (~session_stims["previous_omitted"])
        & (session_stims["flashes_since_change"] > 5)
        & (session_stims["flashes_since_last_lick"] > 1)
        & session_stims["lickbout_for_flash_during_response_window"]
    )
    nonchange_nolicks = (
        (~session_stims["is_change"])
        & (~session_stims["omitted"])
        & (~session_stims["previous_omitted"])
        & (session_stims["flashes_since_change"] > 5)
        & (session_stims["flashes_since_last_lick"] > 1)
        & (~session_stims["lickbout_for_flash_during_response_window"])
    )
    return hits, misses, nonchange_licks, nonchange_nolicks


def load_glm_fit(glm_fit_file: Path) -> dict:
    with bz2.BZ2File(glm_fit_file, "rb") as file:
        return cPickle.load(file)


def run_glm_prediction_psths(
    session_id: int,
    glm_fit_dir: Path,
    stim_table_file: Path,
    output_dir: Path,
    min_trials_per_condition: int = 5,
) -> Sequence[Path]:
    glm_fit_file = glm_fit_dir / f"{session_id}.pbz2"
    fit = load_glm_fit(glm_fit_file)

    stim_table = pd.read_csv(stim_table_file)
    session_stims = stim_table[stim_table["session_id"] == int(session_id)]
    condition_filters = get_condition_filters(session_stims)

    if any(condition_filter.sum() < min_trials_per_condition for condition_filter in condition_filters):
        print(f"Not enough trials in one of the conditions, skipping session {session_id}...")
        return []

    output_dir.mkdir(parents=True, exist_ok=True)
    unit_ids = fit["spike_count_arr"].unit_id.values
    written_files = []
    for unit_index, unit_id in enumerate(unit_ids):
        fit_df = pd.DataFrame(
            {
                "prediction": fit["full_model_prediction"][:, unit_index],
                "activity": fit["spike_count_arr"][:, unit_index].values,
                "time": fit["bin_centers"],
            }
        )
        unit_predictions = []
        unit_prediction_sems = []
        unit_psths = []
        unit_psth_sems = []
        for condition_filter in condition_filters:
            condition_start_times = session_stims.loc[condition_filter]["start_time"].values
            condition_psth = []
            condition_predicted_psth = []
            for start_time in condition_start_times:
                psth, _ = resample_df_to_times(
                    fit_df,
                    "time",
                    "activity",
                    np.arange(start_time - 0.25, start_time + 1, 0.025),
                )
                predicted_psth, _ = resample_df_to_times(
                    fit_df,
                    "time",
                    "prediction",
                    np.arange(start_time - 0.25, start_time + 1, 0.025),
                )

                condition_psth.append(psth[:40])
                condition_predicted_psth.append(predicted_psth[:40])

            unit_predictions.append(np.mean(condition_predicted_psth, axis=0))
            unit_prediction_sems.append(
                np.std(condition_predicted_psth, axis=0) / np.sqrt(len(condition_predicted_psth))
            )
            unit_psths.append(np.mean(condition_psth, axis=0))
            unit_psth_sems.append(
                np.std(condition_psth, axis=0) / np.sqrt(len(condition_psth))
            )

        output_file = output_dir / f"{unit_id}_{session_id}.npz"
        np.savez(
            output_file,
            prediction=np.concatenate(unit_predictions),
            psth=np.concatenate(unit_psths),
            prediction_sem=np.concatenate(unit_prediction_sems),
            psth_sem=np.concatenate(unit_psth_sems),
            trial_counts=[condition_filter.sum() for condition_filter in condition_filters],
        )
        written_files.append(output_file)

    return written_files


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--session-id", type=int, required=True)
    parser.add_argument("--glm-fit-dir", type=Path, required=True)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--min-trials-per-condition", type=int, default=5)
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    run_glm_prediction_psths(
        session_id=args.session_id,
        glm_fit_dir=args.glm_fit_dir,
        stim_table_file=args.stim_table_file,
        output_dir=args.output_dir,
        min_trials_per_condition=args.min_trials_per_condition,
    )


if __name__ == "__main__":
    main()
