"""Frozen MuQ hidden-state extraction and temporal masked pooling."""

import os
from collections.abc import Sequence
from dataclasses import dataclass
from importlib import import_module
from typing import Final, Protocol, override

import torch

from raina_laya.features.audio_features.domain.bins import GlobalBin

from .audio import AudioContext


@dataclass(frozen=True, slots=True)
class MuQBackbone:
    """Pinned identity and hidden-state path for the temporal backbone."""

    repository: str
    revision: str
    inner_path: str
    layer_order: tuple[int, int]


BACKBONE: Final = MuQBackbone(
    repository="OpenMuQ/MuQ-MuLan-large",
    revision="2e01c796b71dca71b45251384c04cd7b237c9020",
    inner_path="model.mulan.audio.model",
    layer_order=(10, 9),
)
EXPECTED_HIDDEN_WIDTH: Final = 1024
HIDDEN_STATE_COUNT: Final = 13
EXPECTED_BATCH_SIZE: Final = 1
HIDDEN_STATE_DIMENSIONS: Final = 3
ALLOWED_EXTRACTION_DEVICES: Final = frozenset({"0", "1", "2", "3"})


class HiddenStateOutput(Protocol):
    """Structural output contract of the pinned inner MuQ model."""

    @property
    def hidden_states(self) -> Sequence[torch.Tensor]:
        """Return every hidden layer in model order."""
        ...


class InnerAudioModel(Protocol):
    """Callable hidden-state source at ``model.mulan.audio.model``."""

    def __call__(
        self,
        waveform: torch.Tensor,
        *,
        output_hidden_states: bool,
    ) -> HiddenStateOutput:
        """Return temporal hidden states for one fixed audio context."""
        ...


@dataclass(frozen=True, slots=True)
class LoadedBackbone:
    """Loaded root module paired with its exact pinned inner model."""

    root: torch.nn.Module
    inner_model: InnerAudioModel


class MuQLoader(Protocol):
    """Loader contract preserving the pinned revision keyword."""

    def __call__(
        self,
        repository: str,
        *,
        revision: str,
        local_files_only: bool = False,
    ) -> LoadedBackbone:
        """Load the root and select ``model.mulan.audio.model``."""
        ...


@dataclass(frozen=True, slots=True)
class _InnerModelAdapter:
    model: torch.nn.Module

    def __call__(
        self,
        waveform: torch.Tensor,
        *,
        output_hidden_states: bool,
    ) -> HiddenStateOutput:
        return self.model(
            waveform,
            output_hidden_states=output_hidden_states,
        )


@dataclass(frozen=True, slots=True)
class PooledWindow:
    """Pooled tokens plus the generated frame alignment evidence."""

    features: torch.Tensor
    frame_timestamps: torch.Tensor
    valid_frame_mask: torch.Tensor


@dataclass(frozen=True, slots=True)
class FrameWindow:
    """Selected FP32 layers and their original context-frame alignment."""

    layer_10: torch.Tensor
    layer_9: torch.Tensor
    frame_timestamps: torch.Tensor
    valid_frame_mask: torch.Tensor


@dataclass(frozen=True, slots=True)
class LayerFrames:
    """Every FP32 hidden state of one context and its original frame alignment."""

    layers: tuple[torch.Tensor, ...]
    frame_timestamps: torch.Tensor
    valid_frame_mask: torch.Tensor


@dataclass(frozen=True, slots=True)
class FeatureExtractionError(Exception):
    """Base error for temporal feature extraction contract failures."""


@dataclass(frozen=True, slots=True)
class CudaExtractionRequiredError(FeatureExtractionError):
    """Real MuQ extraction requires one approved physical CUDA device."""

    reason: str

    @override
    def __str__(self) -> str:
        """Describe the failed real-extraction precondition."""
        return (
            "real MuQ extraction requires CUDA_VISIBLE_DEVICES "
            f"in {sorted(ALLOWED_EXTRACTION_DEVICES)}: {self.reason}"
        )


@dataclass(frozen=True, slots=True)
class HiddenStateContractError(FeatureExtractionError):
    """Inner MuQ hidden states do not match the pinned temporal contract."""

    reason: str

    @override
    def __str__(self) -> str:
        """Describe the malformed hidden-state result."""
        return f"invalid MuQ hidden states: {self.reason}"


@dataclass(frozen=True, slots=True)
class EmptyPoolingBinError(FeatureExtractionError):
    """A post-merge global bin contains no valid hidden-state frame."""

    start: float
    end: float

    @override
    def __str__(self) -> str:
        """Describe the empty pooling interval."""
        return f"global bin [{self.start}, {self.end}) has no valid MuQ frame"


def _selected_hidden_states(
    hidden_states: Sequence[torch.Tensor],
    expected_width: int,
    expected_batch_size: int = EXPECTED_BATCH_SIZE,
) -> tuple[torch.Tensor, torch.Tensor]:
    if len(hidden_states) <= max(BACKBONE.layer_order):
        raise HiddenStateContractError(reason="layers 10 and 9 must both exist")
    layer_10 = hidden_states[BACKBONE.layer_order[0]]
    layer_9 = hidden_states[BACKBONE.layer_order[1]]
    if (
        layer_10.ndim != HIDDEN_STATE_DIMENSIONS
        or layer_9.ndim != HIDDEN_STATE_DIMENSIONS
    ):
        raise HiddenStateContractError(reason="selected layers must have shape [1,F,D]")
    expected_shape = (expected_batch_size, layer_10.shape[1], expected_width)
    if (
        tuple(layer_10.shape) != expected_shape
        or tuple(layer_9.shape) != expected_shape
    ):
        raise HiddenStateContractError(
            reason=(
                "layers 10 and 9 must align as "
                f"[{expected_batch_size},F,{expected_width}]; got "
                f"{tuple(layer_10.shape)} and "
                f"{tuple(layer_9.shape)}"
            ),
        )
    return layer_10.to(dtype=torch.float32), layer_9.to(dtype=torch.float32)


def pool_hidden_states(
    hidden_states: Sequence[torch.Tensor],
    *,
    frame_timestamps: torch.Tensor,
    valid_frame_mask: torch.Tensor,
    bins: Sequence[GlobalBin],
    expected_width: int = EXPECTED_HIDDEN_WIDTH,
) -> torch.Tensor:
    """Masked-mean L10 and L9 separately within each global bin."""
    layer_10, layer_9 = _selected_hidden_states(hidden_states, expected_width)
    return _pool_layers(
        layer_10,
        layer_9,
        frame_timestamps=frame_timestamps,
        valid_frame_mask=valid_frame_mask,
        bins=bins,
    )


def _pool_layers(
    layer_10: torch.Tensor,
    layer_9: torch.Tensor,
    *,
    frame_timestamps: torch.Tensor,
    valid_frame_mask: torch.Tensor,
    bins: Sequence[GlobalBin],
) -> torch.Tensor:
    """Pool selected layers after the caller has finalized global bins."""
    frame_count = layer_10.shape[1]
    if (
        frame_timestamps.ndim != 1
        or valid_frame_mask.ndim != 1
        or frame_timestamps.shape[0] != frame_count
        or valid_frame_mask.shape[0] != frame_count
    ):
        raise HiddenStateContractError(
            reason="timestamps and validity must align with the model frame count",
        )
    if valid_frame_mask.dtype is not torch.bool:
        raise HiddenStateContractError(reason="frame validity mask must be boolean")

    pooled: list[torch.Tensor] = []
    for interval in bins:
        selected = (
            valid_frame_mask
            & (frame_timestamps >= interval.start)
            & (frame_timestamps < interval.end)
        )
        if not torch.any(selected):
            raise EmptyPoolingBinError(start=interval.start, end=interval.end)
        pooled.append(
            torch.cat(
                (
                    layer_10[0, selected].mean(dim=0),
                    layer_9[0, selected].mean(dim=0),
                ),
            ),
        )
    return torch.stack(pooled).to(dtype=torch.float32)


def extract_inner_model_frames(
    model: InnerAudioModel,
    context: AudioContext,
    *,
    window_start: float,
    expected_width: int = EXPECTED_HIDDEN_WIDTH,
) -> FrameWindow:
    """Run one context once and retain the selected aligned frames for pooling."""
    with torch.autocast(device_type=context.waveform.device.type, enabled=False):
        output = model(
            context.waveform.unsqueeze(0).to(dtype=torch.float32),
            output_hidden_states=True,
        )
    layer_10, layer_9 = _selected_hidden_states(output.hidden_states, expected_width)
    return _aligned_frames(layer_10, layer_9, context, window_start)


def extract_inner_model_frame_batch(
    model: InnerAudioModel,
    contexts: Sequence[AudioContext],
    *,
    window_starts: Sequence[float],
    expected_width: int = EXPECTED_HIDDEN_WIDTH,
) -> tuple[FrameWindow, ...]:
    """Batch fixed contexts, then restore the unchanged per-window alignment."""
    if not contexts or len(contexts) != len(window_starts):
        raise HiddenStateContractError(reason="contexts and starts must align")
    first = contexts[0]
    if any(
        context.waveform.shape != first.waveform.shape
        or context.sample_rate != first.sample_rate
        or context.waveform.device != first.waveform.device
        for context in contexts
    ):
        raise HiddenStateContractError(reason="batched contexts must be compatible")
    with torch.autocast(device_type=first.waveform.device.type, enabled=False):
        output = model(
            torch.stack(tuple(context.waveform for context in contexts)).to(
                dtype=torch.float32,
            ),
            output_hidden_states=True,
        )
    layer_10, layer_9 = _selected_hidden_states(
        output.hidden_states,
        expected_width,
        len(contexts),
    )
    return tuple(
        _aligned_frames(
            layer_10[index : index + 1],
            layer_9[index : index + 1],
            context,
            start,
        )
        for index, (context, start) in enumerate(
            zip(contexts, window_starts, strict=True),
        )
    )


def _frame_alignment(
    frame_count: int,
    device: torch.device,
    context: AudioContext,
    window_start: float,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Return (global frame timestamps, valid-frame mask) of one model context."""
    if frame_count <= 0:
        raise HiddenStateContractError(reason="selected layers contain no frames")
    context_seconds = context.waveform.numel() / context.sample_rate
    frame_width = context_seconds / frame_count
    local_centers = (
        torch.arange(frame_count, dtype=torch.float32, device=device) + 0.5
    ) * frame_width
    valid_seconds = context.valid_samples / context.sample_rate
    valid_frame_mask = local_centers < valid_seconds
    frame_timestamps = local_centers + window_start
    return frame_timestamps, valid_frame_mask


def _aligned_frames(
    layer_10: torch.Tensor,
    layer_9: torch.Tensor,
    context: AudioContext,
    window_start: float,
) -> FrameWindow:
    frame_timestamps, valid_frame_mask = _frame_alignment(
        layer_10.shape[1], layer_10.device, context, window_start
    )
    return FrameWindow(
        layer_10=layer_10,
        layer_9=layer_9,
        frame_timestamps=frame_timestamps,
        valid_frame_mask=valid_frame_mask,
    )


def pool_frame_window(
    frames: FrameWindow,
    bins: Sequence[GlobalBin],
) -> torch.Tensor:
    """Pool retained frames under the merged final bin ownership."""
    return _pool_layers(
        frames.layer_10,
        frames.layer_9,
        frame_timestamps=frames.frame_timestamps,
        valid_frame_mask=frames.valid_frame_mask,
        bins=bins,
    )


def run_inner_model(
    model: InnerAudioModel,
    context: AudioContext,
    *,
    window_start: float,
    bins: Sequence[GlobalBin],
    expected_width: int = EXPECTED_HIDDEN_WIDTH,
) -> PooledWindow:
    """Run one prepared context and derive timestamps from its actual frame count."""
    frames = extract_inner_model_frames(
        model,
        context,
        window_start=window_start,
        expected_width=expected_width,
    )
    features = pool_frame_window(frames, bins)
    return PooledWindow(
        features=features,
        frame_timestamps=frames.frame_timestamps,
        valid_frame_mask=frames.valid_frame_mask,
    )


def _all_hidden_states(
    hidden_states: Sequence[torch.Tensor],
    expected_width: int,
) -> tuple[torch.Tensor, ...]:
    if len(hidden_states) != HIDDEN_STATE_COUNT:
        raise HiddenStateContractError(
            reason=(
                f"expected {HIDDEN_STATE_COUNT} hidden states; got {len(hidden_states)}"
            ),
        )
    if any(state.ndim != HIDDEN_STATE_DIMENSIONS for state in hidden_states):
        raise HiddenStateContractError(reason="hidden states must have shape [1,F,D]")
    expected_shape = (EXPECTED_BATCH_SIZE, hidden_states[0].shape[1], expected_width)
    if any(tuple(state.shape) != expected_shape for state in hidden_states):
        raise HiddenStateContractError(
            reason=f"hidden states must all align as {list(expected_shape)}",
        )
    return tuple(state.to(dtype=torch.float32) for state in hidden_states)


def extract_inner_model_layer_frames(
    model: InnerAudioModel,
    context: AudioContext,
    *,
    window_start: float,
    expected_width: int = EXPECTED_HIDDEN_WIDTH,
) -> LayerFrames:
    """Run one context once and retain every hidden state with its alignment."""
    with torch.autocast(device_type=context.waveform.device.type, enabled=False):
        output = model(
            context.waveform.unsqueeze(0).to(dtype=torch.float32),
            output_hidden_states=True,
        )
    layers = _all_hidden_states(output.hidden_states, expected_width)
    frame_timestamps, valid_frame_mask = _frame_alignment(
        layers[0].shape[1], layers[0].device, context, window_start
    )
    return LayerFrames(
        layers=layers,
        frame_timestamps=frame_timestamps,
        valid_frame_mask=valid_frame_mask,
    )


def pool_layer_frames(
    frames: LayerFrames,
    bins: Sequence[GlobalBin],
) -> torch.Tensor:
    """Masked-mean every hidden state within each bin as FP32 [bins, 13, 1024].

    Each layer is the mean over the same gathered rows as `_pool_layers`, so layer 10
    and layer 9 equal the two halves of the legacy joint token bit for bit. The row
    indices are computed once per bin to avoid one host synchronization per layer.
    """
    frame_count = frames.layers[0].shape[1]
    if (
        frames.frame_timestamps.ndim != 1
        or frames.valid_frame_mask.ndim != 1
        or frames.frame_timestamps.shape[0] != frame_count
        or frames.valid_frame_mask.shape[0] != frame_count
    ):
        raise HiddenStateContractError(
            reason="timestamps and validity must align with the model frame count",
        )
    if frames.valid_frame_mask.dtype is not torch.bool:
        raise HiddenStateContractError(reason="frame validity mask must be boolean")

    pooled: list[torch.Tensor] = []
    for interval in bins:
        selected = (
            frames.valid_frame_mask
            & (frames.frame_timestamps >= interval.start)
            & (frames.frame_timestamps < interval.end)
        )
        rows = selected.nonzero().squeeze(1)
        if rows.numel() == 0:
            raise EmptyPoolingBinError(start=interval.start, end=interval.end)
        pooled.append(
            torch.stack(
                [layer[0].index_select(0, rows).mean(dim=0) for layer in frames.layers],
            ),
        )
    return torch.stack(pooled).to(dtype=torch.float32)


def _default_loader(
    repository: str,
    *,
    revision: str,
    local_files_only: bool = False,
) -> LoadedBackbone:
    module = import_module("muq")
    root = module.MuQMuLan.from_pretrained(
        repository,
        revision=revision,
        local_files_only=local_files_only,
    )
    inner_model = root.mulan.audio.model
    if not isinstance(inner_model, torch.nn.Module):
        raise HiddenStateContractError(
            reason="model.mulan.audio.model must be a torch module",
        )
    return LoadedBackbone(
        root=root,
        inner_model=_InnerModelAdapter(model=inner_model),
    )


def load_real_inner_model(
    loader: MuQLoader = _default_loader,
    *,
    local_files_only: bool = False,
) -> InnerAudioModel:
    """Load and freeze the pinned inner model on one of four CUDA devices."""
    if not torch.cuda.is_available():
        raise CudaExtractionRequiredError(reason="CUDA is unavailable")
    visible_devices = os.environ.get("CUDA_VISIBLE_DEVICES")
    if visible_devices not in ALLOWED_EXTRACTION_DEVICES:
        raise CudaExtractionRequiredError(
            reason=f"CUDA_VISIBLE_DEVICES is {visible_devices!r}",
        )

    loaded = loader(
        BACKBONE.repository,
        revision=BACKBONE.revision,
        local_files_only=local_files_only,
    )
    target = (
        loaded.inner_model.model
        if isinstance(loaded.inner_model, _InnerModelAdapter)
        else loaded.root
    )
    # Injected legacy test loaders may expose a callable rather than a module.
    # The real pinned adapter retains only its exact inner audio module.
    _ = target.eval()
    _ = target.to(device="cuda", dtype=torch.float32)
    for parameter in target.parameters():
        parameter.requires_grad = False
    return loaded.inner_model
