"""
Run pooled sensory/action decoding outputs used by Figure3_FigureS8.

This wrapper preserves the original pooled decoder defaults while requiring
explicit table, tensor, and output paths for public reruns.
"""

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

import _bootstrap  # noqa: F401
import decoding_utils as du


def pooled_decoding(
    label: str,
    region: str,
    cluster: str,
    unit_sample_size: int,
    n_pseudo_flashes: int,
    n_unit_samples: int,
    condition: str,
    stim_table_file: Path,
    unit_table_file: Path,
    active_tensor_file: Path,
    passive_tensor_file: Optional[Path],
    output_dir: Path,
    clustering: str = "new",
) -> None:
    if condition == "passive" and passive_tensor_file is None:
        raise ValueError("--passive-tensor-file is required when --condition passive")

    output_dir.mkdir(parents=True, exist_ok=True)
    print(f"calling pooled decoding, {label}, {region} {cluster}")
    du.pooledDecoding(
        label,
        region,
        cluster,
        unit_sample_size,
        n_pseudo_flashes,
        n_unit_samples,
        condition=condition,
        clustering=clustering,
        stim_table_file=stim_table_file,
        unit_table_file=unit_table_file,
        active_tensor_file=active_tensor_file,
        passive_tensor_file=passive_tensor_file,
        outputDir=output_dir,
    )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--label", required=True)
    parser.add_argument("--region", required=True)
    parser.add_argument("--cluster", required=True)
    parser.add_argument("--unit-sample-size", "--unitSampleSize", dest="unit_sample_size", type=int, required=True)
    parser.add_argument("--n-pseudo-flashes", "--nPseudoFlashes", dest="n_pseudo_flashes", type=int, required=True)
    parser.add_argument("--n-unit-samples", "--nUnitSamples", dest="n_unit_samples", type=int, required=True)
    parser.add_argument("--condition", choices=("active", "passive"), default="active")
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--active-tensor-file", type=Path, required=True)
    parser.add_argument("--passive-tensor-file", type=Path, default=None)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--clustering", default="new")
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    pooled_decoding(
        label=args.label,
        region=args.region,
        cluster=args.cluster,
        unit_sample_size=args.unit_sample_size,
        n_pseudo_flashes=args.n_pseudo_flashes,
        n_unit_samples=args.n_unit_samples,
        condition=args.condition,
        stim_table_file=args.stim_table_file,
        unit_table_file=args.unit_table_file,
        active_tensor_file=args.active_tensor_file,
        passive_tensor_file=args.passive_tensor_file,
        output_dir=args.output_dir,
        clustering=args.clustering,
    )


if __name__ == "__main__":
    main()
