# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0

"""Bounded compressed-byte and decoder cache for episode video streaming."""

from __future__ import annotations

import io
import logging
import threading
from collections import OrderedDict
from collections.abc import Callable
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass, field
from pathlib import Path
from types import SimpleNamespace
from typing import TYPE_CHECKING, BinaryIO, TypedDict

import torch

from lerobot.streaming.manifest import EpisodeVideoManifest
from lerobot.streaming.mp4 import Mp4SampleSlice, synthesize_mp4
from lerobot.streaming.range_fetch import make_range_fetcher

if TYPE_CHECKING:
    from torchcodec.decoders import VideoDecoder

logger = logging.getLogger(__name__)


class _VideoPayload(TypedDict):
    """Synthesized standalone MP4 bytes for one episode camera."""

    bytes: bytes


@dataclass
class _DecoderEntry:
    """Keep serialization attached to the decoder, including after LRU eviction."""

    decoder: VideoDecoder | _PyAVVideoDecoder
    lock: threading.Lock = field(default_factory=threading.Lock)

    def close(self) -> None:
        """Close the underlying decoder while preserving its backend's lease rules."""
        _close_decoder(self.decoder)


class EpisodeByteCache:
    """Fetch, synthesize, and retain episode-local MP4s within bounded caches."""

    def __init__(
        self,
        manifest: EpisodeVideoManifest,
        data_root: str | Path,
        *,
        byte_budget: int = 80 * 1024**3,
        workers: int = 8,
        range_backend: str = "fsspec",
        native_http_connections: int | None = None,
        native_http_timeout: float = 60.0,
        native_http_retries: int = 4,
        native_http_subranges: int = 1,
        max_open_decoders: int = 64,
        video_backend: str = "torchcodec",
        tolerance_s: float = 1e-4,
        token: str | bool | None = None,
    ) -> None:
        """Configure byte fetching, synthesis, decoder limits, and backend selection.

        Arguments are not re-validated: ``StreamingLeRobotDataset``, the only production caller,
        passes resolved values (e.g. ``video_reader`` is already mapped to ``pyav``).
        """
        self.manifest = manifest
        self.fetcher = make_range_fetcher(
            data_root,
            range_backend=range_backend,
            workers=workers,
            native_http_connections=native_http_connections,
            native_http_timeout=native_http_timeout,
            native_http_retries=native_http_retries,
            native_http_subranges=native_http_subranges,
            token=token,
        )
        self.byte_budget = byte_budget
        self.max_open_decoders = max_open_decoders
        self.video_backend = video_backend
        self.tolerance_s = tolerance_s
        self._pool = ThreadPoolExecutor(max_workers=workers)
        self._cache: OrderedDict[tuple[int, str], _VideoPayload] = OrderedDict()
        self._decoders: OrderedDict[tuple[int, str], _DecoderEntry] = OrderedDict()
        self._futures: dict[tuple[int, str], Future[None]] = {}
        self._reservations: OrderedDict[int, int] = OrderedDict()
        self._reserved_bytes = 0
        self._retained_episodes: dict[int, int] = {}
        self._decoder_fallback_count = 0
        self._fallback_warning_emitted = False
        self._bytes = 0
        self._lock = threading.RLock()
        self._space_available = threading.Condition(self._lock)
        self._closed = False

    def close(self) -> None:
        """Close fetchers, executors, decoders, and cached state."""
        with self._space_available:
            self._closed = True
            self._space_available.notify_all()
        self._pool.shutdown(wait=True, cancel_futures=True)
        with self._lock:
            decoders = list(self._decoders.values())
            self._cache.clear()
            self._decoders.clear()
            self._futures.clear()
            self._retained_episodes.clear()
            self._bytes = 0
            self._reservations.clear()
            self._reserved_bytes = 0
        for decoder in decoders:
            _close_decoder(decoder)
        self.fetcher.close()

    def __enter__(self) -> EpisodeByteCache:
        """Return this cache as a context manager."""
        return self

    def __exit__(self, *_exc: object) -> None:
        """Close the cache when leaving its context."""
        self.close()

    def submit_prefetch(self, episode_index: int) -> bool:
        """Schedule all cameras if the episode fits; never wait for speculative work.

        A False result means the caller may retry after releasing another episode.
        """
        with self._space_available:
            if not self._reserve_locked(episode_index):
                return False
            for camera_key in self.manifest.video_keys:
                self._submit_locked(episode_index, camera_key)
            return True

    def retain_episode(self, episode_index: int, *, wait: bool = False) -> None:
        """Reserve all camera bytes and prevent eviction until released.

        With wait=True, admission waits for another thread's fetch or decode lease
        to finish. Callers must not wait on leases that only they can release.
        """
        with self._space_available:
            while not self._reserve_locked(episode_index):
                if not wait:
                    raise MemoryError(f"Episode {episode_index} does not fit the available byte budget")
                self._space_available.wait()
            self._retained_episodes[episode_index] = self._retained_episodes.get(episode_index, 0) + 1

    def release_episode(self, episode_index: int) -> None:
        """Release one retention lease and wake blocked admissions."""
        with self._space_available:
            count = self._retained_episodes.get(episode_index, 0)
            if count <= 1:
                self._retained_episodes.pop(episode_index, None)
            else:
                self._retained_episodes[episode_index] = count - 1
            self._space_available.notify_all()

    @property
    def resident_bytes(self) -> int:
        """Return all completed payload bytes, including unconsumed prefetches."""
        with self._lock:
            return self._bytes

    @property
    def reserved_bytes(self) -> int:
        """Return bytes reserved for pending, completed, and leased episode payloads.

        This bounds cached compressed payloads, not decoder memory, decoded tensors,
        or temporary range-fetch and synthesis buffers.
        """
        with self._lock:
            return self._reserved_bytes

    @property
    def open_decoder_count(self) -> int:
        """Return the number of cached video decoders."""
        with self._lock:
            return len(self._decoders)

    @property
    def decoder_fallback_count(self) -> int:
        """Return the number of TorchCodec-to-PyAV fallbacks."""
        with self._lock:
            return self._decoder_fallback_count

    def get_bytes(self, episode_index: int, camera_key: str) -> bytes:
        """Return the synthesized MP4 bytes for one episode camera."""
        return self._get_entry(episode_index, camera_key)["bytes"]

    def _get_decoder_entry(self, episode_index: int, camera_key: str) -> _DecoderEntry:
        """Reuse or open a decoder, resolving concurrent opens before LRU eviction."""
        key = (episode_index, camera_key)
        entry = self._get_entry(episode_index, camera_key)
        with self._lock:
            decoder = self._decoders.get(key)
            if decoder is not None:
                self._decoders.move_to_end(key)
                return decoder

        decoder = _DecoderEntry(self._open_decoder(key, entry["bytes"]))
        with self._lock:
            existing = self._decoders.get(key)
            if existing is not None:
                self._decoders.move_to_end(key)
                _close_decoder(decoder)
                return existing
            self._decoders[key] = decoder
            while len(self._decoders) > self.max_open_decoders:
                _, evicted_decoder = self._decoders.popitem(last=False)
                _close_decoder(evicted_decoder)
        return decoder

    def _open_decoder(self, key: tuple[int, str], data: bytes) -> VideoDecoder | _PyAVVideoDecoder:
        """Open an episode decoder, falling back to PyAV if TorchCodec rejects the bytes."""
        try:
            if self.video_backend == "torchcodec":
                # Raw bytes: TorchCodec then reads from memory without a Python read callback,
                # which would need the GIL for every FFmpeg read.
                return open_video_decoder(data)
            return open_video_decoder(io.BytesIO(data), backend=self.video_backend)
        except Exception as primary_error:
            if self.video_backend != "torchcodec":
                raise
            try:
                decoder = open_video_decoder(io.BytesIO(data), backend="pyav")
            except Exception as fallback_error:
                raise RuntimeError(
                    "Both TorchCodec and PyAV rejected synthesized episode video "
                    f"{key}: TorchCodec error: {primary_error}"
                ) from fallback_error
            with self._lock:
                self._decoder_fallback_count += 1
                should_warn = not self._fallback_warning_emitted
                self._fallback_warning_emitted = True
            if should_warn:
                logger.warning(
                    "TorchCodec rejected a synthesized episode MP4; using the bounded PyAV "
                    "decoder fallback for affected videos. First error: %s",
                    primary_error,
                )
            else:
                logger.debug("Using PyAV decoder fallback for synthesized episode video %s", key)
            return decoder

    def get_frames(self, episode_index: int, camera_key: str, timestamps: list[float]) -> torch.Tensor:
        """Decode source-timeline timestamps from an episode-local MP4.

        Args:
            episode_index (`int`):
                Episode identifier present in the manifest.
            camera_key (`str`):
                Video feature key identifying the camera.
            timestamps (`list[float]`):
                Nonempty request in source-video seconds, including the episode's video offset.

        Returns:
            `torch.Tensor`: Uint8 RGB frames in request order with shape (frames, channels, height, width).

        Raises:
            ValueError: A returned frame's presentation timestamp exceeds the configured tolerance.

        Note:
            The episode's bytes remain retained for the request, and access to its decoder is serialized.
        """
        self.retain_episode(episode_index)
        try:
            return self._get_frames(episode_index, camera_key, timestamps)
        finally:
            self.release_episode(episode_index)

    def _get_frames(self, episode_index: int, camera_key: str, timestamps: list[float]) -> torch.Tensor:
        """Translate source timestamps and serialize access to the leased decoder."""
        span = self.manifest.lookup(episode_index, camera_key)
        local_ts = [ts - span.source_start_pts for ts in timestamps]
        entry, release = self._decoder_for_frames(episode_index, camera_key)
        decoder = entry.decoder
        try:
            with entry.lock:
                if self.video_backend == "pyav" or isinstance(decoder, _PyAVVideoDecoder):
                    return decoder.get_frames_played_at(local_ts, tolerance_s=self.tolerance_s).data
                metadata = decoder.metadata
                fps = getattr(metadata, "average_fps", None)
                if fps is None:
                    duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9)
                    fps = metadata.num_frames / duration
                num_frames = getattr(metadata, "num_frames", None)
                if num_frames is None:
                    duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9)
                    num_frames = round(duration * fps)
                last_index = max(0, int(num_frames) - 1)
                indices = [min(max(round(ts * fps), 0), last_index) for ts in local_ts]
                frames = decoder.get_frames_at(indices=indices)
                # Estimated indices and approximate seeking can drift on variable-rate videos.
                # Check the actual PTS in request order, including clamped and repeated frames.
                # Plain floats: a few scalars per camera are cheaper than building small tensors.
                query_ts = [float(ts) for ts in local_ts]
                decoded_ts = frames.pts_seconds.tolist()
                if not all(abs(q - d) <= self.tolerance_s for q, d in zip(query_ts, decoded_ts, strict=True)):
                    raise ValueError(
                        f"TorchCodec frame timestamps exceed tolerance {self.tolerance_s}: "
                        f"episode={episode_index}, camera={camera_key}, "
                        f"queries={query_ts}, decoded={decoded_ts}"
                    )
                return frames.data
        finally:
            if release is not None:
                release()

    def _decoder_for_frames(
        self, episode_index: int, camera_key: str
    ) -> tuple[_DecoderEntry, Callable[[], None] | None]:
        """Lease a decoder and retry if eviction wins the race with acquisition."""
        key = (episode_index, camera_key)
        while True:
            entry = self._get_decoder_entry(episode_index, camera_key)
            decoder = entry.decoder
            acquire = getattr(decoder, "acquire", None)
            if acquire is None:
                return entry, None
            try:
                acquire()
            except RuntimeError:
                # The decoder was evicted between lookup and lease acquisition. Remove a stale
                # cached reference if it raced with close, then retry with a fresh decoder.
                with self._lock:
                    if self._decoders.get(key) is entry:
                        self._decoders.pop(key)
                continue
            return entry, decoder.release

    def _submit_locked(self, episode_index: int, camera_key: str) -> Future[None]:
        """Schedule one camera fetch at most once while the cache lock is held."""
        key = (episode_index, camera_key)
        future = self._futures.get(key)
        if future is None:
            future = self._pool.submit(self._fetch_and_store, episode_index, camera_key)
            self._futures[key] = future
            future.add_done_callback(self._notify_fetch_done)
        return future

    def _notify_fetch_done(self, _future: Future[None]) -> None:
        """Wake admissions blocked on a pending episode fetch."""
        with self._space_available:
            self._space_available.notify_all()

    def _reserve_locked(self, episode_index: int) -> bool:
        """Reserve an episode's indexed bytes, evicting only unleased completed work."""
        if self._closed:
            raise RuntimeError("Episode byte cache is closed")
        if episode_index in self._reservations:
            self._reservations.move_to_end(episode_index)
            return True
        size = self.manifest.episode_byte_size(episode_index)
        if size > self.byte_budget:
            raise MemoryError(
                f"Episode {episode_index} exceeds the byte budget ({size} > {self.byte_budget})"
            )
        while self._reserved_bytes + size > self.byte_budget:
            evicted = next(
                (
                    episode
                    for episode in self._reservations
                    if episode not in self._retained_episodes
                    and all(future.done() for key, future in self._futures.items() if key[0] == episode)
                ),
                None,
            )
            if evicted is None:
                return False
            self._evict_episode_locked(evicted)
        self._reservations[episode_index] = size
        self._reserved_bytes += size
        return True

    def _get_entry(self, episode_index: int, camera_key: str) -> _VideoPayload:
        """Wait for one camera payload while retaining its episode reservation."""
        self.retain_episode(episode_index)
        try:
            key = (episode_index, camera_key)
            with self._lock:
                future = self._futures.get(key)
                if future is None:
                    future = self._submit_locked(episode_index, camera_key)
            future.result()
            with self._lock:
                return self._cache[key]
        finally:
            self.release_episode(episode_index)

    def _fetch_and_store(self, episode_index: int, camera_key: str) -> None:
        # Futures carry no payload: the accounted cache is the sole owner of completed bytes.
        """Store a synthesized payload under its episode reservation."""
        entry = self._fetch_and_synthesize(episode_index, camera_key)
        with self._lock:
            size = len(entry["bytes"])
            episode_bytes = sum(
                len(value["bytes"]) for key, value in self._cache.items() if key[0] == episode_index
            )
            if episode_bytes + size > self._reservations[episode_index]:
                raise ValueError(f"Synthesized episode {episode_index} exceeds its indexed byte reservation")
            self._cache[episode_index, camera_key] = entry
            self._bytes += size

    def _evict_episode_locked(self, episode_index: int) -> None:
        """Remove an unleased episode's bytes, futures and decoders under the cache lock."""
        self._reserved_bytes -= self._reservations.pop(episode_index)
        for camera_key in self.manifest.video_keys:
            key = (episode_index, camera_key)
            entry = self._cache.pop(key, None)
            if entry is not None:
                self._bytes -= len(entry["bytes"])
            self._futures.pop(key, None)
            decoder = self._decoders.pop(key, None)
            if decoder is not None:
                _close_decoder(decoder)

    def _fetch_and_synthesize(self, episode_index: int, camera_key: str) -> _VideoPayload:
        """Fetch a camera span and wrap it in a standalone MP4."""
        span = self.manifest.lookup(episode_index, camera_key)
        file_record = self.manifest.file_lookup(span.file_id)
        sample_slice = Mp4SampleSlice(
            sample_lo=span.sample_lo,
            sample_hi=span.sample_hi,
            byte_offset=span.mdat_offset,
            byte_length=span.mdat_length,
            source_start_pts=span.source_start_pts,
        )
        payload = self.fetcher.read_range(file_record.file_path, span.mdat_offset, span.mdat_length)
        if len(payload) != span.mdat_length:
            raise OSError(
                f"Short read for {file_record.file_path}: expected {span.mdat_length}, got {len(payload)}"
            )
        return {"bytes": synthesize_mp4(file_record.mp4, sample_slice, payload)}


class _PyAVVideoDecoder:
    """Small seekable PyAV adapter matching the TorchCodec calls used by the byte cache."""

    def __init__(self, file_like_or_bytesio: BinaryIO) -> None:
        """Open a seekable in-memory video and initialize decoder lifetime accounting."""
        import av

        self._source = file_like_or_bytesio
        self._container = av.open(file_like_or_bytesio)
        self._stream = self._container.streams.video[0]
        average_rate = self._stream.average_rate
        if average_rate is None:
            raise ValueError("PyAV video stream does not expose an average frame rate")
        self._fps = float(average_rate)
        duration = (
            float(self._stream.duration * self._stream.time_base)
            if self._stream.duration is not None
            else 0.0
        )
        self.metadata = SimpleNamespace(
            average_fps=self._fps,
            num_frames=int(self._stream.frames or round(duration * self._fps)),
            begin_stream_seconds=0.0,
            end_stream_seconds=duration,
        )
        self._decode_lock = threading.Lock()
        self._state_lock = threading.Lock()
        self._users = 0
        self._close_requested = False
        self._closed = False

    def acquire(self) -> None:
        """Retain the decoder for an active request, rejecting a closing decoder."""
        with self._state_lock:
            if self._close_requested or self._closed:
                raise RuntimeError("PyAV decoder is closing")
            self._users += 1

    def release(self) -> None:
        """Release a request lease and finish any deferred close after the last user."""
        with self._state_lock:
            self._users -= 1
            if self._users < 0:
                raise RuntimeError("Unbalanced PyAV decoder release")
            if self._users == 0 and self._close_requested:
                self._close_resources()

    def get_frames_at(self, *, indices: list[int]) -> SimpleNamespace:
        """Decode zero-based frame indices in request order as uint8 RGB tensors."""
        if not indices:
            return SimpleNamespace(data=torch.empty((0, 3, 0, 0), dtype=torch.uint8))
        timestamps = [index / self._fps for index in indices]
        return self._get_frames_played_at(timestamps, tolerance_s=0.5 / self._fps + 1e-6)

    def get_frames_played_at(
        self,
        timestamps: list[float],
        *,
        tolerance_s: float,
    ) -> SimpleNamespace:
        """Decode local timestamps in seconds within the requested tolerance."""
        return self._get_frames_played_at(timestamps, tolerance_s=tolerance_s)

    def _get_frames_played_at(
        self,
        timestamps: list[float],
        *,
        tolerance_s: float,
    ) -> SimpleNamespace:
        """Seek once and decode through the requested window under the decode lock."""
        first_ts = min(timestamps)
        last_ts = max(timestamps)
        loaded_frames: list[torch.Tensor] = []
        loaded_ts: list[float] = []
        with self._decode_lock:
            self._container.seek(
                round(first_ts / self._stream.time_base) - 1,
                backward=True,
                any_frame=False,
                stream=self._stream,
            )
            for frame in self._container.decode(self._stream):
                if frame.pts is None:
                    continue
                current_ts = float(frame.pts * self._stream.time_base)
                array = frame.to_ndarray(format="rgb24")
                loaded_frames.append(torch.from_numpy(array).permute(2, 0, 1).contiguous())
                loaded_ts.append(current_ts)
                if current_ts >= last_ts:
                    break

        if not loaded_frames:
            raise ValueError(f"PyAV decoded no frames for timestamps {timestamps}")
        query_ts = torch.tensor(timestamps, dtype=torch.float64)
        loaded_ts_tensor = torch.tensor(loaded_ts, dtype=torch.float64)
        distances = torch.cdist(query_ts[:, None], loaded_ts_tensor[:, None], p=1)
        minimum, closest = distances.min(1)
        if not (minimum <= tolerance_s).all():
            raise ValueError(
                f"PyAV frame timestamps exceed tolerance {tolerance_s}: "
                f"queries={query_ts}, decoded={loaded_ts_tensor}"
            )
        return SimpleNamespace(data=torch.stack([loaded_frames[index] for index in closest]))

    def close(self) -> None:
        """Request closure, deferring resource release until active leases finish."""
        with self._state_lock:
            self._close_requested = True
            if self._users == 0:
                self._close_resources()

    def _close_resources(self) -> None:
        """Close the container and its source once, under the state lock."""
        if self._closed:
            return
        self._container.close()
        close = getattr(self._source, "close", None)
        if close is not None:
            close()
        self._closed = True


def _close_decoder(decoder: object) -> None:
    """Close a backend handle when supported without masking another failure."""
    close = getattr(decoder, "close", None)
    if close is not None:
        try:
            close()
        except Exception:
            logger.debug("Failed to close video decoder", exc_info=True)


def open_video_decoder(
    file_like_or_bytesio: bytes | BinaryIO, *, backend: str = "torchcodec"
) -> VideoDecoder | _PyAVVideoDecoder:
    """Open a TorchCodec or PyAV decoder over synthesized MP4 bytes."""
    if backend == "pyav":
        if isinstance(file_like_or_bytesio, bytes):
            file_like_or_bytesio = io.BytesIO(file_like_or_bytesio)
        return _PyAVVideoDecoder(file_like_or_bytesio)
    if backend != "torchcodec":
        raise ValueError(f"Unsupported video backend: {backend}")
    from torchcodec.decoders import VideoDecoder

    return VideoDecoder(file_like_or_bytesio, seek_mode="approximate")
