"""Deterministic temporal windows and globally owned pooling bins."""

from collections.abc import Sequence
from dataclasses import dataclass
from math import isfinite
from typing import Final, override

__all__ = [
    "ContextWindow",
    "FrameBinAssignment",
    "GlobalBin",
    "InvalidBinPlanError",
    "InvalidDurationError",
    "NoValidFrameBinError",
    "TemporalBinPolicy",
    "TemporalPlan",
    "merge_empty_frame_bins",
    "plan_temporal_bins",
]


@dataclass(frozen=True, slots=True)
class TemporalBinPolicy:
    """Timing constants for context planning and global pooling."""

    context_seconds: float = 10.0
    pool_bin_seconds: float = 2.0
    min_tail_bin_seconds: float = 0.5


DEFAULT_POLICY: Final = TemporalBinPolicy()


@dataclass(frozen=True, slots=True)
class ContextWindow:
    """One fixed-duration model context on the source timeline."""

    start: float
    end: float


@dataclass(frozen=True, slots=True)
class GlobalBin:
    """One global pooling interval and its earliest owning context."""

    start: float
    end: float
    center: float
    valid_ratio: float
    owner_window_start: float


@dataclass(frozen=True, slots=True)
class FrameBinAssignment:
    """One frame timestamp assigned to one half-open global bin."""

    timestamp: float
    interval: GlobalBin


@dataclass(frozen=True, slots=True)
class TemporalPlan:
    """Complete context and uniquely owned global-bin plan."""

    windows: tuple[ContextWindow, ...]
    bins: tuple[GlobalBin, ...]

    def bins_owned_by(self, window_start: float) -> tuple[GlobalBin, ...]:
        """Return bins assigned to one context window."""
        return tuple(
            interval
            for interval in self.bins
            if interval.owner_window_start == window_start
        )

    def assign_frame_timestamps(
        self,
        frame_timestamps: Sequence[float],
    ) -> tuple[FrameBinAssignment, ...]:
        """Assign in-range frame timestamps to exactly one half-open bin."""
        return tuple(
            FrameBinAssignment(timestamp=timestamp, interval=interval)
            for timestamp in frame_timestamps
            for interval in self.bins
            if interval.start <= timestamp < interval.end
        )


@dataclass(frozen=True, slots=True)
class TemporalContractError(Exception):
    """Base error for invalid temporal feature inputs."""


@dataclass(frozen=True, slots=True)
class InvalidDurationError(TemporalContractError):
    """The source duration cannot define a temporal plan."""

    duration_seconds: float

    @override
    def __str__(self) -> str:
        """Describe the rejected duration."""
        return (
            "audio duration must be finite and greater than zero; "
            f"received {self.duration_seconds}"
        )


@dataclass(frozen=True, slots=True)
class InvalidBinPlanError(TemporalContractError):
    """A bin sequence or frame-count mapping violates the domain contract."""

    reason: str

    @override
    def __str__(self) -> str:
        """Describe the rejected bin plan."""
        return f"invalid temporal bin plan: {self.reason}"


@dataclass(frozen=True, slots=True)
class NoValidFrameBinError(TemporalContractError):
    """No pooling bin contains a valid model frame."""

    @override
    def __str__(self) -> str:
        """Describe the empty pooling result."""
        return "temporal bin plan contains no valid model frames"


def _validated_policy(policy: TemporalBinPolicy) -> TemporalBinPolicy:
    values = (
        ("context_seconds", policy.context_seconds),
        ("pool_bin_seconds", policy.pool_bin_seconds),
        ("min_tail_bin_seconds", policy.min_tail_bin_seconds),
    )
    for field, value in values:
        if not isfinite(value) or value <= 0:
            raise InvalidBinPlanError(reason=f"{field} must be finite and positive")
    if policy.min_tail_bin_seconds > policy.pool_bin_seconds:
        raise InvalidBinPlanError(
            reason="min_tail_bin_seconds cannot exceed pool_bin_seconds",
        )
    if policy.pool_bin_seconds > policy.context_seconds:
        raise InvalidBinPlanError(
            reason="pool_bin_seconds cannot exceed context_seconds",
        )
    return policy


def _context_windows(
    duration_seconds: float,
    policy: TemporalBinPolicy,
) -> tuple[ContextWindow, ...]:
    context = policy.context_seconds
    if duration_seconds <= context:
        return (ContextWindow(start=0.0, end=context),)

    full_count = int(duration_seconds // context)
    starts = [float(index) * context for index in range(full_count)]
    if starts[-1] + context != duration_seconds:
        starts.append(duration_seconds - context)
    return tuple(ContextWindow(start=start, end=start + context) for start in starts)


def _global_intervals(
    duration_seconds: float,
    policy: TemporalBinPolicy,
) -> tuple[tuple[float, float], ...]:
    width = policy.pool_bin_seconds
    full_count = int(duration_seconds // width)
    intervals = [
        (float(index) * width, float(index + 1) * width) for index in range(full_count)
    ]
    full_end = float(full_count) * width
    tail = duration_seconds - full_end
    if tail > 0:
        if tail < policy.min_tail_bin_seconds and intervals:
            previous_start, _ = intervals[-1]
            intervals[-1] = (previous_start, duration_seconds)
        else:
            intervals.append((full_end, duration_seconds))
    return tuple(intervals)


def plan_temporal_bins(
    duration_seconds: float,
    policy: TemporalBinPolicy = DEFAULT_POLICY,
) -> TemporalPlan:
    """Plan fixed contexts and assign each global bin to its earliest owner."""
    if not isfinite(duration_seconds) or duration_seconds <= 0:
        raise InvalidDurationError(duration_seconds=duration_seconds)
    validated_policy = _validated_policy(policy)
    windows = _context_windows(duration_seconds, validated_policy)
    bins = tuple(
        GlobalBin(
            start=start,
            end=end,
            center=(start + end) / 2,
            valid_ratio=1.0,
            owner_window_start=next(
                window.start
                for window in windows
                if window.start <= start and end <= window.end
            ),
        )
        for start, end in _global_intervals(duration_seconds, validated_policy)
    )
    return TemporalPlan(windows=windows, bins=bins)


def _validate_frame_bins(
    bins: Sequence[GlobalBin],
    frame_counts: Sequence[int],
) -> None:
    if not bins:
        raise InvalidBinPlanError(reason="at least one bin is required")
    if len(bins) != len(frame_counts):
        raise InvalidBinPlanError(reason="frame counts must match the bin count")
    for index, (interval, frame_count) in enumerate(
        zip(bins, frame_counts, strict=True),
    ):
        if interval.start >= interval.end:
            raise InvalidBinPlanError(reason="every bin must have positive duration")
        if frame_count < 0:
            raise InvalidBinPlanError(reason="frame counts cannot be negative")
        if index > 0 and bins[index - 1].end != interval.start:
            raise InvalidBinPlanError(
                reason="bins must be ordered, contiguous, and non-overlapping",
            )


def _valid_ratio(start: float, end: float, duration_seconds: float) -> float:
    valid_start = max(start, 0.0)
    valid_end = min(end, duration_seconds)
    return max(0.0, valid_end - valid_start) / (end - start)


def merge_empty_frame_bins(
    bins: Sequence[GlobalBin],
    frame_counts: Sequence[int],
    duration_seconds: float,
) -> tuple[GlobalBin, ...]:
    """Merge zero-frame bins with the previous valid bin, then the next."""
    if not isfinite(duration_seconds) or duration_seconds <= 0:
        raise InvalidDurationError(duration_seconds=duration_seconds)
    _validate_frame_bins(bins, frame_counts)
    valid_indices = [
        index for index, frame_count in enumerate(frame_counts) if frame_count > 0
    ]
    if not valid_indices:
        raise NoValidFrameBinError

    merged: list[GlobalBin] = []
    for position, valid_index in enumerate(valid_indices):
        start = bins[0].start if position == 0 else bins[valid_index].start
        next_valid = (
            valid_indices[position + 1]
            if position + 1 < len(valid_indices)
            else len(bins)
        )
        end = bins[next_valid - 1].end
        source = bins[valid_index]
        merged.append(
            GlobalBin(
                start=start,
                end=end,
                center=(start + end) / 2,
                valid_ratio=_valid_ratio(start, end, duration_seconds),
                owner_window_start=source.owner_window_start,
            ),
        )
    return tuple(merged)
