"""
Build Figure 2 event-aligned unit summary pickles.

The original Figure 2 notebook wrote these assets in-place. This helper keeps
the same calculations in a CLI so public notebooks can read precomputed data.
"""

import argparse
import concurrent.futures
import pickle
from functools import partial
from pathlib import Path
from typing import Callable, Iterable, Optional, Sequence, Tuple

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

import decoding_utils as du
import vbn_utils
from analysis_utils import exponential_convolve, makePSTH_numba


DEFAULT_MANIFEST = "visual-behavior-neuropixels_project_manifest_v0.5.0.json"
LICK_STIM_FILTER = (
    "engaged",
    "~is_change",
    "~omitted",
    "~previous_omitted",
    "flashes_since_change>5",
    "lickbout_for_flash_during_response_window",
)


def read_table(path: Path) -> pd.DataFrame:
    """Read a CSV table while dropping notebook-written unnamed index columns."""

    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 parse_session_ids(session_ids: Optional[Sequence[str]]) -> Optional[np.ndarray]:
    if session_ids is None:
        return None
    return np.array([int(session_id) for session_id in session_ids])


def get_tensor_session_ids(active_tensor_file: Path) -> Sequence[str]:
    with h5py.File(active_tensor_file, "r") as tensor:
        return list(tensor.keys())


def build_lick_aligned_summary(
    active_tensor_file: Path,
    stim_table_file: Path,
    unit_table_file: Path,
    output_file: Path,
    session_ids: Optional[Sequence[str]] = None,
    baseline_length: int = 750,
    response_window_length: int = 1500,
) -> dict:
    """Build ``lick_aligned_unit_data.pkl`` using the original notebook logic."""

    units = read_table(unit_table_file)
    stim_table = read_table(stim_table_file)

    unit_filter = du.apply_unit_quality_filter(units)
    unit_ids = units.loc[unit_filter, "unit_id"].values
    if session_ids is None:
        session_ids = get_tensor_session_ids(active_tensor_file)

    lick_aligned_unit_summary = {unit_id: [] for unit_id in unit_ids}
    unit_data, shuffle_data, unit_id_groups = vbn_utils.unit_averaged_psth_lick_aligned(
        str(active_tensor_file),
        stim_table,
        session_ids,
        unit_ids,
        *LICK_STIM_FILTER,
        baseline_length=baseline_length,
        resp_window_length=response_window_length,
    )

    for unit_session_data, shuffle_session_data, unit_ids_for_session in zip(
        unit_data, shuffle_data, unit_id_groups
    ):
        if len(unit_ids_for_session) > 0:
            for unit_response, shuffle_response, unit_id in zip(
                unit_session_data, shuffle_session_data, unit_ids_for_session
            ):
                unit_mean = exponential_convolve(unit_response, 3, symmetrical=True)
                time_above = np.convolve(unit_mean > shuffle_response[2], np.ones(10)).max()
                time_below = np.convolve(unit_mean < shuffle_response[1], np.ones(10)).max()
                passes = np.max((time_above, time_below)) > 9
                lick_aligned_unit_summary[unit_id] = {
                    "mean": unit_mean,
                    "shuffle_mean": shuffle_response[0],
                    "shuffle_ci_low": shuffle_response[1],
                    "shuffle_ci_high": shuffle_response[2],
                    "pass": passes,
                }

    write_pickle(lick_aligned_unit_summary, output_file)
    return lick_aligned_unit_summary


def load_vbn_cache(
    cache_dir: Path,
    cache_source: str = "s3",
    manifest: str = DEFAULT_MANIFEST,
):
    from allensdk.brain_observatory.behavior.behavior_project_cache.behavior_neuropixels_project_cache import (  # noqa: E501
        VisualBehaviorNeuropixelsProjectCache,
    )

    if cache_source == "s3":
        cache = VisualBehaviorNeuropixelsProjectCache.from_s3_cache(cache_dir=str(cache_dir))
    elif cache_source == "local":
        cache = VisualBehaviorNeuropixelsProjectCache.from_local_cache(cache_dir=str(cache_dir))
    else:
        raise ValueError(f"Unknown cache source: {cache_source}")

    if manifest:
        cache.load_manifest(manifest)
    return cache


def find_running_acceleration_deceleration_times(
    session,
    stimulus_block: int = 5,
) -> Tuple[np.ndarray, np.ndarray]:
    """Find run-start and run-stop times, preserving the original notebook helper."""

    running = session.running_speed.copy()
    running.loc[running["speed"] < 0, "speed"] = 0

    rolling_mean_before = running["speed"].rolling(window=30).mean().shift(1)
    rolling_mean_after = running["speed"].rolling(window=30).mean().shift(-29)

    condition = (rolling_mean_before < 1) & (rolling_mean_after > 5)
    indices = np.where(condition)[0]

    acceleration_times = []
    for running_index in indices:
        if (running_index > 30) and (running_index < len(running) - 31):
            window_indices = running.iloc[running_index - 30 : running_index + 30].index.values
            min_index = running.loc[window_indices].idxmin()["speed"]
            window_diffs = running.loc[min_index : window_indices[-1]].diff()
            try:
                max_diff_index = window_diffs.idxmax()["speed"]
                if running.loc[max_diff_index]["speed"] < 1:
                    last_point_below_threshold = running.loc[max_diff_index]["timestamps"]
                else:
                    last_point_index = np.where(running.loc[min_index:max_diff_index] < 1)[0][-1]
                    last_point_below_threshold = running.loc[min_index + last_point_index][
                        "timestamps"
                    ]
                if len(acceleration_times) > 0:
                    if last_point_below_threshold - acceleration_times[-1] < 0.5:
                        continue
                acceleration_times.append(last_point_below_threshold)
            except Exception as exc:
                session_id = session.metadata["ecephys_session_id"]
                print(f"{session_id} generated an exception: {exc}")
                continue

    stimulus_presentations = session.stimulus_presentations
    passive_stimuli = stimulus_presentations[
        stimulus_presentations["stimulus_block"] == stimulus_block
    ]
    passive_start = passive_stimuli["start_time"].iloc[0]
    passive_end = passive_stimuli["end_time"].iloc[-1]

    acceleration_times = np.array(acceleration_times)
    passive_acceleration_times = acceleration_times[
        (acceleration_times > passive_start) & (acceleration_times < passive_end)
    ]

    condition = (rolling_mean_before > 5) & (rolling_mean_after < 1)
    indices = np.where(condition)[0]

    deceleration_times = []
    for running_index in indices:
        if (running_index > 30) and (running_index < len(running) - 31):
            window_indices = running.iloc[running_index - 30 : running_index + 30].index.values
            max_index = running.loc[window_indices].idxmax()["speed"]
            min_index = running.loc[max_index : window_indices[-1]].idxmin()["speed"]
            window_diffs = running.loc[window_indices[0] : min_index].diff()
            min_diff_index = window_diffs.idxmin()["speed"]
            max_decel_point = running.loc[min_diff_index]["timestamps"]
            if len(deceleration_times) > 0:
                if max_decel_point - deceleration_times[-1] < 0.5:
                    continue
            deceleration_times.append(max_decel_point)

    deceleration_times = np.array(deceleration_times)
    passive_deceleration_times = deceleration_times[
        (deceleration_times > passive_start) & (deceleration_times < passive_end)
    ]

    return passive_acceleration_times, passive_deceleration_times


def unit_peth(
    session,
    alignment_func: Callable,
    time_before: float,
    time_after: float,
    binsize: float,
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
    alignment_times_1, alignment_times_2 = alignment_func(session)

    units = session.get_units()
    units = units[units["quality"] == "good"]

    total_time = time_before + time_after
    time_bins = int(total_time / binsize)

    condition_peths = []
    for alignment_times in [alignment_times_1, alignment_times_2]:
        if len(alignment_times) < 5:
            condition_peths.append(np.full((len(units), time_bins), np.nan))
            unit_ids = units.index.values
            continue

        peths = []
        unit_ids = []
        for unit_id in units.index.values:
            spike_times = session.spike_times[unit_id]
            peth, _ = makePSTH_numba(
                spike_times,
                alignment_times - time_before,
                total_time,
                binSize=binsize,
            )

            peths.append(peth[:time_bins])

        condition_peths.append(peths)

    return np.array(condition_peths[0]), np.array(condition_peths[1]), unit_ids


def unit_peth_from_cache(
    session_id: int,
    cache_dir: Path,
    cache_source: str,
    manifest: str,
    stimulus_block: int,
    time_before: float,
    time_after: float,
    binsize: float,
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
    cache = load_vbn_cache(cache_dir, cache_source=cache_source, manifest=manifest)
    session = cache.get_ecephys_session(int(session_id))
    alignment_func = partial(
        find_running_acceleration_deceleration_times,
        stimulus_block=stimulus_block,
    )
    return unit_peth(session, alignment_func, time_before, time_after, binsize)


def unit_averaged_psth_time_aligned(
    session_ids: Iterable[int],
    cache_dir: Path,
    cache_source: str = "s3",
    manifest: str = DEFAULT_MANIFEST,
    stimulus_block: int = 5,
    time_before: float = 0.5,
    time_after: float = 0.5,
    binsize: float = 0.001,
    max_workers: int = 20,
) -> Tuple[Sequence[np.ndarray], Sequence[np.ndarray], Sequence[np.ndarray]]:
    """Build run-start aligned PETH arrays without import-time global cache access."""

    session_ids = list(session_ids)
    if max_workers == 1:
        session_data_1 = []
        session_data_2 = []
        unit_ids = []
        for session_id in tqdm.tqdm(session_ids, total=len(session_ids), leave=True):
            try:
                data = unit_peth_from_cache(
                    int(session_id),
                    cache_dir,
                    cache_source,
                    manifest,
                    stimulus_block,
                    time_before,
                    time_after,
                    binsize,
                )
                session_data_1.append(data[0])
                session_data_2.append(data[1])
                unit_ids.append(data[2])
            except Exception as exc:
                print(f"{session_id} generated an exception: {exc}")
        return session_data_1, session_data_2, unit_ids

    session_data_1 = []
    session_data_2 = []
    unit_ids = []
    with concurrent.futures.ProcessPoolExecutor(max_workers=max_workers) as pool:
        future_to_session = {}
        for session_id in session_ids:
            future = pool.submit(
                unit_peth_from_cache,
                int(session_id),
                cache_dir,
                cache_source,
                manifest,
                stimulus_block,
                time_before,
                time_after,
                binsize,
            )
            future_to_session[future] = session_id

        for future in tqdm.tqdm(
            concurrent.futures.as_completed(future_to_session),
            total=len(future_to_session),
            leave=True,
        ):
            session_id = future_to_session[future]
            try:
                data = future.result()
                session_data_1.append(data[0])
                session_data_2.append(data[1])
                unit_ids.append(data[2])
            except Exception as exc:
                print(f"{session_id} generated an exception: {exc}")

    return session_data_1, session_data_2, unit_ids


def summarize_run_start_peths(
    acceleration_peths: Sequence[np.ndarray],
    deceleration_peths: Sequence[np.ndarray],
    unit_id_groups: Sequence[np.ndarray],
) -> dict:
    running_df = {}
    for acceleration_session_peths, deceleration_session_peths, unit_ids_for_session in zip(
        acceleration_peths, deceleration_peths, unit_id_groups
    ):
        if len(unit_ids_for_session) > 0:
            for acceleration_peth, deceleration_peth, unit_id in zip(
                acceleration_session_peths,
                deceleration_session_peths,
                unit_ids_for_session,
            ):
                running_df[unit_id] = {
                    "acceleration": acceleration_peth,
                    "deceleration": deceleration_peth,
                }
    return running_df


def build_run_start_aligned_summary(
    unit_table_file: Path,
    output_file: Path,
    cache_dir: Path,
    cache_source: str = "s3",
    manifest: str = DEFAULT_MANIFEST,
    session_ids: Optional[Sequence[str]] = None,
    stimulus_block: int = 5,
    time_before: float = 0.5,
    time_after: float = 0.5,
    binsize: float = 0.001,
    max_workers: int = 20,
) -> dict:
    """Build ``run_start_aligned_unit_data.pkl`` using the notebook settings."""

    units = read_table(unit_table_file)
    parsed_session_ids = parse_session_ids(session_ids)
    if parsed_session_ids is None:
        if "no_anomalies" not in units.columns:
            raise ValueError("unit table must contain no_anomalies when --session-ids is omitted")
        parsed_session_ids = units.loc[units["no_anomalies"], "ecephys_session_id"].dropna().unique()

    acceleration_peths, deceleration_peths, unit_id_groups = unit_averaged_psth_time_aligned(
        parsed_session_ids,
        cache_dir=cache_dir,
        cache_source=cache_source,
        manifest=manifest,
        stimulus_block=stimulus_block,
        time_before=time_before,
        time_after=time_after,
        binsize=binsize,
        max_workers=max_workers,
    )
    running_df = summarize_run_start_peths(
        acceleration_peths,
        deceleration_peths,
        unit_id_groups,
    )
    write_pickle(running_df, output_file)
    return running_df


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--event", choices=["lick", "run_start"], required=True)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--output-file", type=Path, required=True)
    parser.add_argument("--session-ids", nargs="+", help="Optional session IDs to process")

    parser.add_argument("--active-tensor-file", type=Path, help="Required for --event lick")
    parser.add_argument("--stim-table-file", type=Path, help="Required for --event lick")
    parser.add_argument("--baseline-length", type=int, default=750)
    parser.add_argument("--response-window-length", type=int, default=1500)

    parser.add_argument("--cache-dir", type=Path, help="Required for --event run_start")
    parser.add_argument("--cache-source", choices=["s3", "local"], default="s3")
    parser.add_argument("--manifest", default=DEFAULT_MANIFEST)
    parser.add_argument("--stimulus-block", type=int, default=5)
    parser.add_argument("--time-before", type=float, default=0.5)
    parser.add_argument("--time-after", type=float, default=0.5)
    parser.add_argument("--binsize", type=float, default=0.001)
    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)

    if args.event == "lick":
        if args.active_tensor_file is None or args.stim_table_file is None:
            raise ValueError("--event lick requires --active-tensor-file and --stim-table-file")
        build_lick_aligned_summary(
            active_tensor_file=args.active_tensor_file,
            stim_table_file=args.stim_table_file,
            unit_table_file=args.unit_table_file,
            output_file=args.output_file,
            session_ids=args.session_ids,
            baseline_length=args.baseline_length,
            response_window_length=args.response_window_length,
        )
    elif args.event == "run_start":
        if args.cache_dir is None:
            raise ValueError("--event run_start requires --cache-dir")
        build_run_start_aligned_summary(
            unit_table_file=args.unit_table_file,
            output_file=args.output_file,
            cache_dir=args.cache_dir,
            cache_source=args.cache_source,
            manifest=args.manifest,
            session_ids=args.session_ids,
            stimulus_block=args.stimulus_block,
            time_before=args.time_before,
            time_after=args.time_after,
            binsize=args.binsize,
            max_workers=args.max_workers,
        )


if __name__ == "__main__":
    main()
