"""
Run the compute-heavy image decoder used by Figure6_FigureS13.

This keeps the original decoder settings for the published
``*_nonchangeRS.npy`` outputs, but requires explicit public input/output paths
instead of reading and writing hard-coded shared-drive locations.
"""

import argparse
from pathlib import Path
from typing import Optional, Sequence

import h5py
import numpy as np
import pandas as pd

import _bootstrap  # noqa: F401
import decoding_utils as du


DEFAULT_REGIONS = (
    "VISall",
    "VISp",
    "VISl",
    "VISrl",
    "VISal",
    "VISpm",
    "VISam",
    "LP",
    "LGd",
    "SCMRN",
)
DEFAULT_UNIT_SAMPLE_SIZES = (10, 20, 40, 80)


def image_decoding(
    session_id: int,
    active_tensor_file: Path,
    stim_table_file: Path,
    unit_table_file: Path,
    output_dir: Path,
    regions: Sequence[str] = DEFAULT_REGIONS,
    unit_sample_sizes: Sequence[int] = DEFAULT_UNIT_SAMPLE_SIZES,
    decode_window_end: int = 300,
    decode_full_timecourse_index: int = 2,
    use_nonchange: bool = True,
    class_weight: Optional[str] = "balanced",
    cell_type: str = "RS",
    output_suffix: str = "nonchangeRS",
    random_seed: Optional[int] = None,
) -> dict:
    if random_seed is not None:
        np.random.seed(random_seed)

    stim_table = pd.read_csv(stim_table_file)
    unit_table = pd.read_csv(unit_table_file)
    output_dir.mkdir(parents=True, exist_ok=True)

    with h5py.File(active_tensor_file, "r") as unit_data:
        decoded = du.decodeImage(
            session_id,
            unit_table,
            unit_data,
            stim_table,
            tuple(regions),
            list(unit_sample_sizes),
            decode_window_end,
            decode_full_timecourse_index=decode_full_timecourse_index,
            use_nonchange=use_nonchange,
            class_weight=class_weight,
            cell_type=cell_type,
        )

    np.save(output_dir / f"{session_id}_{output_suffix}.npy", decoded)
    return decoded


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--session-id", type=int, required=True)
    parser.add_argument("--active-tensor-file", type=Path, required=True)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--regions", nargs="+", default=list(DEFAULT_REGIONS))
    parser.add_argument("--unit-sample-sizes", nargs="+", type=int, default=list(DEFAULT_UNIT_SAMPLE_SIZES))
    parser.add_argument("--decode-window-end", type=int, default=300)
    parser.add_argument("--decode-full-timecourse-index", type=int, default=2)
    parser.add_argument("--use-change-flashes", action="store_true")
    parser.add_argument("--class-weight", default="balanced")
    parser.add_argument("--cell-type", default="RS")
    parser.add_argument("--output-suffix", default="nonchangeRS")
    parser.add_argument(
        "--random-seed",
        type=int,
        default=None,
        help="Optional seed for unit subsampling. Defaults to notebook-like unseeded sampling.",
    )
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    image_decoding(
        session_id=args.session_id,
        active_tensor_file=args.active_tensor_file,
        stim_table_file=args.stim_table_file,
        unit_table_file=args.unit_table_file,
        output_dir=args.output_dir,
        regions=args.regions,
        unit_sample_sizes=args.unit_sample_sizes,
        decode_window_end=args.decode_window_end,
        decode_full_timecourse_index=args.decode_full_timecourse_index,
        use_nonchange=not args.use_change_flashes,
        class_weight=args.class_weight,
        cell_type=args.cell_type,
        output_suffix=args.output_suffix,
        random_seed=args.random_seed,
    )


if __name__ == "__main__":
    main()
