"""Public helper module.

Conceptually, this module collects the compact helpers used by structured
single-unit example notebooks.

It exists as a separate unit so trial filtering, baseline handling, raster
construction, and PSTH rendering stay reusable across the positive and
negative example-unit panels.

It connects raw PSTH probe entries to notebook-ready figure objects for one
selected probe/unit pair.
"""

from __future__ import annotations

from typing import Any, Mapping, Sequence

import matplotlib.pyplot as plt
import numpy as np
from matplotlib.lines import Line2D

from .figure_data import mean_sem_over_axis
from .selection import get_completion_mask, get_quiet_mask

__all__ = ["build_example_unit_figure"]


def build_example_unit_figure(
    entry: Mapping[str, Any],
    *,
    unit_idx_matlab: int,
    window_centers: Any,
    quiet_state: str,
    baseline_subtraction: bool,
    completion_state: str,
    trial_types: Sequence[int],
    lick_states: Sequence[int],
    trial_type_names: Sequence[str],
    bin_width: float,
    xticks: Sequence[float],
    xticklabels: Sequence[str],
    t_show: tuple[float, float],
    condition_colors: Any,
    baseline_window: tuple[float, float] = (0.95, 1.0),
    event_times: Sequence[float] = (1.0, 1.03),
    figure_size: tuple[float, float] = (8.27, 3.94),
    dpi: int = 180,
    raster_markersize: float = 4.5,
    line_width: float = 1.2,
) -> tuple[Any, dict[str, Any]]:
    """Build the two-panel PSTH/raster figure for one selected unit.

    Parameters
    ----------
    entry:
        Raw or canonical PSTH entry mapping for one probe.
    unit_idx_matlab:
        One-based unit index, kept in the original notebook convention.
    window_centers:
        Shared PSTH time axis.
    quiet_state, baseline_subtraction, completion_state:
        Trial-selection and baseline settings.
    trial_types, lick_states, trial_type_names:
        Parallel condition definitions and display labels.
    bin_width:
        Bin width in seconds used to convert counts to Hz and to set the raster
        jitter range.
    xticks, xticklabels, t_show:
        Display-axis settings.
    condition_colors:
        RGB array with one row per condition.
    baseline_window:
        Time window used for per-trial baseline subtraction.
    event_times:
        Shared vertical event markers for raster and PSTH panels.
    figure_size, dpi:
        Matplotlib figure sizing controls.
    raster_markersize, line_width:
        Styling controls for the raster and PSTH curves.

    Returns
    -------
    tuple[matplotlib.figure.Figure, dict[str, Any]]
        Figure object and a metadata dictionary describing the selected unit.
    """

    centers = np.asarray(window_centers, dtype=float).reshape(-1)
    trial = np.asarray(entry["trial_type"]).reshape(-1)
    lick = np.asarray(entry["lick_flag"]).reshape(-1)
    spike_counts = np.asarray(entry["spike_counts"], dtype=np.float32)
    if spike_counts.ndim != 3:
        raise ValueError(f"spike_counts must be 3D, got {spike_counts.shape}")

    unit_idx = int(unit_idx_matlab) - 1
    if not (0 <= unit_idx < spike_counts.shape[2]):
        raise IndexError(
            f"Unit index {unit_idx_matlab} out of range for probe with {spike_counts.shape[2]} units"
        )

    fig, axs = plt.subplots(
        2,
        1,
        figsize=figure_size,
        dpi=dpi,
        gridspec_kw={"height_ratios": [1.15, 1.0]},
    )
    fig.subplots_adjust(left=0.08, right=0.98, top=0.80, bottom=0.14, hspace=0.12)

    baseline_first_bin = int(np.argmin(np.abs(centers - float(baseline_window[0]))))
    baseline_last_bin = int(np.argmin(np.abs(centers - float(baseline_window[1]))))
    completion_mask = get_completion_mask(entry, str(completion_state))
    quiet_mask = get_quiet_mask(entry, str(quiet_state))

    trial_offset = 0
    ytick_positions = [1]
    condition_counts: list[int] = []
    rng = np.random.default_rng()

    for condition_index, (trial_type, lick_state, trial_type_name) in enumerate(
        zip(trial_types, lick_states, trial_type_names)
    ):
        current_trial_ind = quiet_mask & completion_mask & (trial == int(trial_type)) & (lick == int(lick_state))
        curr_sp_trials = spike_counts[:, current_trial_ind, unit_idx]

        if curr_sp_trials.ndim == 1:
            curr_sp_trials = curr_sp_trials.reshape(-1, 1)

        if baseline_subtraction and curr_sp_trials.size > 0:
            baseline_mean = np.repeat(
                np.mean(curr_sp_trials[baseline_first_bin : baseline_last_bin + 1, :], axis=0, keepdims=True),
                curr_sp_trials.shape[0],
                axis=0,
            )
            curr_sp_trials = curr_sp_trials - baseline_mean

        signal2plot = curr_sp_trials / float(bin_width)
        meansig, semsig = mean_sem_over_axis(signal2plot, axis=1)
        axs[1].fill_between(
            centers,
            meansig - semsig,
            meansig + semsig,
            color=condition_colors[condition_index],
            alpha=0.20,
            linewidth=0,
        )
        axs[1].plot(centers, meansig, color=condition_colors[condition_index], linewidth=line_width)

        # Raster points are jittered inside each time bin so stacked spikes do
        # not collapse into one vertical column.
        n_trials = int(curr_sp_trials.shape[1])
        for trial_i in range(n_trials):
            spike_inds = np.flatnonzero(curr_sp_trials[:, trial_i] > 0)
            if spike_inds.size == 0:
                continue
            base_times = centers[spike_inds]
            jitter = (rng.random(base_times.size) - 0.5) * float(bin_width)
            spk_times = base_times + jitter
            y_vals = (trial_i + 1 + trial_offset) * np.ones_like(spk_times)
            axs[0].plot(spk_times, y_vals, ".", color=condition_colors[condition_index], markersize=raster_markersize)

        trial_offset += n_trials
        ytick_positions.append(trial_offset)
        condition_counts.append(n_trials)

    for ax in axs:
        for event_time in event_times:
            ax.axvline(float(event_time), color="0.3", linewidth=0.9)
        ax.set_xlim(*t_show)
        ax.spines["top"].set_visible(False)
        ax.spines["right"].set_visible(False)
        ax.tick_params(axis="both", labelsize=8.5, width=1.0, length=3.2)

    axs[0].set_ylim(0, trial_offset + 1)
    axs[0].set_ylabel("Trial number", fontsize=9.5)
    axs[0].set_yticks(ytick_positions)
    axs[0].set_yticklabels([str(v) for v in ytick_positions])
    axs[0].set_xticks(list(xticks))
    axs[0].set_xticklabels([])

    axs[1].set_xlabel("Time (s)", fontsize=9.5)
    axs[1].set_ylabel("Firing rate (Hz)", fontsize=9.5)
    axs[1].set_xticks(list(xticks))
    axs[1].set_xticklabels(list(xticklabels))

    current_area = str(entry["probe_location"])
    session_id = str(entry["session_id"])
    axs[0].set_title(
        f"{current_area} session: {session_id}   unit: {unit_idx_matlab}",
        fontsize=10,
        pad=6,
        fontweight="semibold",
    )

    legend_handles = [
        Line2D([0], [0], color=condition_colors[index], linewidth=line_width, label=str(label))
        for index, label in enumerate(trial_type_names)
    ]
    legend = fig.legend(
        handles=legend_handles,
        loc="upper center",
        bbox_to_anchor=(0.5, 0.985),
        ncol=len(legend_handles),
        frameon=True,
        fancybox=False,
        framealpha=1.0,
        facecolor="white",
        edgecolor="black",
        fontsize=8.2,
        handlelength=1.9,
        handletextpad=0.6,
        columnspacing=1.4,
        borderpad=0.45,
        prop={"weight": "bold", "size": 8.2},
    )
    legend.get_frame().set_linewidth(1.3)

    meta = {
        "unit_idx_matlab": int(unit_idx_matlab),
        "unit_idx_python": int(unit_idx),
        "area": current_area,
        "session_id": session_id,
        "condition_trial_counts": condition_counts,
    }
    return fig, meta
