"""
Build the unit RF stats table.

For each session in the VBN cache, computes receptive field metrics for every
unit using ReceptiveFieldMapping_VBN, then concatenates results into a single
table keyed by unit_id.

Inputs
------
cache_dir: path to the VBN cache

Output
------
unit_rf_table.csv  (unit_id + RF metric columns)

Usage
-----
python make_unit_rf_table.py \
    --cache-dir /path/to/vbn_s3_cache \
    --output-file /path/to/supplemental_tables/units_with_rf_stats.csv
"""

import argparse
from pathlib import Path

import pandas as pd


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


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

# ---------------------------------------------------------------------------
# Main build function
# ---------------------------------------------------------------------------

def build_unit_rf_table(cache, output_file, session_ids=None):
    """Build and save unit RF metrics from VBN ecephys sessions."""
    from brain_observatory_utilities.datasets.electrophysiology.\
        receptive_field_mapping import ReceptiveFieldMapping_VBN

    session_table = cache.get_ecephys_session_table()
    session_ids = session_ids if session_ids is not None else session_table.index.values
    print(f"Processing {len(session_ids)} sessions...")

    session_dfs = []
    for session_ind, session_id in enumerate(session_ids):
        if session_ind % 10 == 0:
            print(f"  session {session_ind}/{len(session_ids)}")

        try:
            session = cache.get_ecephys_session(ecephys_session_id=session_id)
            rf = ReceptiveFieldMapping_VBN(session)
            rf_metrics = rf.metrics
            session_dfs.append(rf_metrics)
        except Exception as e:
            print(f"  session {session_id} failed: {e}")
            continue

    if not session_dfs:
        raise RuntimeError("No RF metrics were built.")

    print("Concatenating...")
    output = pd.concat(session_dfs, ignore_index=True)

    n_before = len(output)
    output = output.drop_duplicates(subset="unit_id")
    if len(output) < n_before:
        print(f"  Dropped {n_before - len(output)} duplicate unit rows.")

    output_file = Path(output_file)
    output_file.parent.mkdir(parents=True, exist_ok=True)
    print(f"Saving {len(output)} rows to {output_file}")
    output.to_csv(output_file, index=False)
    print("Done.")
    return output


def parse_args():
    parser = argparse.ArgumentParser(
        description="Build a unit RF metrics table from VBN ecephys 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", "--output_file", dest="output_file", required=True,
                        help="Output CSV path, for example units_with_rf_stats.csv.")
    parser.add_argument("--session-ids", "--session_ids", dest="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.")
    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_rf_table(
        cache=cache,
        output_file=args.output_file,
        session_ids=args.session_ids,
    )


if __name__ == "__main__":
    main()
