"""
Run session-level decoding outputs used by Figure3_FigureS8.

This replaces the legacy hard-coded session decoding launcher with a
public-runnable CLI while preserving the original decoder defaults.
"""

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

import h5py
import pandas as pd

import _bootstrap  # noqa: F401
import decoding_utils as du


DEFAULT_REGIONS = ("LP", "VISp")
DEFAULT_CLUSTERS = ("sensory",)
DEFAULT_UNIT_SAMPLE_SIZES = (5,)


def _class_weight(value: Optional[str]) -> Optional[str]:
    if value is None or str(value).lower() in {"none", "null"}:
        return None
    return value


def session_decoding(
    session_id: int,
    label: str,
    active_tensor_file: Path,
    stim_table_file: Path,
    unit_table_file: Path,
    output_dir: Path,
    regions: Sequence[str] = DEFAULT_REGIONS,
    clusters: Sequence[str] = DEFAULT_CLUSTERS,
    unit_sample_sizes: Sequence[int] = DEFAULT_UNIT_SAMPLE_SIZES,
    decode_window_end: int = 750,
    class_weight: Optional[str] = "balanced",
    clustering: str = "new",
) -> None:
    output_dir.mkdir(parents=True, exist_ok=True)
    stim_table = pd.read_csv(stim_table_file)
    unit_table = pd.read_csv(unit_table_file)

    with h5py.File(active_tensor_file, "r") as unit_data:
        for cluster in clusters:
            du.sessionDecoding(
                session_id,
                label,
                cluster,
                unit_table,
                unit_data,
                stim_table,
                tuple(regions),
                list(unit_sample_sizes),
                decode_window_end,
                class_weight=_class_weight(class_weight),
                outputDir=output_dir,
                clustering=clustering,
            )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--session-id", "--session_id", dest="session_id", type=int, required=True)
    parser.add_argument("--label", "--to-decode", "--to_decode", dest="label", 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("--clusters", nargs="+", default=list(DEFAULT_CLUSTERS))
    parser.add_argument("--unit-sample-sizes", nargs="+", type=int, default=list(DEFAULT_UNIT_SAMPLE_SIZES))
    parser.add_argument("--decode-window-end", type=int, default=750)
    parser.add_argument("--class-weight", default="balanced")
    parser.add_argument("--clustering", default="new")
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    session_decoding(
        session_id=args.session_id,
        label=args.label,
        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,
        clusters=args.clusters,
        unit_sample_sizes=args.unit_sample_sizes,
        decode_window_end=args.decode_window_end,
        class_weight=args.class_weight,
        clustering=args.clustering,
    )


if __name__ == "__main__":
    main()
