#!/usr/bin/env python

# 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

import dataclasses
import struct
import threading
import time
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from types import SimpleNamespace

import numpy as np
import pytest
import torch

from lerobot.streaming.episode_cache import EpisodeByteCache
from lerobot.streaming.manifest import EpisodeVideoManifest
from lerobot.streaming.mp4 import (
    _box,
    _co64,
    _dinf,
    _hdlr,
    _mdhd,
    _mvhd,
    _stco,
    _stsc_one_sample_per_chunk,
    _stss,
    _stsz,
    _stts,
    _tkhd,
    _vmhd,
    fetch_mp4_index,
    parse_mp4_index,
    synthesize_mp4,
    synthesized_mp4_size,
)
from lerobot.streaming.range_fetch import ThreadLocalRangeFetcher


def _minimal_mp4(sample_offsets: list[int], *, use_co64: bool = False) -> bytes:
    ftyp = _box(b"ftyp", b"isom\0\0\2\0isomiso2mp41")
    sizes = np.array([10, 10, 10], dtype=np.int64)
    durations = np.array([1000, 1000, 1000], dtype=np.int64)
    stsd_body = struct.pack(">II", 0, 1) + struct.pack(">I4s", 16, b"avc1") + b"\0" * 8
    offsets = _co64(sample_offsets) if use_co64 else _stco(sample_offsets)
    stbl = _box(
        b"stbl",
        _box(b"stsd", stsd_body)
        + _stts(durations)
        + _stsc_one_sample_per_chunk(len(sizes))
        + _stsz(sizes)
        + offsets
        + _stss(np.array([1], dtype=np.int64)),
    )
    minf = _box(b"minf", _vmhd() + _dinf() + stbl)
    mdia = _box(b"mdia", _mdhd(1000, 3000) + _hdlr() + minf)
    trak = _box(b"trak", _tkhd(1, 3000, 64, 48) + mdia)
    moov = _box(b"moov", _mvhd(1000, 3000, 2) + trak)
    mdat_payload_start = 10_000
    free_size = mdat_payload_start - 8 - len(ftyp) - len(moov)
    assert free_size >= 8
    free = _box(b"free", b"\0" * (free_size - 8))
    return ftyp + moov + free + _box(b"mdat", b"x" * 128)


def test_episode_slice_uses_min_max_sample_offsets_for_reordered_chunks():
    mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025]))

    sample_slice = mp4.sample_slice(0.0, 2.0, keyframe_pad_s=0, keyframe_pad_fraction=0)

    assert sample_slice.byte_offset == 10_000
    assert sample_slice.byte_length == 60
    assert sample_slice.sample_lo == 0
    assert sample_slice.sample_hi == 2


def test_synthesized_mp4_rebases_one_chunk_per_sample_offsets():
    mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025]))
    sample_slice = mp4.sample_slice(0.0, 2.0, keyframe_pad_s=0, keyframe_pad_fraction=0)

    mini = synthesize_mp4(mp4, sample_slice, b"x" * sample_slice.byte_length)
    mini_index = parse_mp4_index("mini.mp4", mini)

    expected = np.array([0, 50, 25], dtype=np.int64) + mini_index.mdat_payload_offset
    np.testing.assert_array_equal(mini_index.sample_offsets, expected)
    np.testing.assert_array_equal(mini_index.sample_sizes, np.array([10, 10, 10]))


def test_synthesized_mp4_size_matches_materialized_bytes():
    mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025]))
    sample_slice = mp4.sample_slice(0.0, 2.0, keyframe_pad_s=0, keyframe_pad_fraction=0)

    mini = synthesize_mp4(mp4, sample_slice, b"x" * sample_slice.byte_length)

    assert synthesized_mp4_size(mp4, sample_slice) == len(mini)


def test_parser_accepts_co64_chunk_offsets():
    mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025], use_co64=True))

    np.testing.assert_array_equal(mp4.sample_offsets, np.array([10_000, 10_050, 10_025]))


def _sized_mp4(sample_count: int, *, moov_first: bool) -> bytes:
    """Encode a valid MP4 whose movie box grows with the sample count, before or after the payload."""
    ftyp = _box(b"ftyp", b"isom\0\0\2\0isomiso2mp41")
    sizes = np.full(sample_count, 10, dtype=np.int64)

    def moov_for(payload_start: int) -> bytes:
        stsd_body = struct.pack(">II", 0, 1) + struct.pack(">I4s", 16, b"avc1") + b"\0" * 8
        offsets = _stco([payload_start + 10 * i for i in range(sample_count)])
        stbl = _box(
            b"stbl",
            _box(b"stsd", stsd_body)
            + _stts(np.full(sample_count, 1000, dtype=np.int64))
            + _stsc_one_sample_per_chunk(sample_count)
            + _stsz(sizes)
            + offsets
            + _stss(np.array([1], dtype=np.int64)),
        )
        minf = _box(b"minf", _vmhd() + _dinf() + stbl)
        mdia = _box(b"mdia", _mdhd(1000, 1000 * sample_count) + _hdlr() + minf)
        trak = _box(b"trak", _tkhd(1, 1000 * sample_count, 64, 48) + mdia)
        return _box(b"moov", _mvhd(1000, 1000 * sample_count, 2) + trak)

    mdat = _box(b"mdat", b"x" * (10 * sample_count))
    if moov_first:
        moov_size = len(moov_for(0))
        return ftyp + moov_for(len(ftyp) + moov_size + 8) + mdat
    return ftyp + mdat + moov_for(len(ftyp) + 8)


def _counting_reader(source: bytes) -> tuple[list[tuple[int, int]], Callable[[str, int, int], bytes]]:
    calls: list[tuple[int, int]] = []

    def read_range(_path: str, offset: int, length: int) -> bytes:
        calls.append((offset, length))
        return source[offset : offset + length]

    return calls, read_range


def _assert_same_index(actual, expected) -> None:
    for field in dataclasses.fields(expected):
        got, want = getattr(actual, field.name), getattr(expected, field.name)
        if isinstance(want, np.ndarray):
            np.testing.assert_array_equal(got, want)
        else:
            assert got == want, field.name


@pytest.mark.parametrize("header_probe_bytes", [64, 4096, 1 << 20])
def test_header_probe_reads_a_large_faststart_moov_exactly(header_probe_bytes):
    source = _sized_mp4(20_000, moov_first=True)
    calls, read_range = _counting_reader(source)

    index = fetch_mp4_index("v.mp4", read_range, file_size=len(source), header_probe_bytes=header_probe_bytes)

    _assert_same_index(index, parse_mp4_index("v.mp4", source))
    assert len(calls) <= 2
    assert sum(length for _offset, length in calls) <= max(header_probe_bytes, index.mdat_payload_offset + 16)


def test_header_probe_reads_a_trailing_moov_from_the_tail_without_the_payload():
    source = _sized_mp4(200_000, moov_first=False)
    calls, read_range = _counting_reader(source)

    index = fetch_mp4_index("v.mp4", read_range, file_size=len(source), header_probe_bytes=4096)

    _assert_same_index(index, parse_mp4_index("v.mp4", source))
    payload_end = index.mdat_payload_offset + index.mdat_payload_size
    assert calls == [(0, 4096), (payload_end, len(source) - payload_end)]


def _fetch_episode(cache: EpisodeByteCache, episode_index: int) -> None:
    """Fetch every camera payload of an episode into the cache."""
    for camera_key in cache.manifest.video_keys:
        cache.get_bytes(episode_index, camera_key)


def _open_decoder(cache: EpisodeByteCache, episode_index: int, camera_key: str):
    """Open (or reuse) a cached decoder through the cache's own LRU path."""
    cache.retain_episode(episode_index)
    try:
        return cache._get_decoder_entry(episode_index, camera_key).decoder
    finally:
        cache.release_episode(episode_index)


def _fake_cache(
    monkeypatch,
    tmp_path,
    *,
    byte_budget=8,
    max_open_decoders=1,
    video_backend="torchcodec",
):
    manifest = EpisodeVideoManifest(video_keys=["camera"], files=[], spans={})
    monkeypatch.setattr(manifest, "episode_byte_size", lambda _episode: 5)
    cache = EpisodeByteCache(
        manifest,
        tmp_path,
        byte_budget=byte_budget,
        workers=1,
        max_open_decoders=max_open_decoders,
        video_backend=video_backend,
    )
    monkeypatch.setattr(
        cache,
        "_fetch_and_synthesize",
        lambda episode_index, _camera_key: {"bytes": bytes([episode_index]) * 5},
    )
    return cache


def test_byte_cache_does_not_evict_retained_episode(monkeypatch, tmp_path):
    with _fake_cache(monkeypatch, tmp_path, byte_budget=10) as cache:
        cache.retain_episode(0)
        _fetch_episode(cache, 0)
        _fetch_episode(cache, 1)
        _fetch_episode(cache, 2)

        assert (0, "camera") in cache._cache
        assert (1, "camera") not in cache._cache
        assert cache.resident_bytes <= cache.byte_budget


def test_byte_cache_rejects_retained_set_larger_than_budget(monkeypatch, tmp_path):
    with (
        _fake_cache(monkeypatch, tmp_path, byte_budget=4) as cache,
        pytest.raises(MemoryError, match="byte budget"),
    ):
        cache.retain_episode(0)


def test_decoder_count_has_independent_limit(monkeypatch, tmp_path):
    opened = []

    class FakeDecoder:
        pass

    def open_decoder(_data):
        decoder = FakeDecoder()
        opened.append(decoder)
        return decoder

    monkeypatch.setattr("lerobot.streaming.episode_cache.open_video_decoder", open_decoder)
    with _fake_cache(monkeypatch, tmp_path, byte_budget=20, max_open_decoders=1) as cache:
        first = _open_decoder(cache, 0, "camera")
        second = _open_decoder(cache, 1, "camera")

        assert first is not second
        assert cache.open_decoder_count == 1


def test_decoder_eviction_and_cache_shutdown_close_backend_resources(monkeypatch, tmp_path):
    opened = []

    class FakeDecoder:
        def __init__(self):
            self.closed = False

        def close(self):
            self.closed = True

    def open_decoder(_data):
        decoder = FakeDecoder()
        opened.append(decoder)
        return decoder

    monkeypatch.setattr("lerobot.streaming.episode_cache.open_video_decoder", open_decoder)
    with _fake_cache(monkeypatch, tmp_path, byte_budget=20, max_open_decoders=1) as cache:
        _open_decoder(cache, 0, "camera")
        _open_decoder(cache, 1, "camera")

        assert opened[0].closed
        assert not opened[1].closed

    assert opened[1].closed


def test_decoder_falls_back_to_pyav_when_torchcodec_rejects_mini_mp4(monkeypatch, tmp_path):
    opened_backends = []

    class FakeDecoder:
        pass

    def open_decoder(_data, *, backend="torchcodec"):
        opened_backends.append(backend)
        if backend == "torchcodec":
            raise ValueError("No valid stream found")
        return FakeDecoder()

    monkeypatch.setattr("lerobot.streaming.episode_cache.open_video_decoder", open_decoder)
    with _fake_cache(monkeypatch, tmp_path, video_backend="torchcodec") as cache:
        decoder = _open_decoder(cache, 0, "camera")

        assert isinstance(decoder, FakeDecoder)
        assert opened_backends == ["torchcodec", "pyav"]
        assert cache.decoder_fallback_count == 1


def test_torchcodec_frame_indices_are_clamped_to_decoder_bounds(monkeypatch, tmp_path):
    requested_indices = []

    class FakeDecoder:
        metadata = type("Metadata", (), {"average_fps": 30.0, "num_frames": 10})()

        def get_frames_at(self, *, indices):
            requested_indices.extend(indices)
            return SimpleNamespace(data=indices, pts_seconds=torch.tensor([index / 30 for index in indices]))

    with _fake_cache(monkeypatch, tmp_path) as cache:
        monkeypatch.setattr(
            cache.manifest,
            "lookup",
            lambda *_args: type("Span", (), {"source_start_pts": 0.0})(),
        )
        monkeypatch.setattr(cache, "_open_decoder", lambda *_args: FakeDecoder())

        cache.tolerance_s = 0.034
        cache.get_frames(0, "camera", [-0.02, 10 / 30])

    assert requested_indices == [0, 9]


@pytest.mark.parametrize("decoded_ts", [[0.12, 0.0, 0.12], [float("nan"), 0.0, 0.1]])
def test_torchcodec_rejects_out_of_tolerance_timestamps(
    monkeypatch: pytest.MonkeyPatch, tmp_path: Path, decoded_ts: list[float]
) -> None:
    """Average-fps frame selection must not silently return drifted or invalid PTS."""
    decoder = SimpleNamespace(
        metadata=SimpleNamespace(average_fps=30.0, num_frames=10),
        get_frames_at=lambda **_: SimpleNamespace(
            data=torch.zeros(3, 3, 2, 2, dtype=torch.uint8),
            pts_seconds=torch.tensor(decoded_ts, dtype=torch.float64),
        ),
    )
    with _fake_cache(monkeypatch, tmp_path) as cache:
        monkeypatch.setattr(cache.manifest, "lookup", lambda *_: SimpleNamespace(source_start_pts=10.0))
        monkeypatch.setattr(cache, "_open_decoder", lambda *_: decoder)
        with pytest.raises(ValueError, match="tolerance"):
            cache.get_frames(0, "camera", [10.1, 10.0, 10.1])
        assert not cache._retained_episodes


def test_torchcodec_timestamp_check_preserves_order_and_float64_precision(
    monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
    """Validate corresponding PTS without reordering duplicates or quantizing long timelines."""
    timestamps = [7200.00002, 0.0, 7200.00002]
    pixels = torch.arange(3, dtype=torch.uint8).view(3, 1, 1, 1)
    decoder = SimpleNamespace(
        metadata=SimpleNamespace(average_fps=30.0, num_frames=216002),
        get_frames_at=lambda **_: SimpleNamespace(
            data=pixels, pts_seconds=torch.tensor(timestamps, dtype=torch.float64)
        ),
    )
    with _fake_cache(monkeypatch, tmp_path) as cache:
        cache.tolerance_s = 1e-6
        monkeypatch.setattr(cache.manifest, "lookup", lambda *_: SimpleNamespace(source_start_pts=100.0))
        monkeypatch.setattr(cache, "_open_decoder", lambda *_: decoder)
        assert cache.get_frames(0, "camera", [value + 100 for value in timestamps]) is pixels


def test_torchcodec_does_not_hide_out_of_range_queries_by_clamping(
    monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
    """Clamping an estimated frame index must still enforce the original query tolerance."""
    decoder = SimpleNamespace(
        metadata=SimpleNamespace(average_fps=30.0, num_frames=10),
        get_frames_at=lambda **_: SimpleNamespace(data=None, pts_seconds=torch.tensor([0.3])),
    )
    with _fake_cache(monkeypatch, tmp_path) as cache:
        monkeypatch.setattr(cache.manifest, "lookup", lambda *_: SimpleNamespace(source_start_pts=0.0))
        monkeypatch.setattr(cache, "_open_decoder", lambda *_: decoder)
        with pytest.raises(ValueError, match="tolerance"):
            cache.get_frames(0, "camera", [1.0])


def test_frame_reads_serialize_access_to_each_decoder(monkeypatch, tmp_path):
    state_lock = threading.Lock()
    active = 0
    max_active = 0

    class FakeDecoder:
        metadata = type("Metadata", (), {"average_fps": 30.0, "num_frames": 10})()

        def get_frames_at(self, *, indices):
            nonlocal active, max_active
            with state_lock:
                active += 1
                max_active = max(max_active, active)
            try:
                time.sleep(0.01)
                return SimpleNamespace(
                    data=indices, pts_seconds=torch.tensor([index / 30 for index in indices])
                )
            finally:
                with state_lock:
                    active -= 1

    with _fake_cache(monkeypatch, tmp_path) as cache:
        monkeypatch.setattr(
            cache.manifest,
            "lookup",
            lambda *_args: type("Span", (), {"source_start_pts": 0.0})(),
        )
        decoder = FakeDecoder()
        monkeypatch.setattr(cache, "_open_decoder", lambda *_args: decoder)

        with ThreadPoolExecutor(max_workers=2) as executor:
            futures = [executor.submit(cache.get_frames, 0, "camera", [0.1]) for _ in range(2)]
            for future in futures:
                future.result()

    assert max_active == 1


def test_releasing_episode_allows_immediate_eviction(monkeypatch, tmp_path):
    with _fake_cache(monkeypatch, tmp_path, byte_budget=5) as cache:
        cache.retain_episode(0)
        _fetch_episode(cache, 0)
        cache.release_episode(0)
        _fetch_episode(cache, 1)

        assert (0, "camera") not in cache._cache
        assert (1, "camera") in cache._cache


def test_range_fetcher_closes_handles_from_all_worker_threads(tmp_path):
    (tmp_path / "video.mp4").write_bytes(b"0123456789")
    fetcher = ThreadLocalRangeFetcher(tmp_path)
    barrier = threading.Barrier(2)

    def read_from_worker(offset):
        barrier.wait()
        return fetcher.read_range("video.mp4", offset, 1)

    with ThreadPoolExecutor(max_workers=2) as pool:
        futures = [pool.submit(read_from_worker, offset) for offset in range(2)]
        assert [future.result() for future in futures] == [b"0", b"1"]

        handles = list(fetcher._all_handles.values())
    assert len(handles) == 2
    fetcher.close()

    assert not fetcher._all_handles
    assert all(handle.closed for handle in handles)
