"""
Build the Figure 7 flash-decoding metric tables.

This moves the value-producing logic from Figure7.ipynb into a public helper:

- per-session ``*_responseWin_20to100.csv`` confidence tables
- the aggregate ``previous_image_and_change_confidence_response_metrics.csv``

The decoder and response-rate calculations preserve the notebook logic while
replacing hard-coded shared-drive paths with command-line arguments.
"""

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

import h5py
import numpy as np
import pandas as pd
import scipy.stats
import sklearn
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import balanced_accuracy_score
from sklearn.svm import LinearSVC

import decoding_utils as du
import notebook_utils as nu


REGIONS = ("VISall", "VISp", "VISl", "VISrl", "VISal", "VISpm", "VISam", "LGd", "LP", "SCMRN", "Hipp")
RESPONSE_RATE_REGIONS = ("VISall", "VISp", "VISl", "VISrl", "VISal", "VISpm", "VISam", "LGd", "LP", "SCMRN")
METRICS = ("previous_image_confidence", "change_confidence")
DEFAULT_UNIT_SAMPLE_SIZES = (20, 40, -1)
DEFAULT_RESPONSE_WINDOW = slice(20, 100)
RESPONSE_TYPES = (
    "private_nonchange",
    "shared_nonchange",
    "private_hit",
    "shared_hit",
    "private_fa",
    "shared_fa",
    "omission",
    "postomission",
)


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 unit_sample_name(unit_sample_size: int):
    return unit_sample_size if unit_sample_size >= 1 else "all"


def decode_single_class(
    spikes: np.ndarray,
    labels: np.ndarray,
    n_cross_val: int,
    model,
    unit_sample_sizes: Sequence[int],
) -> dict:
    """
    Decode one binary class label from a units-by-trials response matrix.

    This is the Figure7 notebook ``decodeSingleClass`` logic with snake-case
    naming only.
    """
    n_units = spikes.shape[0]
    warnings.filterwarnings("ignore")
    summary = {
        sample_size: {
            metric: []
            for metric in (
                "trainAccuracy",
                "featureWeights",
                "accuracy",
                "prediction",
                "confidence",
                "balanced_accuracy",
            )
        }
        for sample_size in unit_sample_sizes
    }

    for sample_size in unit_sample_sizes:
        if n_units < sample_size:
            continue
        if sample_size > 1:
            if sample_size == n_units:
                unit_samples = [np.arange(n_units)]
            else:
                n_samples = int(math.ceil(math.log(0.01) / math.log(1 - sample_size / n_units)))
                unit_samples = [
                    np.random.choice(n_units, sample_size, replace=False)
                    for _ in range(n_samples)
                ]
        elif sample_size == 1:
            unit_samples = [[i] for i in range(n_units)]
        else:
            unit_samples = [np.arange(n_units)]

        for metric in summary[sample_size]:
            summary[sample_size][metric].append([])

        for unit_sample in unit_samples:
            cv = du.trainDecoder(model, spikes[unit_sample].T, labels, n_cross_val)
            summary[sample_size]["trainAccuracy"][-1].append(np.mean(cv["train_score"]))
            summary[sample_size]["featureWeights"][-1].append(np.mean(cv["coef"], axis=0).squeeze())
            summary[sample_size]["accuracy"][-1].append(np.mean(cv["test_score"]))
            summary[sample_size]["prediction"][-1].append(cv["predict"])
            summary[sample_size]["confidence"][-1].append(cv["decision_function"])
            summary[sample_size]["balanced_accuracy"][-1].append(
                sklearn.metrics.balanced_accuracy_score(
                    labels.astype(bool),
                    cv["predict"].astype(bool),
                )
            )

        for metric in summary[sample_size]:
            if metric == "prediction":
                summary[sample_size][metric][-1] = scipy.stats.mode(
                    summary[sample_size][metric][-1],
                    axis=0,
                )[0][0]
            elif metric != "featureWeights":
                summary[sample_size][metric][-1] = np.median(
                    summary[sample_size][metric][-1],
                    axis=0,
                )

    warnings.filterwarnings("default")
    return summary


def add_previous_image_skipping_omitted(stim: pd.DataFrame) -> pd.DataFrame:
    previous_image_ids = np.insert(stim["image_name"].values, 0, "omitted")[:-1]
    previous_image_ids_skipping_omitted = []
    for index, image_id in enumerate(previous_image_ids):
        if image_id != "omitted":
            previous_image_ids_skipping_omitted.append(image_id)
        else:
            if index == 0:
                previous_image_ids_skipping_omitted.append(np.nan)
            else:
                previous_image_ids_skipping_omitted.append(previous_image_ids[index - 1])

    stim = stim.copy()
    stim["previous_image_skipping_omitted"] = np.array(previous_image_ids_skipping_omitted)
    return stim


def build_session_flash_metrics(
    session_id: int,
    unit_table: pd.DataFrame,
    stim_table: pd.DataFrame,
    unit_data,
    output_dir: Path,
    regions: Sequence[str] = REGIONS,
    unit_sample_sizes: Sequence[int] = DEFAULT_UNIT_SAMPLE_SIZES,
    response_window: slice = DEFAULT_RESPONSE_WINDOW,
) -> pd.DataFrame:
    stim = stim_table[(stim_table["session_id"] == session_id) & stim_table["active"]].reset_index()
    stim = add_previous_image_skipping_omitted(stim)
    image_ids = stim["image_name"].values
    unique_images = [
        image
        for image in np.unique(stim["image_name"].values)
        if image not in ["im083_r", "im111_r", "omitted"]
    ] + ["im083_r", "im111_r"]

    session_group = unit_data[str(session_id)]
    units = unit_table.set_index("unit_id").loc[session_group["unitIds"][:]]
    spikes = session_group["spikes"]
    high_quality = du.apply_unit_quality_filter(units, no_abnorm=False)
    unit_sample_names = [unit_sample_name(size) for size in unit_sample_sizes]

    for region in regions:
        in_region = du.getUnitsInRegion(units, region)
        final_unit_filter = high_quality & in_region

        cols_to_add = [
            f"{metric}_{region}_{sample_name}"
            for sample_name in unit_sample_names
            for metric in METRICS
        ]
        stim[cols_to_add] = np.nan

        if np.sum(final_unit_filter) < 10:
            continue

        filtered_spikes = np.zeros(
            (final_unit_filter.sum(), spikes.shape[1], spikes.shape[2]),
            dtype=bool,
        )
        for index, unit_index in enumerate(np.where(final_unit_filter)[0]):
            filtered_spikes[index] = spikes[unit_index, :, :]

        flash_response = filtered_spikes[:, :, response_window].mean(axis=2)

        image_decoder_summary = {}
        for image in unique_images:
            labels = image_ids == image
            model = LinearSVC(C=1.0, max_iter=int(1e4), class_weight="balanced")
            image_decoder_summary[image] = decode_single_class(
                flash_response,
                labels,
                5,
                model,
                unit_sample_sizes,
            )

        model = LinearSVC(C=1.0, max_iter=int(1e4), class_weight="balanced")
        change_results = decode_single_class(
            flash_response,
            stim["is_change"].values,
            5,
            model,
            unit_sample_sizes,
        )

        for unit_sample, sample_name in zip(unit_sample_sizes, unit_sample_names):
            for index, row in stim.iterrows():
                previous_image_id = row["previous_image_skipping_omitted"]
                if previous_image_id in image_decoder_summary:
                    previous_results = image_decoder_summary[previous_image_id]
                    if len(previous_results[unit_sample]["confidence"]) > 0:
                        confidence = previous_results[unit_sample]["confidence"][0][index]
                        stim.at[index, f"previous_image_confidence_{region}_{sample_name}"] = confidence

            if len(change_results[unit_sample]["confidence"]) > 0:
                stim[f"change_confidence_{region}_{sample_name}"] = change_results[unit_sample]["confidence"][0]

    cols_to_save = (
        ["session_id", "is_change", "previous_image_skipping_omitted", "image_name"]
        + [
            f"previous_image_confidence_{region}_{sample_name}"
            for region in regions
            for sample_name in unit_sample_names
        ]
        + [
            f"change_confidence_{region}_{sample_name}"
            for region in regions
            for sample_name in unit_sample_names
        ]
    )
    session_metrics = stim[cols_to_save]
    output_dir.mkdir(parents=True, exist_ok=True)
    session_metrics.to_csv(
        output_dir / f"{session_id}_responseWin_{response_window.start}to{response_window.stop}.csv",
        index=False,
    )
    return session_metrics


def build_per_session_flash_metrics(
    unit_table_file: Path,
    stim_table_file: Path,
    active_tensor_file: Path,
    sessions_table_file: Path,
    output_dir: Path,
    session_ids: Optional[Sequence[int]] = None,
    regions: Sequence[str] = REGIONS,
    unit_sample_sizes: Sequence[int] = DEFAULT_UNIT_SAMPLE_SIZES,
    response_window: slice = DEFAULT_RESPONSE_WINDOW,
    random_seed: Optional[int] = None,
) -> None:
    if random_seed is not None:
        np.random.seed(random_seed)

    unit_table = read_table(unit_table_file)
    stim_table = read_table(stim_table_file)
    sessions = read_table(sessions_table_file)
    if session_ids is not None:
        session_ids = [int(session_id) for session_id in session_ids]
        sessions = sessions[sessions["ecephys_session_id"].isin(session_ids)]

    with h5py.File(active_tensor_file, "r") as unit_data:
        for _, session in sessions.iterrows():
            build_session_flash_metrics(
                int(session["ecephys_session_id"]),
                unit_table,
                stim_table,
                unit_data,
                output_dir,
                regions=regions,
                unit_sample_sizes=unit_sample_sizes,
                response_window=response_window,
            )


def load_flash_metrics_with_stim_table(
    flash_metrics_dir: Path,
    stim_table: pd.DataFrame,
    response_window: slice = DEFAULT_RESPONSE_WINDOW,
) -> pd.DataFrame:
    pattern = f"responseWin_{response_window.start}to{response_window.stop}"
    stim_files = [
        path
        for path in sorted(flash_metrics_dir.iterdir())
        if pattern in path.name and path.suffix == ".csv"
    ]

    dataframes = []
    for stim_file in stim_files:
        df = pd.read_csv(stim_file)
        session_id = int(stim_file.name.split(".")[0].split("_")[0])
        session_stim = stim_table[stim_table["session_id"] == session_id]
        df = df.merge(
            session_stim.reset_index(),
            left_index=True,
            right_index=True,
            suffixes=("", "_stimtable"),
        )
        dataframes.append(df)

    stimtable_with_flash_metrics = pd.concat(dataframes)
    return stimtable_with_flash_metrics.drop(
        columns=[col for col in stimtable_with_flash_metrics if "Unnamed" in col]
    )


def train_model(model, features: np.ndarray, labels: np.ndarray, n_splits: int) -> dict:
    class_values = np.unique(labels)
    n_classes = len(class_values)
    n_samples = len(labels)
    cv = {"estimator": [sklearn.base.clone(model) for _ in range(n_splits)]}
    cv["train_balanced_accuracy"] = []
    cv["test_balanced_accuracy"] = []
    cv["predict"] = np.full(n_samples, "", dtype="O")
    cv["predict_proba"] = np.full((n_samples, n_classes), np.nan)
    cv["coef"] = []
    model_methods = dir(model)
    train_indices, test_indices = du.getTrainTestSplits(labels, n_splits, hasClasses=False)
    for estimator, train, test in zip(cv["estimator"], train_indices, test_indices):
        estimator.fit(features[train], labels[train])
        cv["train_balanced_accuracy"].append(
            balanced_accuracy_score(labels[train], estimator.predict(features[train]))
        )
        cv["test_balanced_accuracy"].append(
            balanced_accuracy_score(labels[test], estimator.predict(features[test]))
        )
        cv["predict"][test] = estimator.predict(features[test])
        for method in ("predict_proba",):
            if method in model_methods:
                cv[method][test] = getattr(estimator, method)(features[test])
        for attr in ("coef_",):
            if attr in estimator.__dict__:
                cv[attr[:-1]].append(getattr(estimator, attr))
    return cv


def build_session_model_results(
    stimtable_with_flash_metrics: pd.DataFrame,
    sessions: pd.DataFrame,
    areas: Sequence[str] = REGIONS,
    metrics_to_run: Sequence[str] = METRICS,
    sample_sizes: Sequence[str] = ("20", "40", "all"),
) -> pd.DataFrame:
    session_model_results = {
        session_id: {}
        for session_id in stimtable_with_flash_metrics["session_id"].unique()
    }
    columns_to_predict = ["is_change", "lickbout_for_flash_during_response_window"]
    filters = [lambda table: [True] * len(table), lambda table: ~table["is_change"]]

    warnings.filterwarnings("ignore")
    for session_id in stimtable_with_flash_metrics["session_id"].unique():
        stim = stimtable_with_flash_metrics[
            stimtable_with_flash_metrics["session_id"] == session_id
        ]
        stim = stim[
            stim["engaged"]
            & stim["no_abnorm"]
            & ~stim["grace_period_after_hit"]
        ]

        for sample_size in sample_sizes:
            for area in areas:
                for metrics in metrics_to_run:
                    if isinstance(metrics, str):
                        metric_columns = (
                            metrics + "_" + area + "_" + sample_size
                            if "lick" not in metrics
                            else metrics
                        )
                        metric_name = metric_columns
                    else:
                        metric_columns = [
                            metric + "_" + area if "lick" not in metric else metric
                            for metric in metrics
                        ]
                        metric_name = "combo"
                        for metric in metrics_to_run:
                            metric_name = metric_name + "__" + metric

                    for column_to_predict, row_filter in zip(columns_to_predict, filters):
                        curated_table = stim[row_filter(stim)]
                        if len(curated_table) == 0:
                            continue

                        features = curated_table[metric_columns].to_numpy().astype(float)
                        no_nan_indices = [
                            row_index
                            for row_index, row in enumerate(features)
                            if not np.any(np.isnan(row))
                        ]
                        if len(no_nan_indices) == 0:
                            session_model_results[session_id].update(
                                {
                                    "train_balanced_accuracy" + "_" + metric_name: np.nan,
                                    "test_balanced_accuracy" + "_" + metric_name: np.nan,
                                }
                            )
                            continue

                        features_nonan = features[no_nan_indices].reshape(-1, 1)
                        features_nonan = (
                            features_nonan - features_nonan.mean(axis=0)
                        ) / features_nonan.std(axis=0)
                        labels = curated_table[column_to_predict].values
                        labels_nonan = labels[no_nan_indices]
                        model = LogisticRegression(
                            random_state=0,
                            solver="liblinear",
                            class_weight="balanced",
                        )
                        result = train_model(model, features_nonan, labels_nonan, 5)
                        session_model_results[session_id].update(
                            {
                                f"{column_to_predict}_train_balanced_accuracy_{metric_name}": np.mean(
                                    result["train_balanced_accuracy"]
                                ),
                                f"{column_to_predict}_test_balanced_accuracy_{metric_name}": np.mean(
                                    result["test_balanced_accuracy"]
                                ),
                            }
                        )

    warnings.filterwarnings("default")
    return (
        pd.DataFrame.from_dict(session_model_results, orient="index")
        .merge(sessions, left_index=True, right_on="ecephys_session_id")
        .set_index("ecephys_session_id")
    )


def add_engagement_metrics(sessions: pd.DataFrame, stim_table: pd.DataFrame) -> pd.DataFrame:
    sessions = sessions.copy()
    dprimes = []
    hit_counts = []
    for _, session in sessions.iterrows():
        session_id = session["ecephys_session_id"]
        dprimes.append(nu.get_session_engaged_dprime(stim_table, session_id))
        hit_counts.append(nu.get_session_engaged_hit_count(stim_table, session_id))

    sessions["engaged_dprime"] = dprimes
    sessions["engaged_hitcount"] = hit_counts
    return sessions


def make_empty_response_rate_row() -> dict:
    return {
        f"{image_set}_{experience_level}_{metric}_{area}_{sample_size}_{response_type}": np.nan
        for image_set in ["G", "H"]
        for experience_level in ["Familiar", "Novel"]
        for metric in METRICS
        for area in RESPONSE_RATE_REGIONS
        for sample_size in ["20", "40", "all"]
        for response_type in RESPONSE_TYPES
    }


def build_response_rate_summary(
    stimtable_with_flash_metrics: pd.DataFrame,
    sessions: pd.DataFrame,
    stim_table: pd.DataFrame,
) -> pd.DataFrame:
    sessions = add_engagement_metrics(sessions, stim_table)
    good_sessions = sessions[
        sessions["abnormal_histology"].isnull()
        & sessions["abnormal_activity"].isnull()
    ]
    good_behavior_sessions = good_sessions[
        (good_sessions["engaged_dprime"] >= 1)
        & (good_sessions["engaged_hitcount"] >= 50)
    ]

    response_rate_summary = {
        session_id: make_empty_response_rate_row()
        for session_id in good_behavior_sessions["ecephys_session_id"].values
    }

    for _, session in good_behavior_sessions.iterrows():
        experience_level = session["experience_level"]
        image_set = session["image_set"]
        session_id = session["ecephys_session_id"]
        session_stim_table = stimtable_with_flash_metrics[
            stimtable_with_flash_metrics["session_id"] == session_id
        ]

        for metric in METRICS:
            for area in RESPONSE_RATE_REGIONS:
                for sample_size in ["20", "40", "all"]:
                    metric_name = f"{metric}_{area}_{sample_size}"
                    prefix = f"{image_set}_{experience_level}_{metric_name}"
                    response_rate_summary[session_id][f"{prefix}_private_nonchange"] = (
                        nu.get_private_nonchange_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_private_hit"] = (
                        nu.get_private_change_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_private_fa"] = (
                        nu.get_private_catch_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_shared_nonchange"] = (
                        nu.get_shared_nonchange_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_shared_hit"] = (
                        nu.get_shared_change_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_shared_fa"] = (
                        nu.get_shared_catch_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_omission"] = (
                        nu.get_omission_mean(session_stim_table, metric_name)
                    )
                    response_rate_summary[session_id][f"{prefix}_postomission"] = (
                        nu.get_post_omission_mean(session_stim_table, metric_name)
                    )

    return pd.DataFrame.from_dict(response_rate_summary, orient="index")


def aggregate_flash_decoding_metrics(
    flash_metrics_dir: Path,
    stim_table_file: Path,
    sessions_table_file: Path,
    response_rate_output_file: Path,
    session_model_output_file: Optional[Path] = None,
    response_window: slice = DEFAULT_RESPONSE_WINDOW,
) -> pd.DataFrame:
    stim_table = read_table(stim_table_file)
    sessions = read_table(sessions_table_file)
    stimtable_with_flash_metrics = load_flash_metrics_with_stim_table(
        flash_metrics_dir,
        stim_table,
        response_window=response_window,
    )

    if session_model_output_file is not None:
        session_model_results = build_session_model_results(
            stimtable_with_flash_metrics,
            sessions,
        )
        session_model_output_file.parent.mkdir(parents=True, exist_ok=True)
        session_model_results.to_csv(session_model_output_file)

    response_rate_df = build_response_rate_summary(
        stimtable_with_flash_metrics,
        sessions,
        stim_table,
    )
    response_rate_output_file.parent.mkdir(parents=True, exist_ok=True)
    response_rate_df.to_csv(response_rate_output_file)
    return response_rate_df


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    mode = parser.add_mutually_exclusive_group(required=True)
    mode.add_argument("--per-session", action="store_true")
    mode.add_argument("--aggregate", action="store_true")
    parser.add_argument("--unit-table-file", type=Path)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--active-tensor-file", type=Path)
    parser.add_argument("--sessions-table-file", type=Path, required=True)
    parser.add_argument(
        "--flash-metrics-dir",
        type=Path,
        required=True,
        help="Output directory for --per-session; input directory for --aggregate.",
    )
    parser.add_argument("--output-file", type=Path)
    parser.add_argument("--session-model-output-file", type=Path)
    parser.add_argument("--session-ids", nargs="+", type=int)
    parser.add_argument("--regions", nargs="+", default=list(REGIONS))
    parser.add_argument("--unit-sample-sizes", nargs="+", type=int, default=list(DEFAULT_UNIT_SAMPLE_SIZES))
    parser.add_argument("--response-window-start", type=int, default=DEFAULT_RESPONSE_WINDOW.start)
    parser.add_argument("--response-window-stop", type=int, default=DEFAULT_RESPONSE_WINDOW.stop)
    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)
    response_window = slice(args.response_window_start, args.response_window_stop)

    if args.per_session:
        missing = [
            name
            for name in ("unit_table_file", "active_tensor_file")
            if getattr(args, name) is None
        ]
        if missing:
            raise ValueError(f"--per-session requires: {', '.join('--' + name.replace('_', '-') for name in missing)}")
        build_per_session_flash_metrics(
            unit_table_file=args.unit_table_file,
            stim_table_file=args.stim_table_file,
            active_tensor_file=args.active_tensor_file,
            sessions_table_file=args.sessions_table_file,
            output_dir=args.flash_metrics_dir,
            session_ids=args.session_ids,
            regions=args.regions,
            unit_sample_sizes=args.unit_sample_sizes,
            response_window=response_window,
            random_seed=args.random_seed,
        )
    else:
        if args.output_file is None:
            raise ValueError("--aggregate requires --output-file")
        aggregate_flash_decoding_metrics(
            flash_metrics_dir=args.flash_metrics_dir,
            stim_table_file=args.stim_table_file,
            sessions_table_file=args.sessions_table_file,
            response_rate_output_file=args.output_file,
            session_model_output_file=args.session_model_output_file,
            response_window=response_window,
        )


if __name__ == "__main__":
    main()
