"""
Build Figure 4/5 modulation summary pickles.

This extracts the value-producing portions of the Figure 4/5 notebook while
omitting plotting-only code. The saved pickle structures match the notebook
outputs for change/state and novelty modulation summaries.
"""

import argparse
import pickle
from pathlib import Path
from typing import Dict, Optional, Sequence, Tuple

import numpy as np
import pandas as pd

import decoding_utils as du
import notebook_utils as nu
import vbn_utils


BASE_SLICE = slice(0, 50)
RESP_SLICE = slice(70, 150)
CHANGE_STATE_METRIC = "mod_index_norm"
CHANGE_STATE_AREAS = ("LGd", "LP", "VISp", "VISl", "VISrl", "VISal", "VISpm", "VISam", "SCMRN")
CTX_AREAS = ("VISall", "VISp", "VISl", "VISrl", "VISal", "VISpm", "VISam")
LAYERS = ("2/3", "4", "5", "6")
CELL_TYPES = ("RS", "FS", "SST", "VIP")
NOVELTY_FLASHES = ("nonshared_nonchange", "change")


def read_table(path: Path) -> pd.DataFrame:
    table = pd.read_csv(path)
    unnamed_columns = [col for col in table.columns if col.startswith("Unnamed:")]
    if unnamed_columns:
        table = table.drop(columns=unnamed_columns)
    return table


def read_pickle(path: Path):
    with path.open("rb") as file:
        return pickle.load(file)


def write_pickle(data: object, output_file: Path) -> None:
    output_file.parent.mkdir(parents=True, exist_ok=True)
    with output_file.open("wb") as file:
        pickle.dump(data, file)


def base_subtract(responses: np.ndarray) -> np.ndarray:
    return responses - responses[:, BASE_SLICE].mean(axis=1)[:, None]


def make_label(area, layer: str, cell_type: str) -> str:
    layer_str = f"_{layer}" if layer != "all" else ""
    cell_type_str = f"_{cell_type}" if cell_type != "all" else ""
    return f"{area}{layer_str}{cell_type_str}"


def build_response_summary(
    flash_data: dict,
    unit_ids: Sequence[int],
    flashes: Sequence[str] = ("change", "prechange", "shared_nonchange", "nonshared_nonchange"),
) -> Dict[int, Dict[str, Dict[str, float]]]:
    """
    Build the notebook's baseline-subtracted scalar response summary.

    The notebook multiplies PSTHs by 1000, subtracts each unit's mean over bins
    0:50, and averages bins 70:150 for each condition.
    """

    response_summary = {
        unit_id: {condition: {flash: [] for flash in flashes} for condition in ("active", "passive")}
        for unit_id in unit_ids
    }

    for condition in ("active", "passive"):
        for flash in flashes:
            responses = np.stack([flash_data[unit_id][condition][flash] for unit_id in unit_ids]) * 1000
            responses = base_subtract(responses)
            scalar_responses = np.mean(responses[:, RESP_SLICE], axis=1)
            for index, unit_id in enumerate(unit_ids):
                response_summary[unit_id][condition][flash] = scalar_responses[index]

    return response_summary


def get_response_pair(
    response_summary: dict,
    unit_ids: Sequence[int],
    comparison: str,
    context: str,
) -> Tuple[np.ndarray, np.ndarray]:
    if comparison == "change":
        vals1 = [response_summary[unit_id][context]["change"] for unit_id in unit_ids]
        vals2 = [response_summary[unit_id][context]["prechange"] for unit_id in unit_ids]
    elif comparison == "state":
        vals1 = [response_summary[unit_id]["active"][context] for unit_id in unit_ids]
        vals2 = [response_summary[unit_id]["passive"][context] for unit_id in unit_ids]
    else:
        raise ValueError(f"Unknown comparison: {comparison}")
    return np.array(vals1), np.array(vals2)


def get_unit_modulation_values(
    units: pd.DataFrame,
    response_summary: dict,
    areas,
    layers,
    cell_types,
    clusters,
    comparison: str,
) -> dict:
    contexts = ("active", "passive") if comparison == "change" else ("change", "prechange")
    vals_to_return = {context: {CHANGE_STATE_METRIC: {}} for context in contexts}

    for context in contexts:
        for area in vbn_utils.make_iterable(areas):
            for layer in vbn_utils.make_iterable(layers):
                for cell_type in vbn_utils.make_iterable(cell_types):
                    unit_ids = vbn_utils.get_unit_ids(
                        units,
                        area,
                        cell_types=cell_type,
                        layers=layer,
                        clusters=clusters,
                        clustering="new",
                    )
                    vals1, vals2 = get_response_pair(
                        response_summary,
                        unit_ids,
                        comparison=comparison,
                        context=context,
                    )
                    label = make_label(area, layer, cell_type)
                    vals_to_return[context][CHANGE_STATE_METRIC][label] = nu.get_mod_index_norm(
                        vals1,
                        vals2,
                    )

    return vals_to_return


def get_novelty_modulation_values(
    units: pd.DataFrame,
    response_summary: dict,
    areas,
    layers,
    cell_types,
    clusters,
    flashes: Sequence[str] = NOVELTY_FLASHES,
    num_iterations: int = 1000,
) -> Tuple[dict, dict]:
    mods = {state: {flash: {} for flash in flashes} for state in ("active", "passive")}
    familiar_novel_responses = {
        state: {flash: {} for flash in flashes}
        for state in ("active", "passive")
    }

    for state in ("active", "passive"):
        for flash in flashes:
            for area in vbn_utils.make_iterable(areas):
                for layer in vbn_utils.make_iterable(layers):
                    for cell_type in vbn_utils.make_iterable(cell_types):
                        familiar_responses = []
                        novel_responses = []
                        for responses, experience in zip(
                            [familiar_responses, novel_responses],
                            ["Familiar", "Novel"],
                        ):
                            unit_ids = vbn_utils.get_unit_ids(
                                units,
                                area,
                                layers=layer,
                                cell_types=cell_type,
                                clusters=clusters,
                                experience=experience,
                                responsive=False,
                            )
                            responses.extend(
                                [response_summary[unit_id][state][flash] for unit_id in unit_ids]
                            )

                        label = make_label(area, layer, cell_type)
                        mods[state][flash][label] = nu.get_nov_mod_index_norm_bootstrap(
                            familiar_responses,
                            novel_responses,
                            iterations=num_iterations,
                            aggfunc=np.nanmean,
                        )
                        familiar_novel_responses[state][flash][label] = [
                            familiar_responses,
                            novel_responses,
                        ]

    return mods, familiar_novel_responses


def build_change_state_modulation_summaries(
    units: pd.DataFrame,
    response_summary: dict,
) -> Dict[str, object]:
    return {
        "change_modulation_across_areas.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CHANGE_STATE_AREAS,
            layers="all",
            cell_types="all",
            clusters="sensory",
            comparison="change",
        ),
        "state_modulation_across_areas.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CHANGE_STATE_AREAS,
            layers="all",
            cell_types="all",
            clusters="sensory",
            comparison="state",
        ),
        "change_modulation_across_layers.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers=LAYERS,
            cell_types="all",
            clusters="sensory",
            comparison="change",
        ),
        "state_modulation_across_layers.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers=LAYERS,
            cell_types="all",
            clusters="sensory",
            comparison="state",
        ),
        "change_modulation_across_celltypes.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers="all",
            cell_types=CELL_TYPES,
            clusters="sensory",
            comparison="change",
        ),
        "state_modulation_across_celltypes.pkl": get_unit_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers="all",
            cell_types=CELL_TYPES,
            clusters="sensory",
            comparison="state",
        ),
    }


def build_novelty_modulation_summaries(
    units: pd.DataFrame,
    response_summary: dict,
) -> Dict[str, object]:
    return {
        "novelty_modulation_across_areas.pkl": get_novelty_modulation_values(
            units,
            response_summary,
            CHANGE_STATE_AREAS,
            layers="all",
            cell_types="all",
            clusters="all",
            flashes=NOVELTY_FLASHES,
            num_iterations=100,
        ),
        "novelty_modulation_across_layers.pkl": get_novelty_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers=LAYERS,
            cell_types="all",
            clusters="all",
            flashes=NOVELTY_FLASHES,
            num_iterations=1000,
        ),
        "novelty_modulation_across_celltypes.pkl": get_novelty_modulation_values(
            units,
            response_summary,
            CTX_AREAS,
            layers="all",
            cell_types=CELL_TYPES,
            clusters="all",
            flashes=NOVELTY_FLASHES,
            num_iterations=1000,
        ),
    }


def build_modulation_summary_tables(
    unit_table_file: Path,
    flash_data_file: Path,
    output_dir: Path,
    summary: str = "all",
    random_seed: Optional[int] = None,
) -> Dict[str, object]:
    if random_seed is not None:
        np.random.seed(random_seed)

    units = read_table(unit_table_file)
    flash_data = read_pickle(flash_data_file)
    unit_ids = units.loc[du.apply_unit_quality_filter(units), "unit_id"].values
    response_summary = build_response_summary(flash_data, unit_ids)

    outputs = {}
    if summary in ("all", "change-state"):
        outputs.update(build_change_state_modulation_summaries(units, response_summary))
    if summary in ("all", "novelty"):
        outputs.update(build_novelty_modulation_summaries(units, response_summary))

    output_dir.mkdir(parents=True, exist_ok=True)
    for filename, data in outputs.items():
        write_pickle(data, output_dir / filename)

    return outputs


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--flash-data-file", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument(
        "--summary",
        choices=("all", "change-state", "novelty"),
        default="all",
        help="Which set of modulation summaries to write.",
    )
    parser.add_argument(
        "--random-seed",
        type=int,
        default=None,
        help="Optional seed for novelty bootstrap summaries. Defaults to notebook-like unseeded sampling.",
    )
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    outputs = build_modulation_summary_tables(
        unit_table_file=args.unit_table_file,
        flash_data_file=args.flash_data_file,
        output_dir=args.output_dir,
        summary=args.summary,
        random_seed=args.random_seed,
    )
    for filename in outputs:
        print(f"Wrote {args.output_dir / filename}")


if __name__ == "__main__":
    main()
