"""
Aggregate per-unit GLM prediction PSTH NPZ files into the Figure S5/S13/S14 CSV.

The per-unit NPZ files are generated by ``vbn_code/hpc_code/run_GLM_prediction_psths.py``.
This helper moves the notebook aggregation step into a CLI.
"""

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

import numpy as np
import pandas as pd


def load_unit_glm_npz(path: Path) -> dict:
    unit_id = int(path.name.split("_")[0])
    unit_data = np.load(path)
    trial_counts = unit_data["trial_counts"]
    return {
        "unit_id": unit_id,
        "num_hits": trial_counts[0],
        "num_misses": trial_counts[1],
        "num_nonchange_licks": trial_counts[2],
        "num_nonchange_no_licks": trial_counts[3],
        "predicted_psth": unit_data["prediction"],
        "psth": unit_data["psth"],
        "psth_sem": unit_data["psth_sem"],
        "predicted_sems": unit_data["prediction_sem"],
    }


def aggregate_glm_prediction_psths(
    input_dir: Path,
    output_file: Path,
) -> pd.DataFrame:
    rows = [
        load_unit_glm_npz(path)
        for path in sorted(input_dir.glob("*.npz"))
    ]
    glm_results = pd.DataFrame(rows)
    output_file.parent.mkdir(parents=True, exist_ok=True)
    glm_results.to_csv(output_file, index=False)
    return glm_results


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input-dir", type=Path, required=True)
    parser.add_argument("--output-file", type=Path, required=True)
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    aggregate_glm_prediction_psths(args.input_dir, args.output_file)


if __name__ == "__main__":
    main()
