"""
Build VBN per-unit spike tensors.

The output HDF5 matches the figure-notebook expectation:

    <ecephys_session_id>/unitIds  (n_units,)
    <ecephys_session_id>/spikes   (n_units, n_flashes, n_bins), bool

Each `spikes` row is a binary spike-count tensor for one unit, aligned to
stimulus presentation start times. Use `--condition active` for
`vbnAllUnitSpikeTensor.hdf5` and `--condition passive` for
`vbnAllUnitSpikeTensor_passive.hdf5`.
"""

import argparse
from pathlib import Path

import h5py
import numpy as np


DEFAULT_MANIFEST = "visual-behavior-neuropixels_project_manifest_v0.5.0.json"


def get_spike_bins(spikeTimes, startTimes, windowDur, binSize=0.001):
    """Original VBN tensor binning helper."""
    bins = np.arange(0, windowDur + binSize, binSize)
    spikes = np.zeros((len(startTimes), bins.size - 1), dtype=bool)
    for i, start in enumerate(startTimes):
        startInd = np.searchsorted(spikeTimes, start)
        endInd = np.searchsorted(spikeTimes, start + windowDur)
        spikes[i] = np.histogram(spikeTimes[startInd:endInd] - start, bins)[0]
    return spikes


def load_cache(cache_dir, manifest=DEFAULT_MANIFEST, cache_source="local"):
    """Load a VisualBehaviorNeuropixelsProjectCache without import-time side effects."""
    from allensdk.brain_observatory.behavior.behavior_project_cache.\
        behavior_neuropixels_project_cache import VisualBehaviorNeuropixelsProjectCache

    if cache_source == "s3":
        cache = VisualBehaviorNeuropixelsProjectCache.from_s3_cache(cache_dir=cache_dir)
    else:
        cache = VisualBehaviorNeuropixelsProjectCache.from_local_cache(cache_dir=cache_dir)

    if manifest:
        cache.load_manifest(manifest)
    return cache


def get_flash_times(session, condition):
    """Return stimulus start times for the requested tensor condition."""
    stim = session.stimulus_presentations
    if condition == "active":
        return stim.loc[stim["active"], "start_time"].values
    if condition == "passive":
        return stim.loc[stim["stimulus_block"] == 5, "start_time"].values
    raise ValueError(f"Unknown condition: {condition}")


def get_good_units(session):
    """Return units matching permissive quality criteria used by the original tensor."""
    units = session.get_units()
    channels = session.get_channels()
    units = units.merge(channels, left_on="peak_channel_id", right_index=True)
    return units[
        (units["quality"] == "good")
        & (units["snr"] > 1)
        & (units["isi_violations"] < 1)
    ]


def write_session_tensor(
    h5_file,
    session,
    condition,
    window_dur=0.75,
    bin_size=0.001,
    compression="gzip",
    compression_opts=4,
):
    """Write one session group to an open HDF5 file."""
    session_id = session.metadata["ecephys_session_id"]
    flash_times = get_flash_times(session, condition)
    good_units = get_good_units(session)
    spike_times = session.spike_times
    n_bins = int(window_dur / bin_size)

    group = h5_file.create_group(str(session_id))
    group.create_dataset(
        "unitIds",
        data=good_units.index.values,
        compression=compression,
        compression_opts=compression_opts,
    )

    spikes = group.create_dataset(
        "spikes",
        shape=(len(good_units), len(flash_times), n_bins),
        dtype=bool,
        chunks=(1, len(flash_times), n_bins),
        compression=compression,
        compression_opts=compression_opts,
    )

    for i, unit_id in enumerate(good_units.index.values):
        spikes[i] = get_spike_bins(
            spike_times[unit_id],
            flash_times,
            window_dur,
            bin_size,
        )

    return {
        "session_id": session_id,
        "n_units": len(good_units),
        "n_flashes": len(flash_times),
        "n_bins": n_bins,
    }


def build_unit_tensor(
    cache,
    output_file,
    condition,
    session_ids=None,
    window_dur=0.75,
    bin_size=0.001,
    overwrite=False,
):
    """Build an active or passive unit spike tensor for selected sessions."""
    sessions = cache.get_ecephys_session_table(filter_abnormalities=False)
    session_ids = session_ids if session_ids is not None else sessions.index.values
    output_file = Path(output_file)
    output_file.parent.mkdir(parents=True, exist_ok=True)
    mode = "w" if overwrite else "x"

    summaries = []
    with h5py.File(output_file, mode) as h5_file:
        for count, session_id in enumerate(session_ids, start=1):
            print(f"[{count}/{len(session_ids)}] Processing session {session_id}", flush=True)
            session = cache.get_ecephys_session(ecephys_session_id=int(session_id))
            summary = write_session_tensor(
                h5_file,
                session,
                condition=condition,
                window_dur=window_dur,
                bin_size=bin_size,
            )
            summaries.append(summary)
            print(
                "  wrote "
                f"{summary['n_units']} units x {summary['n_flashes']} flashes x "
                f"{summary['n_bins']} bins",
                flush=True,
            )

    return summaries


def parse_args():
    parser = argparse.ArgumentParser(
        description="Build active or passive VBN unit spike tensor HDF5 files."
    )
    parser.add_argument("--cache-dir", required=True, help="Path to the VBN cache directory.")
    parser.add_argument("--output-file", required=True, help="Output HDF5 file path.")
    parser.add_argument(
        "--condition",
        choices=("active", "passive"),
        required=True,
        help="Align to active flashes or passive replay flashes.",
    )
    parser.add_argument(
        "--session-ids",
        type=int,
        nargs="+",
        default=None,
        help="Optional session IDs to process. Defaults to all sessions.",
    )
    parser.add_argument(
        "--manifest",
        default=DEFAULT_MANIFEST,
        help=f"Cache manifest to load. Defaults to {DEFAULT_MANIFEST}.",
    )
    parser.add_argument(
        "--cache-source",
        choices=("local", "s3"),
        default="local",
        help="Use from_local_cache or from_s3_cache. Defaults to local.",
    )
    parser.add_argument("--window-dur", type=float, default=0.75, help="Window duration in seconds.")
    parser.add_argument("--bin-size", type=float, default=0.001, help="Bin size in seconds.")
    parser.add_argument(
        "--overwrite",
        action="store_true",
        help="Overwrite the output file if it already exists.",
    )
    return parser.parse_args()


def main():
    args = parse_args()
    cache = load_cache(
        cache_dir=args.cache_dir,
        manifest=args.manifest,
        cache_source=args.cache_source,
    )
    summaries = build_unit_tensor(
        cache=cache,
        output_file=args.output_file,
        condition=args.condition,
        session_ids=args.session_ids,
        window_dur=args.window_dur,
        bin_size=args.bin_size,
        overwrite=args.overwrite,
    )
    total_units = sum(s["n_units"] for s in summaries)
    print(f"Done. Wrote {len(summaries)} sessions and {total_units} unit tensors.")


if __name__ == "__main__":
    main()
