"""CUDA-only lifetime and full-song prediction orchestration."""

import hashlib
import io
import os
from collections.abc import Sequence
from pathlib import Path
from typing import Final, Protocol

import soundfile
import torch
from huggingface_hub import constants as hf_constants
from huggingface_hub import snapshot_download
from huggingface_hub.errors import LocalEntryNotFoundError
from torch.nn import functional

from raina_laya.features.audio_features.client import (
    BACKBONE,
    CacheBuildItem,
    CachePolicy,
    CudaExtractionRequiredError,
    FeatureExtractionError,
    InnerAudioModel,
    InvalidDurationError,
    NoValidFrameBinError,
    extract_feature_record,
    extract_layer_records,
    load_real_inner_model,
)
from raina_laya.features.grade_model.client import (
    RejectBand,
    TemporalGradeInputError,
    hierarchical_gate_margin,
    hierarchical_log_probs,
)
from raina_laya.features.serving.application.batching import (
    BatchSettings,
    PredictionError,
    PreparedSong,
    head_batches,
    prepare_songs,
)
from raina_laya.features.serving.domain.contracts import (
    BackendLoadError,
    Grade,
    Prediction,
    PredictionExecutionError,
    PredictionInputError,
)
from raina_laya.features.serving.domain.decision import decide
from raina_laya.features.serving.infrastructure.checkpoint import (
    LoadedCheckpoint,
    LoadedMember,
)
from raina_laya.features.serving.infrastructure.checkpoint import (
    load_checkpoint as load_head_checkpoint,
)
from raina_laya.features.serving.infrastructure.package import (
    LoadedPackage,
    load_package,
)

GRADES: Final[tuple[Grade, Grade, Grade]] = ("Fail", "A", "S")
PHYSICAL_CUDA_DEVICE: Final = "1"
LOGICAL_CUDA_DEVICE: Final = "cuda:0"
_CUDA_INIT_ERRORS: Final = (
    AttributeError,
    CudaExtractionRequiredError,
    ImportError,
    OSError,
    RuntimeError,
)


class GradeHead(Protocol):
    """Callable three-grade head used by the serving runtime."""

    def __call__(
        self,
        features: torch.Tensor,
        center_times: torch.Tensor,
        valid_ratios: torch.Tensor,
        padding_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Return one Fail/A/S logit row."""
        ...


class ServingBackend:
    """Own one loaded MuQ/head pair until explicit shutdown."""

    def __init__(
        self,
        *,
        inner_model: InnerAudioModel,
        head: GradeHead,
        device: torch.device,
        checkpoint_sha256: str,
        model_family: str,
    ) -> None:
        """Bind already-loaded components to one inference device."""
        self._inner_model: InnerAudioModel | None = inner_model
        self._head: GradeHead | None = head
        self._device = device
        self.checkpoint_sha256 = checkpoint_sha256
        self.model_family = model_family

    def predict(self, payload: bytes, filename: str) -> Prediction:
        """Decode and classify every temporal bin of one complete song."""
        inner_model = self._inner_model
        head = self._head
        if inner_model is None or head is None:
            raise PredictionExecutionError(detail="backend is closed")
        waveform, sample_rate = _decode_audio(payload)
        input_sha256 = hashlib.sha256(payload).hexdigest()
        duration_seconds = waveform.shape[1] / sample_rate
        item = CacheBuildItem(
            source_identifier=filename,
            source_sha256=input_sha256,
            waveform=waveform.to(device=self._device, dtype=torch.float32),
            sample_rate=sample_rate,
        )
        try:
            with torch.inference_mode():
                record = extract_feature_record(item, inner_model, CachePolicy())
                features = record.features.unsqueeze(0).to(
                    device=self._device,
                    dtype=torch.float32,
                )
                centers = torch.tensor(
                    tuple(interval.center for interval in record.bins),
                    device=self._device,
                    dtype=torch.float32,
                ).unsqueeze(0)
                valid_ratios = torch.tensor(
                    tuple(interval.valid_ratio for interval in record.bins),
                    device=self._device,
                    dtype=torch.float32,
                ).unsqueeze(0)
                padding_mask = record.padding_mask.unsqueeze(0).to(self._device)
                with torch.autocast(
                    device_type=self._device.type,
                    dtype=torch.bfloat16,
                ):
                    logits = head(
                        features,
                        centers,
                        valid_ratios,
                        padding_mask,
                    )
        except (InvalidDurationError, NoValidFrameBinError) as error:
            raise PredictionInputError(detail=str(error)) from error
        except (FeatureExtractionError, TemporalGradeInputError, RuntimeError) as error:
            raise PredictionExecutionError(detail=str(error)) from error
        return _prediction(
            logits,
            input_sha256=input_sha256,
            checkpoint_sha256=self.checkpoint_sha256,
            duration_seconds=duration_seconds,
        )

    def close(self) -> None:
        """Release this runtime's references without touching other CUDA work."""
        self._head = None
        self._inner_model = None

    def predict_batch(
        self,
        requests: Sequence[tuple[bytes, str]],
        *,
        window_batch_size: int = 2,
        decode_workers: int = 2,
    ) -> tuple[Prediction | PredictionError, ...]:
        """Return ordered decisions or explicit per-song errors, never error grades."""
        if self._inner_model is None or self._head is None:
            raise PredictionExecutionError(detail="backend is closed")
        prepared = prepare_songs(
            requests,
            decode=_decode_audio,
            model=self._inner_model,
            device=self._device,
            settings=BatchSettings(window_batch_size, decode_workers),
        )
        outputs = head_batches(prepared, head=self._head, device=self._device)
        results: list[Prediction | PredictionError] = []
        for index, song in enumerate(prepared):
            if not isinstance(song, PreparedSong):
                results.append(song)
                continue
            logits = outputs[index]
            if isinstance(logits, PredictionExecutionError):
                results.append(logits)
                continue
            try:
                results.append(
                    _prediction(
                        logits,
                        input_sha256=song.input_sha256,
                        checkpoint_sha256=self.checkpoint_sha256,
                        duration_seconds=song.duration_seconds,
                    ),
                )
            except PredictionExecutionError as error:
                results.append(error)
        return tuple(results)


class HierarchicalBackend:
    """Own one v5 MuQ model and its gate members until explicit shutdown.

    One member is a single checkpoint; several members score the mean gate margin.
    Each member reads its own MuQ layers from one pass over the song.
    """

    def __init__(  # noqa: PLR0913 - one loaded package bound to one device
        self,
        *,
        inner_model: InnerAudioModel,
        members: Sequence[LoadedMember],
        band: RejectBand | None,
        device: torch.device,
        checkpoint_sha256: str,
        model_family: str,
    ) -> None:
        """Bind already-loaded components to one inference device."""
        self._inner_model: InnerAudioModel | None = inner_model
        self._members: tuple[LoadedMember, ...] | None = tuple(members)
        self._band = band
        self._device = device
        self.checkpoint_sha256 = checkpoint_sha256
        self.model_family = model_family

    def predict(self, payload: bytes, filename: str) -> Prediction:
        """Decode one complete song and decide it from the mean gate margin."""
        inner_model = self._inner_model
        members = self._members
        if inner_model is None or members is None:
            raise PredictionExecutionError(detail="backend is closed")
        waveform, sample_rate = _decode_audio(payload)
        input_sha256 = hashlib.sha256(payload).hexdigest()
        duration_seconds = waveform.shape[1] / sample_rate
        item = CacheBuildItem(
            source_identifier=filename,
            source_sha256=input_sha256,
            waveform=waveform.to(device=self._device, dtype=torch.float32),
            sample_rate=sample_rate,
        )
        try:
            with torch.inference_mode():
                records = extract_layer_records(item, inner_model, CachePolicy())
                bins = records[0].bins
                centers = torch.tensor(
                    tuple(interval.center for interval in bins),
                    device=self._device,
                    dtype=torch.float32,
                ).unsqueeze(0)
                valid_ratios = torch.tensor(
                    tuple(interval.valid_ratio for interval in bins),
                    device=self._device,
                    dtype=torch.float32,
                ).unsqueeze(0)
                padding_mask = records[0].padding_mask.unsqueeze(0).to(self._device)
                rows: list[torch.Tensor] = []
                for member in members:
                    features = torch.cat(
                        [records[layer].features for layer in member.layers], dim=-1
                    )
                    with torch.autocast(
                        device_type=self._device.type,
                        dtype=torch.bfloat16,
                    ):
                        logits = member.head(
                            features.unsqueeze(0).to(
                                device=self._device,
                                dtype=torch.float32,
                            ),
                            centers,
                            valid_ratios,
                            padding_mask,
                        )
                    rows.append(logits.detach().to(device="cpu", dtype=torch.float32))
        except (InvalidDurationError, NoValidFrameBinError) as error:
            raise PredictionInputError(detail=str(error)) from error
        except (FeatureExtractionError, TemporalGradeInputError, RuntimeError) as error:
            raise PredictionExecutionError(detail=str(error)) from error
        return _hierarchical_prediction(
            rows,
            self._band,
            input_sha256=input_sha256,
            checkpoint_sha256=self.checkpoint_sha256,
            duration_seconds=duration_seconds,
        )

    def close(self) -> None:
        """Release this runtime's references without touching other CUDA work."""
        self._members = None
        self._inner_model = None


def _decode_audio(payload: bytes) -> tuple[torch.Tensor, int]:
    if not payload:
        raise PredictionInputError(detail="audio payload is empty")
    try:
        waveform, sample_rate = soundfile.read(
            io.BytesIO(payload),
            always_2d=True,
            dtype="float32",
        )
    except soundfile.LibsndfileError as error:
        raise PredictionInputError(detail=f"cannot decode {error}") from error
    decoded = torch.from_numpy(waveform.T.copy()).to(dtype=torch.float32)
    if decoded.numel() == 0:
        raise PredictionInputError(detail="decoded audio contains no samples")
    if not torch.isfinite(decoded).all():
        raise PredictionInputError(detail="decoded audio must contain finite samples")
    if sample_rate <= 0:
        raise PredictionInputError(detail="decoded sample rate must be positive")
    return decoded, int(sample_rate)


def _prediction(
    logits: torch.Tensor,
    *,
    input_sha256: str,
    checkpoint_sha256: str,
    duration_seconds: float,
) -> Prediction:
    values = logits.detach().to(device="cpu", dtype=torch.float32)
    if values.shape != (1, len(GRADES)) or not torch.isfinite(values).all():
        raise PredictionExecutionError(
            detail="head must return one finite Fail/A/S logit row",
        )
    log_scores = functional.logsigmoid(values)
    probabilities = functional.softmax(log_scores, dim=-1)
    grade_index = int(values.argmax(dim=-1).item())
    grade = GRADES[grade_index]
    return Prediction(
        input_sha256=input_sha256,
        checkpoint_sha256=checkpoint_sha256,
        grade=grade,
        passed=grade != "Fail",
        logits={
            grade_name: float(values[0, index].item())
            for index, grade_name in enumerate(GRADES)
        },
        grade_probabilities={
            grade_name: float(probabilities[0, index].item())
            for index, grade_name in enumerate(GRADES)
        },
        duration_seconds=duration_seconds,
    )


def _hierarchical_prediction(
    rows: Sequence[torch.Tensor],
    band: RejectBand | None,
    *,
    input_sha256: str,
    checkpoint_sha256: str,
    duration_seconds: float,
) -> Prediction:
    """Average member gate margins and conditional logits, then apply the band."""
    if any(
        row.shape != (1, len(GRADES)) or not torch.isfinite(row).all() for row in rows
    ):
        raise PredictionExecutionError(
            detail="head must return one finite gate and conditional logit row",
        )
    # Same arithmetic as the research ensemble: [-m/2, +m/2, c] from the member means.
    margin = torch.stack([hierarchical_gate_margin(row) for row in rows]).mean(dim=0)
    conditional = torch.stack([row[:, 2] for row in rows]).mean(dim=0)
    log_probabilities = hierarchical_log_probs(
        torch.stack((-margin / 2, margin / 2, conditional), dim=1)
    )
    gate_margin = float(margin.item())
    decision, grade = decide(gate_margin, float(conditional.item()), band)
    return Prediction(
        input_sha256=input_sha256,
        checkpoint_sha256=checkpoint_sha256,
        grade=grade,
        passed=None if decision == "reject" else decision == "pass",
        logits={
            grade_name: float(log_probabilities[0, index].item())
            for index, grade_name in enumerate(GRADES)
        },
        grade_probabilities={
            grade_name: float(log_probabilities[0, index].exp().item())
            for index, grade_name in enumerate(GRADES)
        },
        duration_seconds=duration_seconds,
        decision=decision,
        gate_margin=gate_margin,
        reject_band=band,
    )


def _cached_backbone() -> None:
    try:
        _ = snapshot_download(
            repo_id=BACKBONE.repository,
            revision=BACKBONE.revision,
            local_files_only=True,
        )
    except LocalEntryNotFoundError as error:
        raise BackendLoadError(
            detail=(
                "pinned MuQ snapshot is not complete in the local cache; "
                "automatic download is disabled"
            ),
        ) from error


def _require_hf_offline() -> None:
    if os.environ.get("HF_HUB_OFFLINE") != "1" or not hf_constants.HF_HUB_OFFLINE:
        raise BackendLoadError(
            detail=(
                "Hugging Face offline mode is required; set HF_HUB_OFFLINE=1 "
                "before starting the process"
            ),
        )


def _load_inner_model(physical_cuda_device: str) -> InnerAudioModel:
    _require_hf_offline()
    if os.environ.get("CUDA_VISIBLE_DEVICES") != physical_cuda_device:
        raise BackendLoadError(
            detail=(
                "CUDA_VISIBLE_DEVICES must expose only physical GPU "
                f"{physical_cuda_device}"
            ),
        )
    if not torch.cuda.is_available():
        raise BackendLoadError(detail="CUDA is unavailable")
    _cached_backbone()
    try:
        return load_real_inner_model(local_files_only=True)
    except _CUDA_INIT_ERRORS as error:
        raise BackendLoadError(
            detail=f"cannot initialize CUDA models: {error}",
        ) from error


def _load_cuda_components(
    checkpoint: LoadedCheckpoint,
    *,
    physical_cuda_device: str = PHYSICAL_CUDA_DEVICE,
) -> ServingBackend:
    inner_model = _load_inner_model(physical_cuda_device)
    try:
        head = checkpoint.head.eval().to(
            device=LOGICAL_CUDA_DEVICE,
            dtype=torch.float32,
        )
    except _CUDA_INIT_ERRORS as error:
        raise BackendLoadError(
            detail=f"cannot initialize CUDA models: {error}",
        ) from error
    return ServingBackend(
        inner_model=inner_model,
        head=head,
        device=torch.device(LOGICAL_CUDA_DEVICE),
        checkpoint_sha256=checkpoint.sha256,
        model_family=checkpoint.model_family,
    )


def _load_hierarchical_components(
    package: LoadedPackage,
    *,
    physical_cuda_device: str,
) -> HierarchicalBackend:
    inner_model = _load_inner_model(physical_cuda_device)
    try:
        for member in package.members:
            _ = member.head.eval().to(device=LOGICAL_CUDA_DEVICE, dtype=torch.float32)
    except _CUDA_INIT_ERRORS as error:
        raise BackendLoadError(
            detail=f"cannot initialize CUDA models: {error}",
        ) from error
    return HierarchicalBackend(
        inner_model=inner_model,
        members=package.members,
        band=package.band,
        device=torch.device(LOGICAL_CUDA_DEVICE),
        checkpoint_sha256=package.sha256,
        model_family=package.model_family,
    )


def load_backend(
    checkpoint: Path,
    *,
    physical_cuda_device: str = PHYSICAL_CUDA_DEVICE,
) -> ServingBackend | HierarchicalBackend:
    """Load a v4 checkpoint file or a v5 package directory and the CUDA models."""
    if type(physical_cuda_device) is not str or physical_cuda_device not in (
        "0",
        "1",
        "3",
    ):
        raise BackendLoadError(detail="physical_cuda_device must be '0', '1', or '3'")
    if checkpoint.is_dir():
        return _load_hierarchical_components(
            load_package(checkpoint),
            physical_cuda_device=physical_cuda_device,
        )
    return _load_cuda_components(
        load_head_checkpoint(checkpoint),
        physical_cuda_device=physical_cuda_device,
    )
