"""Public helper module.

Conceptually, this module contains the compact PCA-specific helpers used by
structured trajectory notebooks.

It exists as a separate unit so PCA preprocessing, trajectory extraction, and
small plotting utilities remain reusable instead of living inside notebooks.

It connects canonical PSTH entries and shared selection helpers to the PCA
figures that render low-dimensional neural trajectories.
"""

from __future__ import annotations

from pathlib import Path
from typing import Any, Dict, Mapping, Optional, Sequence, Set, Tuple

import numpy as np
from scipy.interpolate import CubicSpline

from .core import to_1d
from .selection import get_ccf_mask, get_celltype_mask, get_completion_mask, get_quiet_mask

__all__ = [
    "angle_map",
    "area_sessions_from_attractor",
    "bin_spike_counts",
    "build_area_condition_matrix",
    "compute_vector_field",
    "extract_condition_segment",
    "interpolate_trajectory",
    "linear_color_shade",
    "load_pc_projection_bundle",
    "normalize_activity_matrix",
    "orient_pca_projection",
    "pca_numpy",
    "session_profile_from_angle_map",
    "session_ids_for_area",
    "simultaneous_session_map",
]


def bin_spike_counts(spike_counts: np.ndarray, bin_size_ms: int, original_bin_size_ms: int) -> np.ndarray:
    """Re-bin a `(time, trials, units)` spike tensor to a coarser temporal grid."""

    spike_counts = np.asarray(spike_counts, dtype=float)
    if spike_counts.ndim != 3:
        raise ValueError(f"spike_counts must be 3D, got shape {spike_counts.shape}")
    if bin_size_ms <= 0 or original_bin_size_ms <= 0:
        raise ValueError("bin sizes must be positive")
    if bin_size_ms % original_bin_size_ms != 0:
        raise ValueError("bin_size_ms must be an integer multiple of original_bin_size_ms")

    bin_factor = int(bin_size_ms / original_bin_size_ms)
    if spike_counts.shape[0] % bin_factor != 0:
        raise ValueError(
            f"Time dimension ({spike_counts.shape[0]}) is not divisible by bin factor ({bin_factor})"
        )

    new_time_bins = spike_counts.shape[0] // bin_factor
    reshaped = spike_counts.reshape(
        (bin_factor, new_time_bins, spike_counts.shape[1], spike_counts.shape[2]),
        order="F",
    )
    return np.sum(reshaped, axis=0)


def linear_color_shade(start_color: Sequence[float], end_color: Sequence[float], n: int) -> np.ndarray:
    """Build a linear RGB gradient with `n` rows."""

    if n <= 0:
        return np.zeros((0, 3), dtype=float)
    start = np.asarray(start_color, dtype=float).reshape(1, 3)
    end = np.asarray(end_color, dtype=float).reshape(1, 3)
    return np.column_stack(
        [
            np.linspace(start[0, channel], end[0, channel], n)
            for channel in range(3)
        ]
    )


def load_pc_projection_bundle(path: str | Path) -> dict[str, Any]:
    """Load the saved PC-projection payload used by attractor notebooks.

    Parameters
    ----------
    path:
        NPZ file containing `attractor_results` and the associated time axes.

    Returns
    -------
    dict[str, Any]
        Mapping with `attractor_results`, `params`, `time_axis`, and
        `time_axis_full`.
    """

    path = Path(path)
    with np.load(path, allow_pickle=True) as data:
        if "attractor_results" not in data:
            raise KeyError(f"{path.name} missing key 'attractor_results'")
        if "time_axis" not in data:
            raise KeyError(f"{path.name} missing key 'time_axis'")

        attractor_results = np.asarray(data["attractor_results"]).reshape(-1)[0]
        if not isinstance(attractor_results, dict):
            attractor_results = dict(attractor_results)

        params = {}
        if "params" in data:
            params_obj = np.asarray(data["params"]).reshape(-1)[0]
            params = params_obj if isinstance(params_obj, dict) else dict(params_obj)

        time_axis = np.asarray(data["time_axis"], dtype=float).reshape(-1)
        time_axis_full = np.asarray(data.get("time_axis_full", data["time_axis"]), dtype=float).reshape(-1)

    return {
        "attractor_results": attractor_results,
        "params": params,
        "time_axis": time_axis,
        "time_axis_full": time_axis_full,
    }


def area_sessions_from_attractor(attractor_results: Mapping[str, Any], area_name: str) -> list[dict[str, Any]]:
    """Return the stored session list for one area in an attractor payload."""

    area_obj = attractor_results[str(area_name)]
    return list(area_obj["sessions"])


def session_ids_for_area(attractor_results: Mapping[str, Any], area_name: str) -> list[str]:
    """Return one session-id string per stored session for one area."""

    ids: list[str] = []
    for sess in area_sessions_from_attractor(attractor_results, area_name):
        conds = list(sess["conditions"])
        ids.append(str(conds[0]["sessionID"]))
    return ids


def simultaneous_session_map(
    attractor_results: Mapping[str, Any],
    regionlist: Sequence[str],
    chosen_simultaneous_session: int,
) -> tuple[str, dict[str, int], list[str]]:
    """Map a global simultaneous-session index onto each area's local index.

    Parameters
    ----------
    attractor_results:
        Saved attractor payload keyed by area.
    regionlist:
        Ordered list of areas that must all share the returned session.
    chosen_simultaneous_session:
        One-based simultaneous-session index, matching the original notebook
        convention.

    Returns
    -------
    tuple[str, dict[str, int], list[str]]
        Chosen global session id, one local zero-based session index per area,
        and the full ordered list of common session ids.
    """

    areas = [str(area) for area in regionlist]
    all_session_ids = {area: session_ids_for_area(attractor_results, area) for area in areas}
    common_ids = list(all_session_ids[areas[0]])
    for area in areas[1:]:
        valid_ids = set(all_session_ids[area])
        common_ids = [session_id for session_id in common_ids if session_id in valid_ids]

    if not common_ids:
        raise RuntimeError("No simultaneous common sessions across all requested areas")
    if not (1 <= int(chosen_simultaneous_session) <= len(common_ids)):
        raise IndexError(
            f"chosen_simultaneous_session={chosen_simultaneous_session} outside 1..{len(common_ids)}"
        )

    chosen_id = common_ids[int(chosen_simultaneous_session) - 1]
    session_index_by_area = {
        area: all_session_ids[area].index(chosen_id)
        for area in areas
    }
    return chosen_id, session_index_by_area, common_ids


def compute_vector_field(
    trial_traj: Any,
    *,
    analysis_bins: Any,
    pc1_range: tuple[float, float],
    pc2_range: tuple[float, float],
    n_grid: int = 40,
    bandwidth_factor: float = 0.1,
) -> dict[str, np.ndarray]:
    """Estimate a Gaussian-smoothed vector field in PC space.

    Parameters
    ----------
    trial_traj:
        Trial trajectory tensor with shape `(n_trials, n_pcs, n_time_bins)`.
    analysis_bins:
        Integer bins used to build the local velocity samples.
    pc1_range, pc2_range:
        Plotting ranges for the PC1 and PC2 axes.
    n_grid:
        Number of grid points per axis.
    bandwidth_factor:
        Fraction of the smallest plot span used to set the Gaussian smoothing
        radius.

    Returns
    -------
    dict[str, np.ndarray]
        Grid coordinates and estimated velocity components keyed by
        `pc1_grid`, `pc2_grid`, `dpc1_grid`, and `dpc2_grid`.
    """

    trial_traj = np.asarray(trial_traj, dtype=float)
    if trial_traj.ndim != 3 or trial_traj.shape[1] < 2:
        raise ValueError(f"trial_traj must be (n_trials, >=2, n_time), got {trial_traj.shape}")

    positions: list[list[float]] = []
    velocities: list[list[float]] = []
    n_trials = int(trial_traj.shape[0])
    n_time = int(trial_traj.shape[2])
    valid_bins = np.asarray([b for b in np.asarray(analysis_bins, dtype=int).reshape(-1) if 0 <= b < n_time], dtype=int)
    if valid_bins.size < 2:
        raise ValueError("Need at least two valid analysis bins")

    for itrial in range(n_trials):
        for idx in range(valid_bins.size - 1):
            itime = int(valid_bins[idx])
            itime_next = int(valid_bins[idx + 1])
            pc1_t = float(trial_traj[itrial, 0, itime])
            pc2_t = float(trial_traj[itrial, 1, itime])
            pc1_next = float(trial_traj[itrial, 0, itime_next])
            pc2_next = float(trial_traj[itrial, 1, itime_next])
            if not np.all(np.isfinite([pc1_t, pc2_t, pc1_next, pc2_next])):
                continue
            positions.append([pc1_t, pc2_t])
            velocities.append([pc1_next - pc1_t, pc2_next - pc2_t])

    positions_arr = np.asarray(positions, dtype=float)
    velocities_arr = np.asarray(velocities, dtype=float)
    pc1_grid, pc2_grid = np.meshgrid(
        np.linspace(float(pc1_range[0]), float(pc1_range[1]), int(n_grid)),
        np.linspace(float(pc2_range[0]), float(pc2_range[1]), int(n_grid)),
    )
    dpc1_grid = np.full((int(n_grid), int(n_grid)), np.nan, dtype=float)
    dpc2_grid = np.full((int(n_grid), int(n_grid)), np.nan, dtype=float)

    if positions_arr.size == 0:
        return {
            "pc1_grid": pc1_grid,
            "pc2_grid": pc2_grid,
            "dpc1_grid": dpc1_grid,
            "dpc2_grid": dpc2_grid,
        }

    smooth_radius = float(bandwidth_factor) * min(
        float(pc1_range[1]) - float(pc1_range[0]),
        float(pc2_range[1]) - float(pc2_range[0]),
    )
    sigma = smooth_radius / 2.0
    cutoff = np.exp(-4.5)

    for i in range(int(n_grid)):
        for j in range(int(n_grid)):
            grid_point = np.array([pc1_grid[i, j], pc2_grid[i, j]], dtype=float)
            distances = np.sqrt(np.sum((positions_arr - grid_point) ** 2, axis=1))
            weights = np.exp(-(distances**2) / (2.0 * sigma**2))
            valid = weights > cutoff
            if int(np.sum(valid)) < 3:
                continue

            weights_valid = weights[valid]
            weights_valid = weights_valid / np.sum(weights_valid)
            dpc1_grid[i, j] = float(np.sum(weights_valid * velocities_arr[valid, 0]))
            dpc2_grid[i, j] = float(np.sum(weights_valid * velocities_arr[valid, 1]))

    return {
        "pc1_grid": pc1_grid,
        "pc2_grid": pc2_grid,
        "dpc1_grid": dpc1_grid,
        "dpc2_grid": dpc2_grid,
    }


def angle_map(field1: Mapping[str, Any], field2: Mapping[str, Any]) -> np.ndarray:
    """Return the local angle difference between two vector fields in degrees."""

    u1 = np.asarray(field1["dpc1_grid"], dtype=float)
    v1 = np.asarray(field1["dpc2_grid"], dtype=float)
    u2 = np.asarray(field2["dpc1_grid"], dtype=float)
    v2 = np.asarray(field2["dpc2_grid"], dtype=float)

    valid = np.isfinite(u1) & np.isfinite(v1) & np.isfinite(u2) & np.isfinite(v2)
    dot12 = u1 * u2 + v1 * v2
    mag1 = np.sqrt(u1**2 + v1**2)
    mag2 = np.sqrt(u2**2 + v2**2)
    denom = mag1 * mag2

    cos_theta = np.full_like(dot12, np.nan, dtype=float)
    valid = valid & np.isfinite(denom) & (denom > 0)
    cos_theta[valid] = dot12[valid] / denom[valid]
    cos_theta = np.clip(cos_theta, -1.0, 1.0)
    return np.degrees(np.arccos(cos_theta))


def session_profile_from_angle_map(
    angle_deg: Any,
    pc1_grid: Any,
    pc2_grid: Any,
    p1: Any,
    p2: Any,
    *,
    ext_factor: float,
    n_bins: int,
    band_width: float,
) -> tuple[np.ndarray, np.ndarray]:
    """Project one angle map onto the line between two whisker landmarks.

    Parameters
    ----------
    angle_deg:
        Two-dimensional angle-difference map in degrees.
    pc1_grid, pc2_grid:
        Mesh-grid coordinates associated with `angle_deg`.
    p1, p2:
        Two-dimensional landmark positions defining the profile axis.
    ext_factor:
        Fractional extension applied before `p1` and after `p2`.
    n_bins:
        Number of profile bins along the normalized line.
    band_width:
        Perpendicular strip width retained around the line.

    Returns
    -------
    tuple[np.ndarray, np.ndarray]
        Normalized profile centers and the corresponding mean-angle profile.
        Returns all-`NaN` values when no valid strip samples remain.
    """

    p1 = np.asarray(p1, dtype=float).reshape(-1)
    p2 = np.asarray(p2, dtype=float).reshape(-1)
    if p1.size != 2 or p2.size != 2:
        raise ValueError("p1 and p2 must each contain exactly two coordinates")

    v = p2 - p1
    length = float(np.linalg.norm(v))
    centers = np.linspace(-float(ext_factor), 1.0 + float(ext_factor), int(n_bins))
    if not np.isfinite(length) or length <= 1e-8:
        return centers, np.full((int(n_bins),), np.nan, dtype=float)

    u = v / length
    gx = np.asarray(pc1_grid, dtype=float).reshape(-1)
    gy = np.asarray(pc2_grid, dtype=float).reshape(-1)
    ang = np.asarray(angle_deg, dtype=float).reshape(-1)
    valid = np.isfinite(gx) & np.isfinite(gy) & np.isfinite(ang)
    gx = gx[valid]
    gy = gy[valid]
    ang = ang[valid]
    if ang.size == 0:
        return centers, np.full((int(n_bins),), np.nan, dtype=float)

    diff_vec = np.column_stack([gx - p1[0], gy - p1[1]])
    s = diff_vec @ u
    s_min = -float(ext_factor) * length
    s_max = (1.0 + float(ext_factor)) * length
    keep = (s >= s_min) & (s <= s_max)
    diff_vec = diff_vec[keep]
    ang = ang[keep]
    s = s[keep]
    if ang.size == 0:
        return centers, np.full((int(n_bins),), np.nan, dtype=float)

    proj = np.outer(s, u)
    diff_perp = diff_vec - proj
    d_perp = np.sqrt(np.sum(diff_perp**2, axis=1))
    strip = d_perp <= float(band_width)
    s_strip = s[strip]
    ang_strip = ang[strip]
    if ang_strip.size == 0:
        return centers, np.full((int(n_bins),), np.nan, dtype=float)

    s_norm = s_strip / length
    edges = np.linspace(-float(ext_factor), 1.0 + float(ext_factor), int(n_bins) + 1)
    centers = 0.5 * (edges[:-1] + edges[1:])
    bin_idx = np.digitize(s_norm, edges) - 1
    profile = np.full((int(n_bins),), np.nan, dtype=float)
    for idx in range(int(n_bins)):
        values = ang_strip[bin_idx == idx]
        if values.size:
            profile[idx] = float(np.nanmean(values))
    return centers, profile


def pca_numpy(X: np.ndarray) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
    """Compute a simple SVD-based PCA on row-wise observations."""

    X = np.asarray(X, dtype=float)
    if X.ndim != 2:
        raise ValueError(f"X must be 2D, got shape {X.shape}")

    mu = np.nanmean(X, axis=0, keepdims=True)
    Xc = np.nan_to_num(X - mu, nan=0.0)
    _, singular_values, vt = np.linalg.svd(Xc, full_matrices=False)
    coeff = vt.T
    score = Xc @ coeff
    denom = max(X.shape[0] - 1, 1)
    latent = (singular_values**2) / denom
    total = np.sum(latent)
    explained = (latent / total * 100.0) if total > 0 else np.zeros_like(latent)
    return coeff, score, latent, explained, mu


def normalize_activity_matrix(activity_matrix: np.ndarray, apply_normalization: bool) -> np.ndarray:
    """Prepare a `(units, time)` activity matrix for PCA as `(time, units)` observations."""

    data = np.asarray(activity_matrix, dtype=float)
    if data.ndim != 2:
        raise ValueError(f"activity_matrix must be 2D, got shape {data.shape}")
    if data.size == 0:
        return np.zeros((0, 0), dtype=float)

    if not apply_normalization:
        return data.T

    max_mat = np.nanmax(data, axis=1)
    min_mat = np.nanmin(data, axis=1)
    diff = max_mat - min_mat
    keep_units = np.isfinite(diff) & (diff != 0)
    data = data[keep_units, :]
    if data.size == 0:
        return np.zeros((0, 0), dtype=float)

    norm = np.nanmax(data, axis=1) - np.nanmin(data, axis=1)
    with np.errstate(invalid="ignore", divide="ignore"):
        return (data / norm[:, None]).T


def orient_pca_projection(projected: np.ndarray, trial_type_size: int, delay_start: int, delay_end: int) -> np.ndarray:
    """Flip the first two PCs so the reference trajectory keeps a stable orientation."""

    projected = np.asarray(projected, dtype=float).copy()
    if projected.ndim != 2 or projected.shape[1] < 2 or trial_type_size <= 0:
        return projected

    ref_segment = projected[:trial_type_size, 0:2]
    if ref_segment.shape[0] >= delay_end:
        ref_segment = ref_segment[(delay_start - 1) : delay_end, :]
    if ref_segment.shape[0] < 2:
        return projected

    delta = ref_segment[-1, :] - ref_segment[0, :]
    if np.isfinite(delta[0]) and delta[0] < 0:
        projected[:, 0] *= -1.0
    if np.isfinite(delta[1]) and delta[1] < 0:
        projected[:, 1] *= -1.0
    return projected


def extract_condition_segment(
    projected: np.ndarray,
    condition_index: int,
    trial_type_size: int,
    delay_start: int,
    delay_end: int,
) -> np.ndarray:
    """Return the delay-epoch PC1/PC2 trajectory for one condition."""

    projected = np.asarray(projected, dtype=float)
    start_idx = int(condition_index) * int(trial_type_size)
    end_idx = start_idx + int(trial_type_size)
    segment = projected[start_idx:end_idx, 0:2]
    if segment.shape[0] < delay_end:
        return np.zeros((0, 2), dtype=float)
    return segment[(delay_start - 1) : delay_end, :].copy()


def interpolate_trajectory(segment: np.ndarray, source_step_ms: int, target_step_ms: int) -> np.ndarray:
    """Interpolate a two-dimensional trajectory to a finer temporal grid."""

    segment = np.asarray(segment, dtype=float)
    if segment.ndim != 2 or segment.shape[1] != 2:
        raise ValueError(f"segment must have shape (n, 2), got {segment.shape}")
    if segment.shape[0] < 3:
        return segment.copy()
    if target_step_ms <= 0 or source_step_ms <= 0:
        raise ValueError("step sizes must be positive")
    if target_step_ms >= source_step_ms:
        return segment.copy()

    duration_s = (segment.shape[0] - 1) * (source_step_ms / 1000.0)
    x_base = np.linspace(0.0, duration_s, segment.shape[0])
    x_new = np.arange(0.0, duration_s + target_step_ms / 2000.0, target_step_ms / 1000.0)

    cs1 = CubicSpline(x_base, segment[:, 0])
    cs2 = CubicSpline(x_base, segment[:, 1])
    return np.column_stack([cs1(x_new), cs2(x_new)])


def build_area_condition_matrix(
    entries: Sequence[Mapping[str, Any]],
    *,
    area_name: str,
    area_list: Optional[Mapping[str, Set[str]]],
    trial_types: Sequence[int],
    lick_states: Sequence[int],
    quiet_state: str,
    completion_state: str,
    cell_type: str,
    pca_endbin: int,
    resolution_change: bool = True,
    new_bin_size_ms: int = 50,
    original_bin_size_ms: int = 10,
) -> Dict[str, Optional[np.ndarray]]:
    """Assemble a concatenated `(units, total_time)` matrix for one area."""

    if len(trial_types) != len(lick_states):
        raise ValueError("trial_types and lick_states must have the same length")

    condition_blocks: list[Optional[np.ndarray]] = []
    trial_type_size = int(pca_endbin)

    for trial_type, lick_state in zip(trial_types, lick_states):
        concat_sig: Optional[np.ndarray] = None

        for entry in entries:
            trial = to_1d(entry["trial_type"], float)
            lick = to_1d(entry["lick_flag"], bool)
            curr_sp = np.asarray(entry["spike_counts"], dtype=float)
            if curr_sp.ndim != 3:
                continue

            n_time, n_trials, _ = curr_sp.shape
            if trial.size != n_trials or lick.size != n_trials:
                continue

            curr_trial_ind = (
                get_quiet_mask(entry, quiet_state)
                & get_completion_mask(entry, completion_state)
                & (lick.astype(int) == int(lick_state))
                & (trial == float(trial_type))
            )
            if not np.any(curr_trial_ind):
                continue

            curr_cell_ind = get_celltype_mask(entry, cell_type) & get_ccf_mask(
                entry,
                area_name,
                area_list=area_list,
                enable_ccf_filter=True,
            )
            if not np.any(curr_cell_ind):
                continue

            if resolution_change:
                curr_sp = bin_spike_counts(curr_sp, new_bin_size_ms, original_bin_size_ms)
                window_centers = np.arange(
                    -1 + new_bin_size_ms / 1000.0,
                    2 + 1e-12,
                    new_bin_size_ms / 1000.0,
                )
            else:
                window_centers = to_1d(entry["trial_timestamps"], float)
                if window_centers.size != n_time:
                    continue

            with np.errstate(all="ignore"):
                curr_sp_trial_mean = np.nanmean(curr_sp[:, curr_trial_ind, :], axis=1)
            if curr_sp_trial_mean.ndim == 1:
                curr_sp_trial_mean = curr_sp_trial_mean.reshape(-1, 1)

            curr_sp_selected = curr_sp_trial_mean[:, curr_cell_ind]
            if curr_sp_selected.size == 0:
                continue

            b1 = int(np.argmin(np.abs(window_centers - (-1.0))))
            b2 = int(np.argmin(np.abs(window_centers - 0.0)))
            if b2 < b1:
                b1, b2 = b2, b1
            baseline = np.mean(curr_sp_selected[b1 : b2 + 1, :], axis=0, keepdims=True)
            curr_sp_selected = curr_sp_selected - baseline

            if concat_sig is None:
                concat_sig = curr_sp_selected
            else:
                concat_sig = np.concatenate([concat_sig, curr_sp_selected], axis=1)

        if concat_sig is None or concat_sig.size == 0:
            condition_blocks.append(None)
            continue

        trial_type_size = min(int(pca_endbin), int(concat_sig.shape[0]))
        block = concat_sig[:trial_type_size, :].T

        condition_blocks.append(block)

    nonempty_blocks = [block for block in condition_blocks if block is not None and block.size > 0]
    if not nonempty_blocks:
        return {
            "activity_matrix": None,
            "condition_blocks": condition_blocks,
            "condition_slices": np.full((len(condition_blocks), 2), -1, dtype=int),
            "trial_type_size": np.asarray([trial_type_size], dtype=int),
        }

    min_units = min(block.shape[0] for block in nonempty_blocks)
    condition_slices = np.full((len(condition_blocks), 2), -1, dtype=int)
    concatenated: list[np.ndarray] = []
    curr_start = 0
    for idx, block in enumerate(condition_blocks):
        if block is None or block.size == 0:
            continue
        block = block[:min_units, :]
        curr_end = curr_start + block.shape[1]
        condition_slices[idx] = np.array([curr_start, curr_end], dtype=int)
        concatenated.append(block)
        condition_blocks[idx] = block
        curr_start = curr_end

    neurons_activity_condition = np.concatenate(concatenated, axis=1)

    return {
        "activity_matrix": neurons_activity_condition,
        "condition_blocks": condition_blocks,
        "condition_slices": condition_slices,
        "trial_type_size": np.asarray([trial_type_size], dtype=int),
    }
