"""
build_opto_metrics.py

Iterates over all optotagging sessions in the VBN cache, computes the four
columns required by decoding_utils.py for unit classification:

    pulse_high_mean_evoked_rate_zscored
    pulse_high_first_spike_latency
    pulse_high_first_spike_jitter
    raised_cosine_high_fraction_time_responsive

Saves the result to unit_opto_metrics.csv.

Usage:
    python build_opto_metrics.py \
        --cache-dir /path/to/vbn_s3_cache \
        --output-file /path/to/supplemental_tables/unit_opto_metrics.csv
"""

import argparse
from pathlib import Path

import numpy as np
import pandas as pd

DEFAULT_MANIFEST = "visual-behavior-neuropixels_project_manifest_v0.5.0.json"
CENSOR_PERIOD = 0.0015
DURATIONS = {'pulse': 0.010 - 2 * CENSOR_PERIOD,
             'raised_cosine': 1 - 2 * CENSOR_PERIOD}
BINSIZES = {'pulse': 0.001, 'raised_cosine': 0.01}


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 compute_session_opto_metrics(session):
    """Return a DataFrame of per-unit opto metrics for one session.

    Computes only the metrics needed for SST/VIP classification in
    decoding_utils.py:
        {stim}_{level}_mean_trial_spike_rate  (pulse, all levels – for evoked rate)
        {stim}_{level}_first_spike_latency    (pulse, all levels)
        {stim}_{level}_first_spike_jitter     (pulse, all levels)
        raised_cosine_{level}_fraction_time_responsive
        pulse_baseline_mean / pulse_baseline_std
    """
    from opto_tagging_utils import (
        mean_trial_spike_rate,
        first_spike_latency,
        first_spike_jitter,
        fraction_time_responsive,
        get_baseline_bin_rates,
    )

    spike_times = session.spike_times
    units = session.get_units()

    opto_table = session.optotagging_table
    start_times = opto_table.groupby(['stimulus_name', 'level'])['start_time'].apply(list)
    conditions = start_times.index.get_level_values(0)
    levels = start_times.index.get_level_values(1)

    # Baseline windows: gaps between consecutive opto stimuli
    censor = 0.005
    bl_starts = opto_table['stop_time'].values[:-1] + censor
    bl_ends = opto_table['start_time'].values[1:] - censor
    min_gap = np.min(bl_ends - bl_starts)
    bl_starts = bl_starts + min_gap / 2

    rows = []
    for unit_id, _ in units.iterrows():
        spikes = spike_times[unit_id]
        row = {'uid': unit_id}

        # Baseline stats (pulse window size) for evoked-rate z-scoring
        pulse_baseline = get_baseline_bin_rates(spikes, bl_starts, bl_ends,
                                                binsize=DURATIONS['pulse'])
        row['pulse_baseline_mean'] = np.mean(pulse_baseline)
        row['pulse_baseline_std'] = np.std(pulse_baseline)

        # Per-condition, per-level metrics
        rc_baseline = get_baseline_bin_rates(spikes, bl_starts, bl_ends,
                                             binsize=BINSIZES['raised_cosine'])

        for starts, condition, level in zip(start_times, conditions, levels):
            starts = np.array(starts)
            duration = DURATIONS[condition]
            col_prefix = f'{condition}_{level}'

            if condition == 'pulse':
                row[f'{col_prefix}_mean_trial_spike_rate'] = mean_trial_spike_rate(
                    spikes, starts + CENSOR_PERIOD, duration)
                row[f'{col_prefix}_first_spike_latency'] = first_spike_latency(
                    spikes, starts + CENSOR_PERIOD, duration)
                row[f'{col_prefix}_first_spike_jitter'] = first_spike_jitter(
                    spikes, starts + CENSOR_PERIOD, duration)

            elif condition == 'raised_cosine':
                row[f'{col_prefix}_fraction_time_responsive'] = fraction_time_responsive(
                    spikes, starts, 0, duration, rc_baseline, BINSIZES['raised_cosine'])

        rows.append(row)

    return pd.DataFrame(rows)


def build_unit_opto_metrics(cache, output_file, session_ids=None):
    from opto_tagging_utils import (
        rename_levels_in_metrics_df,
        get_evoked_rates,
    )

    session_table = cache.get_ecephys_session_table()

    # Only sessions from Cre lines with opsin expression
    opto_sessions = session_table[
        session_table['genotype'].str.contains('Sst|Vip', na=False)
    ]
    if session_ids is not None:
        session_ids = [session_id for session_id in session_ids
                       if session_id in opto_sessions.index]
        opto_sessions = opto_sessions.loc[session_ids]

    print(f"Found {len(opto_sessions)} optotagging sessions")

    all_metrics = []
    for i, session_id in enumerate(opto_sessions.index):
        print(f"  [{i+1}/{len(opto_sessions)}] session {session_id}")
        try:
            session = cache.get_ecephys_session(ecephys_session_id=session_id)
            metrics = compute_session_opto_metrics(session)
            metrics = get_evoked_rates(metrics)
            metrics = rename_levels_in_metrics_df(metrics)
            all_metrics.append(metrics)
        except Exception as e:
            print(f"    SKIPPED: {e}")

    if not all_metrics:
        raise RuntimeError("No optotagging metrics were built.")

    unit_opto_metrics = pd.concat(all_metrics, ignore_index=True)
    output_file = Path(output_file)
    output_file.parent.mkdir(parents=True, exist_ok=True)
    unit_opto_metrics.to_csv(output_file, index=False)
    print(f"Saved {len(unit_opto_metrics)} rows to {output_file}")
    return unit_opto_metrics


def parse_args():
    parser = argparse.ArgumentParser(
        description="Build unit_opto_metrics.csv from VBN optotagging sessions."
    )
    parser.add_argument("--cache-dir", "--cache_dir", dest="cache_dir", required=True,
                        help="Path to the VBN cache directory.")
    parser.add_argument("--output-file", "--save-path", "--save_path", dest="output_file",
                        required=True, help="Output path for unit_opto_metrics.csv.")
    parser.add_argument("--session-ids", "--session_ids", dest="session_ids",
                        type=int, nargs="+", default=None,
                        help="Optional optotagging session IDs to process. Defaults to all.")
    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.")
    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,
    )
    build_unit_opto_metrics(
        cache=cache,
        output_file=args.output_file,
        session_ids=args.session_ids,
    )


if __name__ == '__main__':
    main()
