"""
Build the Figure 6 responsiveness-over-time summary pickle.

This moves the value-producing logic from Figure6_co.ipynb into a public
helper while avoiding the hard-coded stimulus-table path inside
``vbn_utils.calculate_responsiveness_over_time``.
"""

import argparse
import concurrent.futures
import pickle
from pathlib import Path
from typing import Optional, Sequence

import h5py
import numpy as np
import pandas as pd
import tqdm

import vbn_utils


REGIONS = ("VISall", "SCMRN", "LGd", "LP")
DEFAULT_WINDOW_SIZE = 40


def read_table(path: Path) -> pd.DataFrame:
    table = pd.read_csv(path)
    unnamed_columns = [col for col in table.columns if col.startswith("Unnamed:")]
    if unnamed_columns:
        table = table.drop(columns=unnamed_columns)
    return table


def write_pickle(data: object, output_file: Path) -> None:
    output_file.parent.mkdir(parents=True, exist_ok=True)
    with output_file.open("wb") as file:
        pickle.dump(data, file)


def get_image_sets(stim_table: pd.DataFrame) -> dict:
    g_images = (
        ["omitted"]
        + list(
            np.sort(
                stim_table[
                    (stim_table["stimulus_name"].str.contains("_G_"))
                    & (~stim_table["omitted"])
                    & (~stim_table["image_name"].isin(["im083_r", "im111_r"]))
                ]["image_name"].unique()
            )
        )
        + ["im083_r", "im111_r"]
    )
    h_images = (
        ["omitted"]
        + list(
            np.sort(
                stim_table[
                    (stim_table["stimulus_name"].str.contains("_H_"))
                    & (~stim_table["omitted"])
                    & (~stim_table["image_name"].isin(["im083_r", "im111_r"]))
                ]["image_name"].unique()
            )
        )
        + ["im083_r", "im111_r"]
    )
    return {"G": g_images, "H": h_images}


def calculate_session_responsiveness_over_time(
    session_id: int,
    active_tensor_file: Path,
    stim_table: pd.DataFrame,
    unit_ids: Sequence[int],
    image_sets: dict,
    window_size: int = DEFAULT_WINDOW_SIZE,
) -> dict:
    with h5py.File(active_tensor_file, "r") as active_tensor:
        session_tensor = active_tensor[str(session_id)]
        session_unit_ids = session_tensor["unitIds"][()]
        unit_indices = [
            index
            for index, unit_id in enumerate(session_unit_ids)
            if unit_id in unit_ids
        ]
        session_stim_table = stim_table[stim_table["session_id"] == int(session_id)].reset_index()
        image_set = (
            image_sets["G"]
            if "_G_" in session_stim_table["stimulus_name"].iloc[0]
            else image_sets["H"]
        )

        data_dict = {"unit_ids": session_unit_ids[unit_indices]}
        spikes = session_tensor["spikes"]
        for image in image_set:
            filter_stims = vbn_utils.get_nonchange_flashes(session_stim_table, image_id=image)
            stim_sp = np.full((len(unit_indices), len(filter_stims), 750), np.nan)
            pre_stim_sp = np.full((len(unit_indices), len(filter_stims), 750), np.nan)
            for unit_count, unit_index in enumerate(unit_indices):
                stim_sp[unit_count] = spikes[unit_index, filter_stims, :]
                pre_stim_sp[unit_count] = spikes[unit_index, filter_stims - 1, :]

            data_dict[image] = vbn_utils.findResponsiveUnits_overtime(
                pre_stim_sp,
                stim_sp,
                window_duration=window_size,
            )

    return data_dict


def build_responsiveness_over_time_summary(
    unit_table_file: Path,
    stim_table_file: Path,
    active_tensor_file: Path,
    output_file: Path,
    session_ids: Optional[Sequence[int]] = None,
    window_size: int = DEFAULT_WINDOW_SIZE,
    max_workers: int = 20,
) -> dict:
    units = read_table(unit_table_file)
    stim_table = read_table(stim_table_file)
    unit_ids = vbn_utils.get_unit_ids(units, list(REGIONS))

    if session_ids is None:
        session_ids = (
            units.set_index("unit_id")
            .loc[unit_ids]["ecephys_session_id"]
            .dropna()
            .unique()
        )
    session_ids = [int(session_id) for session_id in session_ids]
    image_sets = get_image_sets(stim_table)

    rot_dict = {}
    if max_workers == 1:
        iterator = tqdm.tqdm(session_ids, total=len(session_ids), leave=True)
        for session_id in iterator:
            rot_dict[session_id] = calculate_session_responsiveness_over_time(
                session_id,
                active_tensor_file,
                stim_table,
                unit_ids,
                image_sets,
                window_size=window_size,
            )
    else:
        with concurrent.futures.ProcessPoolExecutor(max_workers=max_workers) as pool:
            futures = {
                pool.submit(
                    calculate_session_responsiveness_over_time,
                    session_id,
                    active_tensor_file,
                    stim_table,
                    unit_ids,
                    image_sets,
                    window_size,
                ): session_id
                for session_id in session_ids
            }
            for future in tqdm.tqdm(
                concurrent.futures.as_completed(futures),
                total=len(futures),
                leave=True,
            ):
                session_id = futures[future]
                try:
                    rot_dict[session_id] = future.result()
                except Exception as exc:
                    print(f"{session_id} generated an exception: {exc}")

    write_pickle(rot_dict, output_file)
    return rot_dict


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--active-tensor-file", type=Path, required=True)
    parser.add_argument("--output-file", type=Path, required=True)
    parser.add_argument("--session-ids", nargs="+", type=int)
    parser.add_argument("--window-size", type=int, default=DEFAULT_WINDOW_SIZE)
    parser.add_argument("--max-workers", type=int, default=20)
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    build_responsiveness_over_time_summary(
        unit_table_file=args.unit_table_file,
        stim_table_file=args.stim_table_file,
        active_tensor_file=args.active_tensor_file,
        output_file=args.output_file,
        session_ids=args.session_ids,
        window_size=args.window_size,
        max_workers=args.max_workers,
    )


if __name__ == "__main__":
    main()
