"""Synthetic-safe batch orchestration for temporal feature caches."""

import json
from collections.abc import Sequence
from dataclasses import dataclass, replace
from pathlib import Path
from typing import assert_never, override

import torch

from raina_laya.features.audio_features.domain.bins import (
    TemporalContractError,
    merge_empty_frame_bins,
    plan_temporal_bins,
)
from raina_laya.features.audio_features.infrastructure.audio import (
    CudaResamplerContext,
    prepare_audio_context,
)
from raina_laya.features.audio_features.infrastructure.cache_schema import (
    DEFAULT_CACHE_POLICY,
    CacheIncompleteError,
    CachePolicy,
    FeatureCacheError,
    PublishStatus,
    TemporalFeatureRecord,
)
from raina_laya.features.audio_features.infrastructure.feature_cache import (
    cache_record_path,
    load_record,
    publish_record,
)
from raina_laya.features.audio_features.infrastructure.features import (
    HIDDEN_STATE_COUNT,
    FeatureExtractionError,
    FrameWindow,
    HiddenStateContractError,
    InnerAudioModel,
    LayerFrames,
    extract_inner_model_frame_batch,
    extract_inner_model_frames,
    extract_inner_model_layer_frames,
    pool_frame_window,
    pool_layer_frames,
)

__all__ = [
    "CacheBuildFailure",
    "CacheBuildItem",
    "CacheBuildResult",
    "FailureReceiptExistsError",
    "build_feature_cache",
    "extract_feature_record",
    "extract_layer_records",
    "write_failure_receipt",
]


@dataclass(frozen=True, slots=True)
class CacheBuildItem:
    """One explicitly approved in-memory source for cache construction."""

    source_identifier: str
    source_sha256: str
    waveform: torch.Tensor
    sample_rate: int


@dataclass(frozen=True, slots=True)
class CacheBuildFailure:
    """One source that failed without being published."""

    source_identifier: str
    reason: str


@dataclass(frozen=True, slots=True)
class CacheBuildResult:
    """Deterministic publication summary with explicit failures."""

    published: int
    skipped: int
    failures: tuple[CacheBuildFailure, ...]

    @property
    def success(self) -> bool:
        """Return whether every approved source published or verified."""
        return not self.failures


@dataclass(frozen=True, slots=True)
class FailureReceiptExistsError(Exception):
    """A failure receipt already exists and cannot be overwritten."""

    path: Path

    @override
    def __str__(self) -> str:
        """Describe the protected receipt path."""
        return f"failure receipt already exists: {self.path}"


def extract_feature_record(
    item: CacheBuildItem,
    model: InnerAudioModel,
    policy: CachePolicy,
    *,
    window_batch_size: int = 1,
    resampler: CudaResamplerContext | None = None,
) -> TemporalFeatureRecord:
    """Extract one complete temporal feature record without publishing it."""
    if window_batch_size < 1:
        raise HiddenStateContractError(reason="window batch size must be positive")
    duration_seconds = item.waveform.shape[-1] / item.sample_rate
    plan = plan_temporal_bins(duration_seconds)
    frames_by_window: dict[float, FrameWindow] = {}
    frame_counts: list[int] = []
    for offset in range(0, len(plan.windows), window_batch_size):
        windows = plan.windows[offset : offset + window_batch_size]
        contexts = tuple(
            prepare_audio_context(
                item.waveform[
                    ...,
                    round(window.start * item.sample_rate) : min(
                        round(window.end * item.sample_rate),
                        item.waveform.shape[-1],
                    ),
                ],
                item.sample_rate,
                resampler=resampler,
            )
            for window in windows
        )
        if window_batch_size == 1:
            frame_batch = (
                extract_inner_model_frames(
                    model,
                    contexts[0],
                    window_start=windows[0].start,
                ),
            )
        else:
            frame_batch = extract_inner_model_frame_batch(
                model,
                contexts,
                window_starts=tuple(window.start for window in windows),
            )
        for window, frames in zip(windows, frame_batch, strict=True):
            frames_by_window[window.start] = frames
            frame_counts.extend(
                int(
                    (
                        frames.valid_frame_mask
                        & (frames.frame_timestamps >= interval.start)
                        & (frames.frame_timestamps < interval.end)
                    )
                    .sum()
                    .item(),
                )
                for interval in plan.bins_owned_by(window.start)
            )
    bins = merge_empty_frame_bins(plan.bins, frame_counts, duration_seconds)
    features: list[torch.Tensor] = []
    for window in plan.windows:
        owned = tuple(
            interval for interval in bins if interval.owner_window_start == window.start
        )
        if owned:
            features.append(
                pool_frame_window(frames_by_window[window.start], owned),
            )
    return TemporalFeatureRecord(
        source_identifier=item.source_identifier,
        source_sha256=item.source_sha256,
        features=torch.cat(features).to(dtype=torch.float32),
        bins=bins,
        padding_mask=torch.tensor(
            tuple(interval.valid_ratio < 1.0 for interval in bins),
            dtype=torch.bool,
        ),
        policy=policy,
    )


def extract_layer_records(
    item: CacheBuildItem,
    model: InnerAudioModel,
    policy: CachePolicy,
    *,
    resampler: CudaResamplerContext | None = None,
) -> tuple[TemporalFeatureRecord, ...]:
    """Run MuQ once per window and return one width-1024 record per hidden state.

    Planning, validity, bin merging and pooling are those of `extract_feature_record`
    with one window per forward pass; record `i` carries hidden state `i` (0..12).
    """
    duration_seconds = item.waveform.shape[-1] / item.sample_rate
    plan = plan_temporal_bins(duration_seconds)
    frames_by_window: dict[float, LayerFrames] = {}
    frame_counts: list[int] = []
    for window in plan.windows:
        context = prepare_audio_context(
            item.waveform[
                ...,
                round(window.start * item.sample_rate) : min(
                    round(window.end * item.sample_rate),
                    item.waveform.shape[-1],
                ),
            ],
            item.sample_rate,
            resampler=resampler,
        )
        frames = extract_inner_model_layer_frames(
            model, context, window_start=window.start
        )
        frames_by_window[window.start] = frames
        frame_counts.extend(
            int(
                (
                    frames.valid_frame_mask
                    & (frames.frame_timestamps >= interval.start)
                    & (frames.frame_timestamps < interval.end)
                )
                .sum()
                .item(),
            )
            for interval in plan.bins_owned_by(window.start)
        )
    bins = merge_empty_frame_bins(plan.bins, frame_counts, duration_seconds)
    pooled: list[torch.Tensor] = []
    for window in plan.windows:
        owned = tuple(
            interval for interval in bins if interval.owner_window_start == window.start
        )
        if owned:
            pooled.append(pool_layer_frames(frames_by_window[window.start], owned))
    stacked = torch.cat(pooled).to(dtype=torch.float32)
    padding_mask = torch.tensor(
        tuple(interval.valid_ratio < 1.0 for interval in bins),
        dtype=torch.bool,
    )
    return tuple(
        TemporalFeatureRecord(
            source_identifier=item.source_identifier,
            source_sha256=item.source_sha256,
            features=stacked[:, layer].contiguous(),
            bins=bins,
            padding_mask=padding_mask,
            policy=replace(policy, layer_order=(layer,)),
        )
        for layer in range(HIDDEN_STATE_COUNT)
    )


def _verified_restart(
    item: CacheBuildItem,
    cache_root: Path,
    policy: CachePolicy,
) -> bool:
    target = cache_record_path(cache_root, item.source_identifier, policy)
    try:
        _ = load_record(
            cache_root,
            item.source_identifier,
            expected_source_sha256=item.source_sha256,
            policy=policy,
        )
    except CacheIncompleteError:
        if target.exists():
            raise
        return False
    return True


def build_feature_cache(
    items: Sequence[CacheBuildItem],
    *,
    model: InnerAudioModel,
    cache_root: Path,
    policy: CachePolicy = DEFAULT_CACHE_POLICY,
) -> CacheBuildResult:
    """Build each approved synthetic record once and report every failure."""
    published = 0
    skipped = 0
    failures: list[CacheBuildFailure] = []
    for item in items:
        try:
            if _verified_restart(item, cache_root, policy):
                skipped += 1
                continue
            record = extract_feature_record(item, model, policy)
            status = publish_record(cache_root, record)
            match status:
                case PublishStatus.PUBLISHED:
                    published += 1
                case PublishStatus.VERIFIED_EXISTS:
                    skipped += 1
                case unreachable:
                    assert_never(unreachable)
        except (
            FeatureExtractionError,
            FeatureCacheError,
            TemporalContractError,
        ) as error:
            failures.append(
                CacheBuildFailure(
                    source_identifier=item.source_identifier,
                    reason=str(error),
                ),
            )
    return CacheBuildResult(
        published=published,
        skipped=skipped,
        failures=tuple(failures),
    )


def write_failure_receipt(path: Path, result: CacheBuildResult) -> None:
    """Create a separate failure receipt without replacing prior evidence."""
    payload = {
        "success": result.success,
        "published": result.published,
        "skipped": result.skipped,
        "failures": [
            {
                "source_identifier": failure.source_identifier,
                "reason": failure.reason,
            }
            for failure in result.failures
        ],
    }
    path.parent.mkdir(parents=True, exist_ok=True)
    try:
        with path.open("x", encoding="utf-8") as receipt:
            json.dump(payload, receipt, sort_keys=True, separators=(",", ":"))
    except FileExistsError:
        raise FailureReceiptExistsError(path=path) from None
