"""Public analysis module.

Conceptually, this module builds coding-direction projections used to compare
task conditions and spontaneous-lick structure across canonical neural entries.

It exists as a separate unit so coding-direction construction, orthogonalized
projection spaces, and holdout bookkeeping stay together instead of being
repeated in notebooks.

It connects canonical PSTH entries to notebook-ready context, stimulus, and
lick projection outputs.
"""



from __future__ import annotations

from typing import Any, Dict, List, Mapping, Optional, Set, Tuple

import numpy as np

from .loading import _entry_get_any, _infer_entry_n_trials, _infer_entry_n_units, _normalize_vector_length
from .math_utils import (
    downsample_balance,
    ensure_trials_by_cells,
    gram_schmidt_columns,
    holdout_split,
    nearest_bin,
    normalize_vec,
    project_trials,
    psth_simple_counts_single_event,
    zscore_cols,
)

from .decoding import get_ccf_mask, get_celltype_mask, get_completion_mask, get_quiet_mask

def run_coding_direction_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    cfg: Mapping[str, Any],
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
    rng_obj: Optional[np.random.Generator] = None,
) -> Dict[str, List[Dict[str, Any]]]:
    """Compute coding-direction projections for canonical PSTH entries.

    Parameters
    ----------
    entries_by_area:
        Mapping from area name to canonical entry list.
    cfg:
        Configuration dictionary with at least:
        `regionlist`, `completion_state`, `quietstate`, `celltype`,
        `windowCenters`, `Win_context`, `Win_lick`, `Win_stim`, `Win_base`.
        Optional keys:
        `min_cells_per_session` (default 5), `holdout` (default 0.7).
    area_list:
        Decoded Area_list structure for optional CCF filtering.
    enable_ccf_filter:
        Whether to apply CCF filter when `area_list` is available.
    rng_obj:
        Random generator used for holdout splits.

    Returns
    -------
    dict[str, list[dict[str, Any]]]
        Area-grouped projection payloads containing context, stimulus, and
        lick projections plus the index masks needed downstream.
    """

    if rng_obj is None:
        rng_obj = np.random.default_rng(0)

    regionlist = list(cfg["regionlist"])
    min_cells = int(cfg.get("min_cells_per_session", 5))
    holdout = float(cfg.get("holdout", 0.7))
    wc = np.asarray(cfg["windowCenters"], dtype=np.float32)

    coding_direction_matrix: Dict[str, List[Dict[str, Any]]] = {area: [] for area in regionlist}

    for current_area in regionlist:
        probe_entries = list(entries_by_area.get(current_area, []))

        for e in probe_entries:
            trial = np.asarray(e["trial"]).reshape(-1)
            lick = np.asarray(e["lick"]).reshape(-1)

            completed_trial_ind = get_completion_mask(e, str(cfg["completion_state"]))
            qind = get_quiet_mask(e, str(cfg["quietstate"]))

            celltype_ind = get_celltype_mask(e, str(cfg["celltype"]))
            ccf_ind = get_ccf_mask(
                e,
                current_area,
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            if np.sum(curr_cell_ind) < min_cells:
                continue

            currsig = np.asarray(e["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class3 = (trial == 5) & (lick == 0)
            # Match the MATLAB reference: the lick/no-lick coding direction is
            # learned only from the whisker trial families {1, 3, 5}, not from
            # every trial in the session.
            class4 = ((trial == 1) | (trial == 3) | (trial == 5)) & (lick == 1)
            class5 = ((trial == 1) | (trial == 3) | (trial == 5)) & (lick == 0)

            class1 = class1 & qind & completed_trial_ind
            class2 = class2 & qind & completed_trial_ind
            class3 = class3 & qind & completed_trial_ind
            class4 = class4 & qind & completed_trial_ind
            class5 = class5 & qind & completed_trial_ind

            id1 = np.where(class1)[0]
            id2 = np.where(class2)[0]
            id3 = np.where(class3)[0]
            id4 = np.where(class4)[0]
            id5 = np.where(class5)[0]

            id1Train, id1Test = holdout_split(id1, holdout=holdout, rng_obj=rng_obj)
            id2Train, id2Test = holdout_split(id2, holdout=holdout, rng_obj=rng_obj)
            id3Train, id3Test = holdout_split(id3, holdout=holdout, rng_obj=rng_obj)
            id4Train, id4Test = holdout_split(id4, holdout=holdout, rng_obj=rng_obj)
            id5Train, id5Test = holdout_split(id5, holdout=holdout, rng_obj=rng_obj)

            if min(id1Train.size, id2Train.size, id3Train.size, id4Train.size, id5Train.size) < 1:
                continue

            n_time, _, n_cells = currsig.shape

            # Build one time-resolved coding vector per signal family before
            # averaging inside the requested analysis windows.
            cd_context = np.zeros((n_time, n_cells), dtype=np.float32)
            for ind_bin in range(n_time):
                x1 = np.atleast_2d(np.squeeze(currsig[ind_bin, id1Train, :]))
                x2 = np.atleast_2d(np.squeeze(currsig[ind_bin, id2Train, :]))
                cd_context[ind_bin, :] = np.nanmean(x1, axis=0) - np.nanmean(x2, axis=0)

            cd_lick = np.zeros((n_time, n_cells), dtype=np.float32)
            for ind_bin in range(n_time):
                x4 = np.atleast_2d(np.squeeze(currsig[ind_bin, id4Train, :]))
                x5 = np.atleast_2d(np.squeeze(currsig[ind_bin, id5Train, :]))
                cd_lick[ind_bin, :] = np.nanmean(x4, axis=0) - np.nanmean(x5, axis=0)

            wc1 = nearest_bin(wc, cfg["Win_context"][0])
            wc2 = nearest_bin(wc, cfg["Win_context"][1])
            if wc2 < wc1:
                wc1, wc2 = wc2, wc1
            coding_direction_context = np.nanmean(cd_context[wc1:wc2 + 1, :], axis=0)
            coding_direction_context = normalize_vec(coding_direction_context)

            wl1 = nearest_bin(wc, cfg["Win_lick"][0])
            wl2 = nearest_bin(wc, cfg["Win_lick"][1])
            if wl2 < wl1:
                wl1, wl2 = wl2, wl1
            coding_direction_lick = np.nanmean(cd_lick[wl1:wl2 + 1, :], axis=0)
            coding_direction_lick = normalize_vec(coding_direction_lick)

            ws1 = nearest_bin(wc, cfg["Win_stim"][0])
            ws2 = nearest_bin(wc, cfg["Win_stim"][1])
            wb1 = nearest_bin(wc, cfg["Win_base"][0])
            wb2 = nearest_bin(wc, cfg["Win_base"][1])
            if ws2 < ws1:
                ws1, ws2 = ws2, ws1
            if wb2 < wb1:
                wb1, wb2 = wb2, wb1

            x3 = np.nanmean(currsig[:, id3Train, :], axis=1)
            if x3.ndim == 1:
                x3 = x3.reshape(-1, 1)
            coding_direction_stim = np.nanmean(x3[ws1:ws2 + 1, :], axis=0) - np.nanmean(x3[wb1:wb2 + 1, :], axis=0)
            coding_direction_stim = normalize_vec(coding_direction_stim)

            # Orthogonalize the three readout axes once so downstream
            # projections separate context, stimulus, and lick structure.
            gs_in = np.column_stack([coding_direction_context, coding_direction_stim, coding_direction_lick])
            v = gram_schmidt_columns(gs_in)
            coding_direction_context = v[:, 0]
            coding_direction_stim = v[:, 1]
            coding_direction_lick = v[:, 2]

            proj_context = project_trials(currsig, coding_direction_context)
            proj_lick = project_trials(currsig, coding_direction_lick)
            proj_stim = project_trials(currsig, coding_direction_stim)

            # Hide the trials used to fit each coding direction so the saved
            # projections can be visualized without reusing the training data.
            id_training_context = (
                np.sort(np.concatenate([id1Train, id2Train]))
                if (id1Train.size + id2Train.size)
                else np.array([], dtype=int)
            )
            id_training_lick = (
                np.sort(np.concatenate([id4Train, id5Train]))
                if (id4Train.size + id5Train.size)
                else np.array([], dtype=int)
            )
            id_training_stim = np.sort(id3Train) if id3Train.size else np.array([], dtype=int)

            if id_training_context.size:
                proj_context[:, id_training_context] = np.nan
            if id_training_lick.size:
                proj_lick[:, id_training_lick] = np.nan
            if id_training_stim.size:
                proj_stim[:, id_training_stim] = np.nan

            coding_direction_matrix[current_area].append(
                {
                    "lickproj": proj_lick,
                    "Contextproj": proj_context,
                    "Stimproj": proj_stim,
                    "index": {
                        "Quiet": qind,
                        "lick": lick,
                        "trial": trial,
                        "completed_trial": completed_trial_ind,
                    },
                }
            )

    return coding_direction_matrix

def run_spontlick_coding_direction_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    cfg: Mapping[str, Any],
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
    rng_obj: Optional[np.random.Generator] = None,
) -> Dict[str, List[Dict[str, Any]]]:
    """Compute spontaneous-lick coding directions for canonical PSTH entries.

    Parameters
    ----------
    entries_by_area:
        Mapping from area name to canonical entry list.
    cfg:
        Configuration dictionary containing context, stimulus, and spontaneous
        lick windows plus bout-quality settings.
    area_list:
        Decoded Area_list structure for optional CCF filtering.
    enable_ccf_filter:
        Whether to apply CCF filter when `area_list` is available.
    rng_obj:
        Random generator used for holdout splits.

    Returns
    -------
    dict[str, list[dict[str, Any]]]
        Area-grouped projection payloads matching the structured notebook
        output format.
    """

    if rng_obj is None:
        rng_obj = np.random.default_rng(0)

    regionlist = list(cfg["regionlist"])
    min_cells = int(cfg.get("min_cells_per_session", 5))
    holdout = float(cfg.get("holdout", 0.7))
    min_spont_licks = int(cfg.get("min_spont_licks", 2))
    wc = np.asarray(cfg["windowCenters"], dtype=np.float32)

    coding_direction_matrix: Dict[str, List[Dict[str, Any]]] = {area: [] for area in regionlist}

    for current_area in regionlist:
        probe_entries = list(entries_by_area.get(current_area, []))

        for e in probe_entries:
            lentry = e.get("lick_entry")
            if lentry is None:
                continue

            trial = np.asarray(e["trial"]).reshape(-1)
            lick = np.asarray(e["lick"]).reshape(-1)

            completed_trial_ind = get_completion_mask(e, str(cfg["completion_state"]))
            qind = get_quiet_mask(e, str(cfg["quietstate"]))

            celltype_ind = get_celltype_mask(e, str(cfg["celltype"]))
            ccf_ind = get_ccf_mask(
                e,
                current_area,
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            if int(np.sum(curr_cell_ind)) < min_cells:
                continue

            currsig = np.asarray(e["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class3 = (trial == 1) | (trial == 3) | (trial == 5)

            class1 = class1 & qind & completed_trial_ind
            class2 = class2 & qind & completed_trial_ind
            class3 = class3 & qind & completed_trial_ind

            id1 = np.where(class1)[0]
            id2 = np.where(class2)[0]
            id3 = np.where(class3)[0]

            id1Train, _ = holdout_split(id1, holdout=holdout, rng_obj=rng_obj)
            id2Train, _ = holdout_split(id2, holdout=holdout, rng_obj=rng_obj)
            id3Train, _ = holdout_split(id3, holdout=holdout, rng_obj=rng_obj)

            if min(id1Train.size, id2Train.size, id3Train.size) < 1:
                continue

            n_time, _, n_cells = currsig.shape

            cd_context = np.zeros((n_time, n_cells), dtype=np.float32)
            for ind_bin in range(n_time):
                x1 = np.atleast_2d(np.squeeze(currsig[ind_bin, id1Train, :]))
                x2 = np.atleast_2d(np.squeeze(currsig[ind_bin, id2Train, :]))
                cd_context[ind_bin, :] = np.nanmean(x1, axis=0) - np.nanmean(x2, axis=0)

            wc1 = nearest_bin(wc, cfg["Win_context"][0])
            wc2 = nearest_bin(wc, cfg["Win_context"][1])
            if wc2 < wc1:
                wc1, wc2 = wc2, wc1
            coding_direction_context = np.nanmean(cd_context[wc1 : wc2 + 1, :], axis=0)
            coding_direction_context = normalize_vec(coding_direction_context)

            # Absolute session timestamps (up to ~10^4 s): keep float64 so the
            # nearest-frame / in-trial comparisons match Code_M double precision.
            lick_times_all = np.asarray(lentry["lick_time"], dtype=np.float64).reshape(-1)
            video_timestamps = np.asarray(lentry["video_timestamp_continues"], dtype=np.float64).reshape(-1)
            jaw_continuous = np.asarray(lentry["jaw_movement_continues"], dtype=np.float64).reshape(-1)
            lick_mask = np.asarray(lentry["lick_mask"]).reshape(-1).astype(bool)

            if min(video_timestamps.size, jaw_continuous.size, lick_mask.size) == 0:
                continue

            n_cont = min(video_timestamps.size, jaw_continuous.size, lick_mask.size)
            video_timestamps = video_timestamps[:n_cont]
            jaw_continuous = jaw_continuous[:n_cont]
            lick_mask = lick_mask[:n_cont]

            start_times = np.asarray(e["start_time"], dtype=np.float64).reshape(-1)
            spontaneous_licks: List[float] = []

            # Keep only isolated off-trial bouts whose post-lick jaw movement
            # exceeds the pre-lick baseline by the requested quality ratio.
            for lick_t in lick_times_all:
                in_any_trial = np.any(
                    (lick_t >= start_times)
                    & (lick_t <= (start_times + float(cfg["trial_epoch_duration_s"])))
                )
                if in_any_trial:
                    continue

                closest_idx = int(np.argmin(np.abs(video_timestamps - lick_t)))

                bout_start = closest_idx
                while bout_start > 0 and lick_mask[bout_start - 1]:
                    bout_start -= 1

                bout_end = closest_idx
                while bout_end < (lick_mask.size - 1) and lick_mask[bout_end + 1]:
                    bout_end += 1

                bout_duration = video_timestamps[bout_end] - video_timestamps[bout_start]
                if bout_duration <= float(cfg["bout_dur"]):
                    continue

                before_window = [
                    lick_t + float(cfg["Win_lick_before_4qc"][0]),
                    lick_t + float(cfg["Win_lick_before_4qc"][1]),
                ]
                after_window = [
                    lick_t + float(cfg["Win_lick_after_4qc"][0]),
                    lick_t + float(cfg["Win_lick_after_4qc"][1]),
                ]

                idx_before = np.where((video_timestamps >= before_window[0]) & (video_timestamps <= before_window[1]))[0]
                idx_after = np.where((video_timestamps >= after_window[0]) & (video_timestamps <= after_window[1]))[0]

                if idx_before.size == 0 or idx_after.size == 0:
                    continue

                mean_before = np.nanmean(jaw_continuous[idx_before])
                mean_after = np.nanmean(jaw_continuous[idx_after])

                if np.isfinite(mean_before) and np.isfinite(mean_after):
                    if mean_after > float(cfg["lick_quality_ratio"]) * mean_before:
                        spontaneous_licks.append(float(lick_t))

            spontaneous_licks = np.asarray(spontaneous_licks, dtype=np.float32)
            if spontaneous_licks.size < min_spont_licks:
                continue

            selected_unit_idx = np.where(curr_cell_ind)[0]
            unit_spike_times = np.asarray(e["unit_spike_times"], dtype=object).reshape(-1)
            selected_unit_spikes = [unit_spike_times[u] for u in selected_unit_idx]

            n_lick = spontaneous_licks.size
            n_sel = len(selected_unit_spikes)
            after_mean_counts = np.full((n_lick, n_sel), np.nan, dtype=np.float32)
            before_mean_counts = np.full((n_lick, n_sel), np.nan, dtype=np.float32)

            # Summarize each spontaneous lick by one before-vs-after firing-rate
            # estimate per selected unit.
            for ilick, lick_t in enumerate(spontaneous_licks):
                for iunit, spike_times in enumerate(selected_unit_spikes):
                    _, _, c_after = psth_simple_counts_single_event(
                        spike_times,
                        float(lick_t),
                        float(cfg["Win_lick"][0]),
                        float(cfg["Win_lick"][1]),
                        float(cfg["BinSize"]),
                        float(cfg["BinStep"]),
                    )
                    _, _, c_before = psth_simple_counts_single_event(
                        spike_times,
                        float(lick_t),
                        float(cfg["Win_nolick"][0]),
                        float(cfg["Win_nolick"][1]),
                        float(cfg["BinSize"]),
                        float(cfg["BinStep"]),
                    )
                    after_mean_counts[ilick, iunit] = np.nanmean(c_after)
                    before_mean_counts[ilick, iunit] = np.nanmean(c_before)

            coding_direction_lick = np.nanmean(after_mean_counts, axis=0) - np.nanmean(before_mean_counts, axis=0)
            coding_direction_lick = normalize_vec(coding_direction_lick)

            ws1 = nearest_bin(wc, cfg["Win_stim"][0])
            ws2 = nearest_bin(wc, cfg["Win_stim"][1])
            wb1 = nearest_bin(wc, cfg["Win_base"][0])
            wb2 = nearest_bin(wc, cfg["Win_base"][1])
            if ws2 < ws1:
                ws1, ws2 = ws2, ws1
            if wb2 < wb1:
                wb1, wb2 = wb2, wb1

            sum_whisker = np.nansum(currsig[ws1 : ws2 + 1, :, :], axis=0)
            sum_baseline = np.nansum(currsig[wb1 : wb2 + 1, :, :], axis=0)
            coding_direction_stim = np.nanmean(sum_whisker[id3Train, :], axis=0) - np.nanmean(sum_baseline[id3Train, :], axis=0)
            coding_direction_stim = normalize_vec(coding_direction_stim)

            proj_context = project_trials(currsig, coding_direction_context)
            proj_lick = project_trials(currsig, coding_direction_lick)
            proj_stim = project_trials(currsig, coding_direction_stim)

            id_training_context = (
                np.sort(np.concatenate([id1Train, id2Train]))
                if (id1Train.size + id2Train.size)
                else np.array([], dtype=int)
            )
            id_training_stim = np.sort(id3Train) if id3Train.size else np.array([], dtype=int)

            if id_training_context.size:
                proj_context[:, id_training_context] = np.nan
            if id_training_stim.size:
                proj_stim[:, id_training_stim] = np.nan

            coding_direction_matrix[current_area].append(
                {
                    "lickproj": proj_lick,
                    "Contextproj": proj_context,
                    "Stimproj": proj_stim,
                    "index": {
                        "Quiet": qind,
                        "lick": lick,
                        "trial": trial,
                        "completed_trial": completed_trial_ind,
                    },
                }
            )

    return coding_direction_matrix
