"""Public helper module.

Conceptually, this module groups low-level behavioral signal helpers used to
build trial-aligned movement summaries and trial annotations.

It exists as a separate unit so signal-window aggregation, quiet-trial
selection, and trial-type labeling can stay reusable across loaders and
analysis pipelines.

It connects raw behavioral time series to the trial-wise masks and aligned
matrices consumed by plotting and decoding modules.
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Iterable, Mapping, MutableMapping, Tuple

import numpy as np


def _nanmedian(x: np.ndarray) -> float:
    x = np.asarray(x, dtype=float)
    if not np.isfinite(x).any():
        return float("nan")
    return float(np.nanmedian(x))


def mad_nan(x: np.ndarray) -> float:
    """Return the NaN-omitting mean absolute deviation of a numeric vector.

    Verified against the Code_M ground truth (psth_10ms.mat): the quiet-trial
    threshold is ``nanmedian(baseline) + mad(baseline)``, and MATLAB's ``mad``
    here effectively *omits* NaN. With 5 out-of-bounds (NaN) baseline trials in
    session sub-PG019/wS2, stripping NaN reproduces MATLAB's quiet masks exactly
    (601 whisker / 540 jaw, 695/695 agreement), whereas *propagating* NaN would
    poison the threshold and mark every trial non-quiet (0 quiet). So strip NaN.
    """

    x = np.asarray(x, dtype=float)
    x = x[np.isfinite(x)]
    if x.size == 0:
        return float("nan")
    mu = float(np.mean(x))
    return float(np.mean(np.abs(x - mu)))


def robust_mode(x: np.ndarray) -> float:
    """Return the most frequent finite value in a 1D numeric vector."""

    x = np.asarray(x, dtype=float).reshape(-1)
    x = x[np.isfinite(x)]
    if x.size == 0:
        return float("nan")

    values, counts = np.unique(x, return_counts=True)
    return float(values[np.argmax(counts)])


def psth_simple(
    spike_times: np.ndarray,
    anchor_times: np.ndarray,
    pre_time: float,
    post_time: float,
    window_size: float,
    window_step: float,
):
    """Compute a simple spike-count PSTH aligned to trial anchors.

    Parameters
    ----------
    spike_times:
        One-dimensional spike-time vector for a single unit.
    anchor_times:
        Trial-alignment timestamps.
    pre_time, post_time:
        Relative window bounds around each anchor in seconds.
    window_size:
        Counting-window size in seconds.
    window_step:
        Step between successive bin edges in seconds.

    Returns
    -------
    tuple[np.ndarray, np.ndarray, np.ndarray]
        Spike rates, window centers, and raw spike counts with shape
        `(bins, trials)`.
    """

    spike_times = np.asarray(spike_times, dtype=float)
    anchor_times = np.asarray(anchor_times, dtype=float)

    if spike_times.size:
        spike_times = np.sort(spike_times)

    relative_edges = np.arange(pre_time, post_time + 1e-12, window_step, dtype=float)
    n_edges = relative_edges.size
    n_bins = n_edges - 1

    spike_counts = np.zeros((n_bins, anchor_times.size), dtype=np.float64)
    if spike_times.size == 0:
        spike_rates = spike_counts / float(window_size)
        window_centers = relative_edges[:n_bins] + float(window_size)
        return spike_rates, window_centers, spike_counts

    # Count spikes trial by trial so each anchor gets its own aligned bin edges.
    for trial_i, anchor in enumerate(anchor_times):
        bin_edges = anchor + relative_edges
        idx = np.searchsorted(spike_times, bin_edges, side="left")
        idx[-1] = np.searchsorted(spike_times, bin_edges[-1], side="right")
        spike_counts[:, trial_i] = np.diff(idx).astype(np.float64)

    spike_rates = spike_counts / float(window_size)
    window_centers = relative_edges[:n_bins] + float(window_size)
    return spike_rates, window_centers, spike_counts


def psth_behavior(
    movement_signals: Mapping[str, np.ndarray],
    signal_time: np.ndarray,
    anchor_times: np.ndarray,
    pre_time: float,
    post_time: float,
    window_size: float,
    window_step: float,
    fs: float,
):
    """Compute trial-aligned moving-window averages for behavioral signals.

    Parameters
    ----------
    movement_signals:
        Mapping from signal name to one-dimensional behavioral trace.
    signal_time:
        Shared timestamp vector for all movement traces.
    anchor_times:
        Trial-alignment timestamps.
    pre_time, post_time:
        Relative analysis window around each anchor in seconds.
    window_size:
        Averaging-window size in seconds.
    window_step:
        Step between successive windows in seconds.
    fs:
        Sampling rate in Hz.

    Returns
    -------
    tuple[dict[str, np.ndarray], dict[str, np.ndarray]]
        Per-signal aligned mean traces with shape `(windows, trials)` and the
        matching window-center vectors.
    """

    signal_time = np.asarray(signal_time, dtype=float)
    anchor_times = np.asarray(anchor_times, dtype=float)

    if signal_time.ndim != 1:
        raise ValueError(f"signal_time must be 1D, got shape {signal_time.shape}")
    if signal_time.size == 0:
        raise ValueError("signal_time is empty")

    def _nearest_frame_idx(times: np.ndarray, t: float) -> int:
        j = int(np.searchsorted(times, t, side="left"))
        if j <= 0:
            return 0
        if j >= times.size:
            return int(times.size - 1)
        if (t - times[j - 1]) <= (times[j] - t):
            return j - 1
        return j

    relative_edges = np.arange(pre_time, post_time - window_size + 1e-12, window_step, dtype=float)
    window_centers = relative_edges + float(window_size)

    window_size_frames = int(round(window_size * fs))
    window_len = window_size_frames + 1
    window_step_frames = int(round(window_step * fs))
    pre_frames = int(round(pre_time * fs))
    post_frames = int(round(post_time * fs))

    relative_frames = np.arange(pre_frames, post_frames - window_size_frames + 1, window_step_frames, dtype=int)
    n_windows = relative_frames.size

    psth_out = {}
    centers_out = {}

    for name, signal in movement_signals.items():
        signal = np.asarray(signal)
        if signal.ndim > 1 and signal.shape[1] == 1:
            signal = signal[:, 0]
        if signal.ndim != 1:
            raise ValueError(f"psth_behavior expects 1D signals; got {name} with shape {signal.shape}")

        mean_signal = np.full((n_windows, anchor_times.size), np.nan, dtype=np.float64)

        # Reuse cumulative sums within each trial span so every moving window
        # can be computed without re-averaging the raw samples.
        for trial_i, anchor in enumerate(anchor_times):
            if anchor < signal_time[0] or anchor > signal_time[-1]:
                continue

            anchor_frame = _nearest_frame_idx(signal_time, float(anchor))
            start_frame = anchor_frame + pre_frames
            end_frame = anchor_frame + post_frames - window_size_frames

            if start_frame < 0 or end_frame < 0:
                continue
            if end_frame + window_size_frames >= signal.size:
                continue

            span = signal[start_frame : end_frame + window_size_frames + 1].astype(float, copy=False)
            finite = np.isfinite(span)
            span_zeronan = np.where(finite, span, 0.0)
            csum = np.cumsum(np.concatenate([[0.0], span_zeronan]))
            ccount = np.cumsum(np.concatenate([[0], finite.astype(np.int32)]))

            starts = np.arange(start_frame, end_frame + 1, window_step_frames, dtype=int)
            offs = starts - start_frame
            sums = csum[offs + window_len] - csum[offs]
            counts = ccount[offs + window_len] - ccount[offs]
            means = sums / float(window_len)
            means[counts < window_len] = np.nan
            mean_signal[: means.size, trial_i] = means.astype(np.float64)

        psth_out[name] = mean_signal
        centers_out[name] = window_centers

    return psth_out, centers_out


@dataclass(frozen=True)
class QuietTrialParams:
    """Parameter bundle for quiet-trial detection on aligned behavior traces."""

    prewhisk_window: Tuple[float, float]
    baseline_window: Tuple[float, float]
    movement_signals: Tuple[str, ...]
    selection_method: str = "mad_all"


def nwb_find_quiet_trial(
    behavior_table: MutableMapping[str, np.ndarray],
    params: QuietTrialParams,
) -> MutableMapping[str, np.ndarray]:
    """Add quiet-trial boolean vectors to a behavior table in place.

    Parameters
    ----------
    behavior_table:
        Mutable trial-aligned behavior payload. Each movement signal is stored
        as a `(time, trials)` array and `trial_timestamps` defines the shared
        alignment axis.
    params:
        Quiet-trial selection windows and movement-signal names.

    Returns
    -------
    MutableMapping[str, np.ndarray]
        The same mapping with one `quiet_trial_<signal>` boolean vector added
        per requested movement signal.
    """

    currwincenter = np.asarray(behavior_table["trial_timestamps"], dtype=float)

    prewin = [
        int(np.argmin(np.abs(currwincenter - params.prewhisk_window[0]))),
        int(np.argmin(np.abs(currwincenter - params.prewhisk_window[1]))),
    ]
    basewin = [
        int(np.argmin(np.abs(currwincenter - params.baseline_window[0]))),
        int(np.argmin(np.abs(currwincenter - params.baseline_window[1]))),
    ]

    prewin.sort()
    basewin.sort()

    def _nanmean_no_warn(x: np.ndarray) -> np.ndarray:
        finite = np.isfinite(x)
        count = finite.sum(axis=0)
        sums = np.where(finite, x, 0.0).sum(axis=0)
        out = np.full(x.shape[1], np.nan, dtype=float)
        mask = count > 0
        out[mask] = sums[mask] / count[mask]
        return out

    for sig_name in params.movement_signals:
        curr_signal = np.asarray(behavior_table[sig_name], dtype=float)
        if curr_signal.ndim != 2:
            raise ValueError(f"Expected behavior_table[{sig_name!r}] to be 2D (time, trials).")

        baseline_slice = curr_signal[basewin[0] : basewin[1] + 1, :]
        prewhisk_slice = curr_signal[prewin[0] : prewin[1] + 1, :]

        baseline_value = (
            np.full(curr_signal.shape[1], np.nan, dtype=float)
            if baseline_slice.size == 0
            else _nanmean_no_warn(baseline_slice)
        )
        prewhisk_value = (
            np.full(curr_signal.shape[1], np.nan, dtype=float)
            if prewhisk_slice.size == 0
            else _nanmean_no_warn(prewhisk_slice)
        )

        threshold = _nanmedian(baseline_value) + mad_nan(baseline_value)

        # `one_by_one` compares each trial to its own baseline, whereas
        # `mad_all` uses one robust session-level threshold across trials.
        if params.selection_method == "one_by_one":
            ind_quiet = prewhisk_value <= baseline_value
        elif params.selection_method == "mad_all":
            ind_quiet = prewhisk_value <= threshold
        else:
            raise ValueError(f"Unknown selection_method: {params.selection_method}")

        behavior_table[f"quiet_trial_{sig_name}"] = ind_quiet.astype(bool)

    return behavior_table


def trial_type_maker(whisker_stim: Iterable, context: Iterable) -> np.ndarray:
    """Map whisker/context conditions to the canonical numeric trial codes.

    Parameters
    ----------
    whisker_stim:
        Trial-wise whisker-stimulation flags.
    context:
        Trial-wise context labels.

    Returns
    -------
    np.ndarray
        One-dimensional numeric trial-type vector.
    """

    whisker_stim = np.asarray(list(whisker_stim)).reshape(-1).astype(bool)
    context = np.asarray(list(context)).reshape(-1)
    if whisker_stim.size != context.size:
        raise ValueError("whisker_stim and context must have the same length")

    context_str = np.asarray([str(x).strip().lower() for x in context], dtype=object)
    # Match MATLAB behavior exactly:
    # - canonical task conditions map to 1..5
    # - rare no_tone + no_stim trials are left at 0 and ignored downstream
    trial_type = np.zeros(context_str.shape, dtype=np.float32)

    go = context_str == "go_tone"
    nogo = context_str == "nogo_tone"
    no_tone = context_str == "no_tone"
    stim = whisker_stim
    nostim = ~whisker_stim

    trial_type[stim & go] = 1
    trial_type[nostim & go] = 2
    trial_type[stim & nogo] = 3
    trial_type[nostim & nogo] = 4
    trial_type[stim & no_tone] = 5

    unmapped = (
        (context_str != "go_tone")
        & (context_str != "nogo_tone")
        & (context_str != "no_tone")
    )
    if np.any(unmapped):
        bad_ctx = np.unique(context_str[unmapped])
        raise ValueError(f"Unmapped trial types for contexts: {bad_ctx.tolist()}")

    return trial_type
