"""
Run pooled decoding dropout and sufficiency jobs.

This wrapper keeps the original subset naming and decoding calls, but makes the
unit table, stimulus table, tensor files, and output directory explicit.
"""

import argparse
import os
from pathlib import Path
from typing import Sequence

import numpy as np
import pandas as pd

import _bootstrap  # noqa: F401
import decoding_utils as du


def run_unit_subset_decoding(
    label: str,
    unit_ids: Sequence[int],
    unit_subset_name: str,
    unit_sample_size: int,
    n_pseudo_flashes: int,
    n_unit_samples: int,
    experience: str,
    output_dir: Path,
    stim_table_file: Path,
    active_tensor_file: Path,
    passive_tensor_file: Path,
    use_max_sample_size_available: bool = False,
) -> None:
    du.pooledDecoding_unit_subsets(
        label,
        unit_ids,
        unit_subset_name,
        unit_sample_size,
        n_pseudo_flashes,
        n_unit_samples,
        condition="active",
        experience=experience,
        use_max_sample_size_available=use_max_sample_size_available,
        outputDir=output_dir,
        stim_table_file=stim_table_file,
        active_tensor_file=active_tensor_file,
        passive_tensor_file=passive_tensor_file,
    )


def unit_subset_dropout_decoding(
    unit_set_ids: Sequence[int],
    unit_subset_ids: Sequence[int],
    unit_subset_name: str,
    unit_sample_size: int,
    n_pseudo_flashes: int,
    n_unit_samples: int,
    experience: str,
    output_dir: Path,
    stim_table_file: Path,
    active_tensor_file: Path,
    passive_tensor_file: Path,
) -> None:
    """
    Run dropout and sufficiency decoders for a selected unit subset.

    ``unit_subset_ids`` are excluded for dropout and exclusively used for
    sufficiency, matching the original wrapper behavior.
    """
    unit_ids_dropout = np.setdiff1d(unit_set_ids, unit_subset_ids)
    for label in ("change", "image"):
        run_unit_subset_decoding(
            label,
            unit_ids_dropout,
            unit_subset_name + "_dropout",
            unit_sample_size,
            n_pseudo_flashes,
            n_unit_samples,
            experience,
            output_dir,
            stim_table_file,
            active_tensor_file,
            passive_tensor_file,
        )
        run_unit_subset_decoding(
            label,
            unit_subset_ids,
            unit_subset_name + "_sufficiency",
            unit_sample_size,
            n_pseudo_flashes,
            n_unit_samples,
            experience,
            output_dir,
            stim_table_file,
            active_tensor_file,
            passive_tensor_file,
            use_max_sample_size_available=True,
        )


def run_subset_decoding(
    unit_table_file: Path,
    stim_table_file: Path,
    active_tensor_file: Path,
    passive_tensor_file: Path,
    output_dir: Path,
    unit_set_region: str = "all",
    unit_set_layer: str = "all",
    unit_set_cell_type: str = "all",
    unit_set_cluster: str = "all",
    unit_subset_region: str = "all",
    unit_subset_layer: str = "all",
    unit_subset_cell_type: str = "all",
    unit_subset_cluster: str = "all",
    unit_sample_size: int = 100,
    n_pseudo_flashes: int = 100,
    n_unit_samples: int = 100,
    experience: str = "all",
) -> None:
    output_dir.mkdir(parents=True, exist_ok=True)
    units = pd.read_csv(unit_table_file).set_index("unit_id")
    units["cortical_layer"].replace("3-Feb", "2/3", inplace=True)

    in_region = du.getUnitsInRegion(
        units,
        unit_set_region,
        cell_type=unit_set_cell_type,
        layer=unit_set_layer,
    )
    high_quality = du.apply_unit_quality_filter(units, no_abnorm=True)
    in_cluster = du.get_units_in_cluster(
        units,
        *du.get_clusters_from_cluster_string(unit_set_cluster),
        clustering="new",
    )
    unit_set_ids = units.index[in_region & high_quality & in_cluster].values
    print(f"num units in set: {len(unit_set_ids)}")

    in_region = du.getUnitsInRegion(
        units,
        unit_subset_region,
        cell_type=unit_subset_cell_type,
        layer=unit_subset_layer,
    )
    high_quality = du.apply_unit_quality_filter(units, no_abnorm=True)
    in_cluster = du.get_units_in_cluster(
        units,
        *du.get_clusters_from_cluster_string(unit_subset_cluster),
        clustering="new",
    )
    unit_subset_ids = units.index[in_region & high_quality & in_cluster].values
    print(f"num units in subset: {len(unit_subset_ids)}")

    if len(unit_subset_ids) < 1:
        print("No units in subset, skipping...")
        return

    unit_set_layer_str = unit_set_layer.replace("/", "")
    unit_subset_layer_str = unit_subset_layer.replace("/", "")
    full_model_substring = (
        f"set_{unit_set_region}_{unit_set_layer_str}_{unit_set_cell_type}_"
        f"{unit_set_cluster}_{experience}_full"
    )
    full_model_run = False
    for _, _, files in os.walk(output_dir):
        for filename in files:
            if full_model_substring in filename and filename.split("_")[-4] == str(unit_sample_size):
                full_model_run = True
                break

    if not full_model_run:
        print("Running full model for comparison...")
        for label in ("change", "image"):
            run_unit_subset_decoding(
                label,
                unit_set_ids,
                f"set_{unit_set_region}_{unit_set_layer_str}_{unit_set_cell_type}_{unit_set_cluster}_{experience}_full",
                unit_sample_size,
                n_pseudo_flashes,
                n_unit_samples,
                experience,
                output_dir,
                stim_table_file,
                active_tensor_file,
                passive_tensor_file,
            )
    else:
        print("Full model already run, skipping...")

    unit_subset_dropout_decoding(
        unit_set_ids,
        unit_subset_ids,
        (
            f"set_{unit_set_region}_{unit_set_layer_str}_{unit_set_cell_type}_"
            f"{unit_set_cluster}_{experience}_subset_{unit_subset_region}_"
            f"{unit_subset_layer_str}_{unit_subset_cell_type}_{unit_subset_cluster}_{experience}"
        ),
        unit_sample_size,
        n_pseudo_flashes,
        n_unit_samples,
        experience,
        output_dir,
        stim_table_file,
        active_tensor_file,
        passive_tensor_file,
    )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--active-tensor-file", type=Path, required=True)
    parser.add_argument("--passive-tensor-file", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--unit-set-region", default="all")
    parser.add_argument("--unit-set-layer", default="all")
    parser.add_argument("--unit-set-cell-type", default="all")
    parser.add_argument("--unit-set-cluster", default="all")
    parser.add_argument("--unit-subset-region", default="all")
    parser.add_argument("--unit-subset-layer", default="all")
    parser.add_argument("--unit-subset-cell-type", default="all")
    parser.add_argument("--unit-subset-cluster", default="all")
    parser.add_argument("--experience", default="all")
    parser.add_argument("--unit-sample-size", type=int, default=100)
    parser.add_argument("--n-pseudo-flashes", type=int, default=100)
    parser.add_argument("--n-unit-samples", type=int, default=100)
    return parser


def main() -> None:
    args = build_parser().parse_args()
    print("Running subset decoding with args:")
    print(args)
    run_subset_decoding(
        unit_table_file=args.unit_table_file,
        stim_table_file=args.stim_table_file,
        active_tensor_file=args.active_tensor_file,
        passive_tensor_file=args.passive_tensor_file,
        output_dir=args.output_dir,
        unit_set_region=args.unit_set_region,
        unit_set_layer=args.unit_set_layer,
        unit_set_cell_type=args.unit_set_cell_type,
        unit_set_cluster=args.unit_set_cluster,
        unit_subset_region=args.unit_subset_region,
        unit_subset_layer=args.unit_subset_layer,
        unit_subset_cell_type=args.unit_subset_cell_type,
        unit_subset_cluster=args.unit_subset_cluster,
        experience=args.experience,
        unit_sample_size=args.unit_sample_size,
        n_pseudo_flashes=args.n_pseudo_flashes,
        n_unit_samples=args.n_unit_samples,
    )


if __name__ == "__main__":
    main()
