"""Synthetic-safe mono resampling and fixed-context audio padding."""

import inspect
from collections import OrderedDict
from dataclasses import dataclass
from math import gcd, isfinite
from threading import get_ident
from typing import Final, override

import torch
import torchaudio
from torchaudio.functional import functional as audio_functional
from torchaudio.functional import resample

__all__ = [
    "AudioContext",
    "AudioDurationError",
    "AudioPreprocessPolicy",
    "AudioShapeError",
    "CudaResamplerContext",
    "InvalidSampleRateError",
    "ResamplerCompatibilityError",
    "prepare_audio_context",
]


@dataclass(frozen=True, slots=True)
class AudioPreprocessPolicy:
    """Target sample rate and fixed model-context duration."""

    sample_rate: int = 24_000
    context_seconds: float = 10.0


DEFAULT_POLICY: Final = AudioPreprocessPolicy()
MONO_DIMENSIONS: Final = 1
CHANNELS_FIRST_DIMENSIONS: Final = 2


@dataclass(frozen=True, slots=True)
class AudioContext:
    """A mono fixed-length context and its actual-sample mask."""

    waveform: torch.Tensor
    valid_sample_mask: torch.Tensor
    sample_rate: int
    valid_samples: int


@dataclass(frozen=True, slots=True)
class AudioContractError(Exception):
    """Base error for invalid decoded audio inputs."""


@dataclass(frozen=True, slots=True)
class ResamplerCompatibilityError(AudioContractError):
    """The scoped resampler cannot preserve the pinned CUDA recipe."""

    reason: str

    @override
    def __str__(self) -> str:
        return f"incompatible CUDA resampler: {self.reason}"


class CudaResamplerContext:
    """Own bounded CUDA-created FP32 kernels for one extraction invocation.

    Torchaudio 2.11's public Resample constructor creates kernels on CPU even
    inside a CUDA device context. This minimal adapter uses the same two pinned
    kernel operations as functional.resample, with its actual device and dtype.
    It never transfers a CPU-created kernel or changes the filter parameters.
    """

    def __init__(
        self, *, max_entries: int = 8, max_kernel_bytes: int = 16 * 1024 * 1024
    ) -> None:
        """Validate the pinned adapter before retaining any CUDA kernel."""
        if max_entries < 1 or max_kernel_bytes < 1:
            raise ResamplerCompatibilityError(reason="kernel bounds must be positive")
        if torchaudio.__version__.split("+", maxsplit=1)[0] != "2.11.0":
            raise ResamplerCompatibilityError(reason="Torchaudio 2.11.0 is required")
        self._get_kernel = audio_functional._get_sinc_resample_kernel  # noqa: SLF001 - pinned adapter matching functional.resample
        self._apply_kernel = audio_functional._apply_sinc_resample_kernel  # noqa: SLF001 - pinned adapter matching functional.resample
        expected = (
            "orig_freq",
            "new_freq",
            "gcd",
            "lowpass_filter_width",
            "rolloff",
            "resampling_method",
            "beta",
            "device",
            "dtype",
        )
        apply_expected = ("waveform", "orig_freq", "new_freq", "gcd", "kernel", "width")
        if (
            tuple(inspect.signature(self._get_kernel).parameters) != expected
            or tuple(inspect.signature(self._apply_kernel).parameters) != apply_expected
        ):
            raise ResamplerCompatibilityError(reason="kernel signature changed")
        self.max_entries = max_entries
        self.max_kernel_bytes = max_kernel_bytes
        self._kernels: OrderedDict[
            tuple[int, int, torch.device, torch.dtype, int, float, str, None],
            tuple[torch.Tensor, int],
        ] = OrderedDict()
        self._owner = get_ident()
        self._closed = False
        self.hits = 0
        self.misses = 0

    @property
    def cached_bytes(self) -> int:
        """Return bytes retained by this invocation's kernel cache."""
        return sum(
            kernel.numel() * kernel.element_size()
            for kernel, _ in self._kernels.values()
        )

    def close(self) -> None:
        """Drop only this context's kernel references, without global CUDA actions."""
        self._kernels.clear()
        self._closed = True

    def resample(
        self, waveform: torch.Tensor, source_rate: int, target_rate: int
    ) -> torch.Tensor:
        """Apply the original FP32 CUDA filter, reusing only same-recipe kernels."""
        if self._closed or get_ident() != self._owner:
            raise ResamplerCompatibilityError(
                reason="context is closed or has another owner"
            )
        if waveform.device.type != "cuda" or waveform.dtype is not torch.float32:
            raise ResamplerCompatibilityError(
                reason="kernels require CUDA float32 inputs"
            )
        if source_rate <= 0 or target_rate <= 0:
            reason = "Original frequency and desired frequecy should be positive"
            raise ValueError(reason)
        if source_rate == target_rate:
            return waveform
        key = (
            source_rate,
            target_rate,
            waveform.device,
            waveform.dtype,
            6,
            0.99,
            "sinc_interp_hann",
            None,
        )
        cached = self._kernels.pop(key, None)
        if cached is None:
            self.misses += 1
            divisor = gcd(source_rate, target_rate)
            kernel, width = self._get_kernel(
                source_rate,
                target_rate,
                divisor,
                lowpass_filter_width=6,
                rolloff=0.99,
                resampling_method="sinc_interp_hann",
                beta=None,
                device=waveform.device,
                dtype=waveform.dtype,
            )
            if kernel.device != waveform.device or kernel.dtype != waveform.dtype:
                raise ResamplerCompatibilityError(
                    reason="kernel device or precision changed"
                )
            size = kernel.numel() * kernel.element_size()
            if size <= self.max_kernel_bytes:
                while self._kernels and (
                    len(self._kernels) >= self.max_entries
                    or self.cached_bytes + size > self.max_kernel_bytes
                ):
                    self._kernels.popitem(last=False)
                self._kernels[key] = (kernel, width)
        else:
            self.hits += 1
            kernel, width = cached
            self._kernels[key] = cached
        return self._apply_kernel(
            waveform,
            source_rate,
            target_rate,
            gcd(source_rate, target_rate),
            kernel,
            width,
        )


@dataclass(frozen=True, slots=True)
class InvalidSampleRateError(AudioContractError):
    """A source or target sample rate is invalid."""

    sample_rate: int

    @override
    def __str__(self) -> str:
        """Describe the rejected sample rate."""
        return f"sample rate must be positive; received {self.sample_rate}"


@dataclass(frozen=True, slots=True)
class AudioShapeError(AudioContractError):
    """A decoded tensor is empty or not mono/channels-first audio."""

    shape: tuple[int, ...]

    @override
    def __str__(self) -> str:
        """Describe the rejected decoded tensor shape."""
        return (
            f"audio must have shape [samples] or [channels, samples]; got {self.shape}"
        )


@dataclass(frozen=True, slots=True)
class AudioDurationError(AudioContractError):
    """A decoded context is empty or longer than the model context."""

    duration_seconds: float
    maximum_seconds: float

    @override
    def __str__(self) -> str:
        """Describe the rejected context duration."""
        return (
            "audio context duration must be finite, positive, and no longer than "
            f"{self.maximum_seconds} seconds; received {self.duration_seconds}"
        )


def _validate_policy(policy: AudioPreprocessPolicy) -> AudioPreprocessPolicy:
    if policy.sample_rate <= 0:
        raise InvalidSampleRateError(sample_rate=policy.sample_rate)
    if not isfinite(policy.context_seconds) or policy.context_seconds <= 0:
        raise AudioDurationError(
            duration_seconds=policy.context_seconds,
            maximum_seconds=policy.context_seconds,
        )
    return policy


def _mono_waveform(waveform: torch.Tensor) -> torch.Tensor:
    if waveform.ndim not in (MONO_DIMENSIONS, CHANNELS_FIRST_DIMENSIONS):
        raise AudioShapeError(shape=tuple(waveform.shape))
    if waveform.ndim == CHANNELS_FIRST_DIMENSIONS and waveform.shape[0] == 0:
        raise AudioShapeError(shape=tuple(waveform.shape))
    mono = waveform if waveform.ndim == MONO_DIMENSIONS else waveform.mean(dim=0)
    if mono.numel() == 0:
        raise AudioShapeError(shape=tuple(waveform.shape))
    return mono.to(dtype=torch.float32)


def prepare_audio_context(
    waveform: torch.Tensor,
    source_sample_rate: int,
    policy: AudioPreprocessPolicy = DEFAULT_POLICY,
    *,
    resampler: CudaResamplerContext | None = None,
) -> AudioContext:
    """Convert decoded channels-first audio to mono, resample, and zero-pad."""
    if source_sample_rate <= 0:
        raise InvalidSampleRateError(sample_rate=source_sample_rate)
    validated_policy = _validate_policy(policy)
    mono = _mono_waveform(waveform)
    duration_seconds = mono.numel() / source_sample_rate
    if (
        not isfinite(duration_seconds)
        or duration_seconds <= 0
        or duration_seconds > validated_policy.context_seconds
    ):
        raise AudioDurationError(
            duration_seconds=duration_seconds,
            maximum_seconds=validated_policy.context_seconds,
        )

    normalized = (
        mono
        if source_sample_rate == validated_policy.sample_rate
        else resampler.resample(mono, source_sample_rate, validated_policy.sample_rate)
        if resampler is not None
        else resample(
            mono,
            orig_freq=source_sample_rate,
            new_freq=validated_policy.sample_rate,
        )
    )
    target_samples = round(
        validated_policy.context_seconds * validated_policy.sample_rate,
    )
    valid_samples = normalized.numel()
    if valid_samples > target_samples:
        raise AudioDurationError(
            duration_seconds=valid_samples / validated_policy.sample_rate,
            maximum_seconds=validated_policy.context_seconds,
        )
    padding_samples = target_samples - valid_samples
    padded = torch.cat(
        (
            normalized,
            torch.zeros(
                padding_samples,
                dtype=normalized.dtype,
                device=normalized.device,
            ),
        ),
    )
    valid_sample_mask = torch.cat(
        (
            torch.ones(valid_samples, dtype=torch.bool, device=normalized.device),
            torch.zeros(
                padding_samples,
                dtype=torch.bool,
                device=normalized.device,
            ),
        ),
    )
    return AudioContext(
        waveform=padded,
        valid_sample_mask=valid_sample_mask,
        sample_rate=validated_policy.sample_rate,
        valid_samples=valid_samples,
    )
