# 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

"""Episode-to-MP4 byte-span manifests and sidecar serialization."""

from __future__ import annotations

import errno
import json
import logging
from collections import defaultdict
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any
from zipfile import BadZipFile

import numpy as np
from numpy.typing import NDArray

from lerobot.streaming.mp4 import (
    DEFAULT_HEADER_PROBE_BYTES,
    Mp4Index,
    Mp4SampleSlice,
    fetch_mp4_index,
    synthesized_mp4_size,
)
from lerobot.streaming.range_fetch import make_range_fetcher
from lerobot.streaming.sidecar_utils import (
    mapped_arrays,
    mapped_sidecar,
    sidecar_payload,
    validate_source_arrays,
)

if TYPE_CHECKING:
    from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
    from lerobot.streaming.sidecar import SidecarSpec


@dataclass(frozen=True)
class EpisodeVideoSpan:
    """Indexed media span for one episode and camera."""

    file_id: int
    mdat_offset: int
    mdat_length: int
    first_pts: float
    last_pts: float
    frame_count: int
    sample_lo: int
    sample_hi: int
    source_start_pts: float


@dataclass(frozen=True)
class VideoFileRecord:
    """Source file metadata and parsed MP4 sample index."""

    file_path: str
    file_size: int
    mp4: Mp4Index


def video_file_groups(
    meta: LeRobotDatasetMetadata, episode_indices: Sequence[int] | None = None
) -> dict[str, list[tuple[int, int]]]:
    """Group episode/camera positions by source, reading only path metadata columns."""
    meta.ensure_readable()
    indices = range(int(meta.total_episodes)) if episode_indices is None else episode_indices
    columns = [f"videos/{key}/{field}" for key in meta.video_keys for field in ("chunk_index", "file_index")]
    if not columns:
        return {}
    video_path = meta.video_path
    if video_path is None:
        raise ValueError("Video features require a video_path template in the dataset metadata")
    table = meta.episodes.select_columns(columns).with_format(None)
    if episode_indices is not None:
        table = table.select(indices)
    groups: dict[str, list[tuple[int, int]]] = defaultdict(list)
    offset = 0
    for batch in table.iter(batch_size=10_000):
        count = len(batch[columns[0]])
        for camera, key in enumerate(meta.video_keys):
            chunks = batch[f"videos/{key}/chunk_index"]
            files = batch[f"videos/{key}/file_index"]
            for position, (chunk, file) in enumerate(zip(chunks, files, strict=True)):
                path = str(Path(video_path.format(video_key=key, chunk_index=chunk, file_index=file)))
                groups[path].append((int(indices[offset + position]), camera))
        offset += count
    return dict(groups)


class EpisodeVideoManifest:
    """Map episode-camera pairs to source byte spans and MP4 metadata."""

    def __init__(
        self,
        *,
        video_keys: list[str],
        files: list[VideoFileRecord],
        spans: dict[str, NDArray[np.generic]],
    ) -> None:
        """Store video keys, indexed source files, and episode span arrays."""
        self.video_keys = list(video_keys)
        self._camera_to_id = {key: idx for idx, key in enumerate(self.video_keys)}
        self.files = files
        self.spans = spans
        self._episode_byte_sizes: dict[int, int] = {}

    @classmethod
    def build(
        cls,
        meta: LeRobotDatasetMetadata,
        data_root: str | Path,
        *,
        episode_indices: list[int] | range | None = None,
        range_backend: str = "fsspec",
        workers: int = 8,
        header_probe_bytes: int = DEFAULT_HEADER_PROBE_BYTES,
        max_probe_bytes: int = 64 * 1024 * 1024,
        keyframe_pad_s: float = 0.1,
        keyframe_pad_fraction: float = 0.05,
        sidecar_path: str | Path | None = None,
        token: str | bool | None = None,
    ) -> EpisodeVideoManifest:
        """Build episode spans from source MP4s or a validated file sidecar."""
        meta.ensure_readable()
        video_keys = list(meta.video_keys)
        if episode_indices is None:
            episode_indices = range(int(meta.total_episodes))
        file_episodes = video_file_groups(meta, episode_indices)
        rel_paths = sorted(file_episodes)
        if sidecar_path is None:
            files = cls._build_file_records(
                rel_paths,
                data_root,
                range_backend=range_backend,
                workers=workers,
                header_probe_bytes=header_probe_bytes,
                max_probe_bytes=max_probe_bytes,
                token=token,
            )
        else:
            records = cls.load_file_sidecar(sidecar_path, file_paths=rel_paths)
            missing = [path for path in rel_paths if path not in records]
            if missing:
                raise ValueError(
                    f"Sidecar {sidecar_path} is missing {len(missing)} files, first: {missing[0]}"
                )
            files = [records[path] for path in rel_paths]

        total = int(meta.total_episodes)
        num_cameras = len(video_keys)
        spans: dict[str, NDArray[np.generic]] = {
            "file_id": np.zeros((total, num_cameras), dtype=np.int32),
            "mdat_offset": np.zeros((total, num_cameras), dtype=np.int64),
            "mdat_length": np.zeros((total, num_cameras), dtype=np.int64),
            "first_pts": np.zeros((total, num_cameras), dtype=np.float64),
            "last_pts": np.zeros((total, num_cameras), dtype=np.float64),
            "frame_count": np.zeros((total, num_cameras), dtype=np.int32),
            "sample_lo": np.zeros((total, num_cameras), dtype=np.int32),
            "sample_hi": np.zeros((total, num_cameras), dtype=np.int32),
            "source_start_pts": np.zeros((total, num_cameras), dtype=np.float64),
        }

        # Finish each source's mapped pages before moving on. Rank-balanced episode
        # order can revisit most files and thrash an index larger than physical RAM.
        timestamp_columns = [
            f"videos/{key}/{field}" for key in video_keys for field in ("from_timestamp", "to_timestamp")
        ]
        timestamps = {
            key: np.asarray(values, dtype=np.float64)
            for key, values in meta.episodes.select_columns(timestamp_columns).with_format(None)[:].items()
        }

        manifest = cls(video_keys=video_keys, files=files, spans=spans)
        for file_id, path in enumerate(rel_paths):
            mp4 = files[file_id].mp4
            for ep_idx, cam_idx in file_episodes.pop(path):
                key = video_keys[cam_idx]
                from_ts = float(timestamps[f"videos/{key}/from_timestamp"][ep_idx])
                to_ts = float(timestamps[f"videos/{key}/to_timestamp"][ep_idx])
                sample_slice = mp4.sample_slice(
                    from_ts,
                    to_ts,
                    keyframe_pad_s=keyframe_pad_s,
                    keyframe_pad_fraction=keyframe_pad_fraction,
                    file_size=files[file_id].file_size,
                )
                spans["file_id"][ep_idx, cam_idx] = file_id
                spans["mdat_offset"][ep_idx, cam_idx] = sample_slice.byte_offset
                spans["mdat_length"][ep_idx, cam_idx] = sample_slice.byte_length
                spans["first_pts"][ep_idx, cam_idx] = from_ts
                spans["last_pts"][ep_idx, cam_idx] = to_ts
                spans["frame_count"][ep_idx, cam_idx] = sample_slice.sample_hi - sample_slice.sample_lo + 1
                spans["sample_lo"][ep_idx, cam_idx] = sample_slice.sample_lo
                spans["sample_hi"][ep_idx, cam_idx] = sample_slice.sample_hi
                spans["source_start_pts"][ep_idx, cam_idx] = sample_slice.source_start_pts
                manifest._episode_byte_sizes[ep_idx] = manifest._episode_byte_sizes.get(
                    ep_idx, 0
                ) + synthesized_mp4_size(mp4, sample_slice)

        return manifest

    @staticmethod
    def _build_file_records(
        rel_paths: list[str],
        data_root: str | Path,
        *,
        range_backend: str,
        workers: int,
        header_probe_bytes: int,
        max_probe_bytes: int,
        token: str | bool | None,
    ) -> list[VideoFileRecord]:
        """Index unique source files concurrently and return records sorted by path."""
        fetcher = make_range_fetcher(
            data_root,
            range_backend=range_backend,
            workers=workers,
            token=token,
        )

        def build_file(path: str) -> VideoFileRecord:
            """Resolve a source size and parse its MP4 sample tables."""
            file_size = fetcher.info_size(path)
            mp4 = fetch_mp4_index(
                path,
                fetcher.read_range,
                file_size=file_size,
                header_probe_bytes=header_probe_bytes,
                max_probe_bytes=max_probe_bytes,
            )
            return VideoFileRecord(path, file_size, mp4)

        try:
            with ThreadPoolExecutor(max_workers=workers) as pool:
                futures = {pool.submit(build_file, path): path for path in rel_paths}
                records = []
                progress_interval = max(1, len(futures) // 20)
                for completed, future in enumerate(as_completed(futures), start=1):
                    records.append(future.result())
                    if completed == len(futures) or completed % progress_interval == 0:
                        logging.info("Indexed %d/%d MP4 files for streaming sidecar", completed, len(futures))
                return sorted(records, key=lambda record: record.file_path)
        finally:
            fetcher.close()

    @classmethod
    def write_file_sidecar(
        cls,
        sidecar_path: str | Path,
        rel_paths: list[str],
        data_root: str | Path,
        *,
        spec: SidecarSpec,
        range_backend: str = "native-http",
        workers: int = 8,
        header_probe_bytes: int = DEFAULT_HEADER_PROBE_BYTES,
        max_probe_bytes: int = 64 * 1024 * 1024,
        token: str | bool | None = None,
    ) -> None:
        """Index source files and write a reusable sidecar."""
        records = cls._build_file_records(
            sorted(set(rel_paths)),
            data_root,
            range_backend=range_backend,
            workers=workers,
            header_probe_bytes=header_probe_bytes,
            max_probe_bytes=max_probe_bytes,
            token=token,
        )
        cls.save_file_sidecar(sidecar_path, records, spec=spec)

    @staticmethod
    def save_file_sidecar(
        sidecar_path: str | Path,
        records: list[VideoFileRecord],
        *,
        spec: SidecarSpec,
    ) -> None:
        """Serialize source records and sample arrays to an NPZ sidecar."""
        sidecar_path = Path(sidecar_path)
        sidecar_path.parent.mkdir(parents=True, exist_ok=True)
        payload = {
            "version": 3,
            "sidecar": spec.with_source_files(
                tuple((record.file_path, record.file_size) for record in records)
            ).to_dict(),
            "files": [
                {"file_path": record.file_path, "file_size": record.file_size, "mp4": record.mp4.to_dict()}
                for record in records
            ],
        }
        arrays: dict[str, Any] = {}
        for file_idx, record in enumerate(records):
            arrays[f"{file_idx}/sample_pts"] = record.mp4.sample_pts
            arrays[f"{file_idx}/sample_durations"] = record.mp4.sample_durations
            arrays[f"{file_idx}/sample_composition_offsets"] = record.mp4.sample_composition_offsets
            arrays[f"{file_idx}/sample_sizes"] = record.mp4.sample_sizes
            arrays[f"{file_idx}/sample_offsets"] = record.mp4.sample_offsets
            arrays[f"{file_idx}/sync_samples"] = record.mp4.sync_samples
        np.savez_compressed(sidecar_path, manifest_json=json.dumps(payload).encode("utf-8"), **arrays)

    @staticmethod
    def load_file_sidecar_metadata(sidecar_path: str | Path) -> dict[str, Any]:
        """Load and validate the sidecar identity metadata."""
        with np.load(sidecar_path, allow_pickle=False) as data:
            payload = json.loads(bytes(data["manifest_json"]).decode("utf-8"))
        if payload.get("version") != 3 or not isinstance(payload.get("sidecar"), dict):
            raise ValueError(f"Unsupported MP4 sidecar schema in {sidecar_path}")
        return payload["sidecar"]

    @staticmethod
    def validate_file_sidecar(
        sidecar_path: str | Path, spec: SidecarSpec, *, prepare_cache: bool = True
    ) -> bool:
        """Return whether a sidecar matches its expected source specification."""
        try:
            from lerobot.streaming.sidecar import SidecarSpec

            path = Path(sidecar_path).expanduser()
            payload = sidecar_payload(path)
            candidate = SidecarSpec.from_dict(payload["sidecar"])
            if not spec.matches(candidate):
                return False
            expected = dict(candidate.source_files)
            actual = {item["file_path"]: int(item["file_size"]) for item in payload["files"]}
            if actual != expected:
                return False
            if prepare_cache:
                mapped_sidecar(path)
            else:
                validate_source_arrays(path, payload)
        except OSError as exc:
            if exc.errno in (errno.EACCES, errno.EROFS, errno.ENOSPC, errno.EDQUOT, errno.ENOTDIR):
                raise OSError(
                    exc.errno,
                    "Cannot access the MP4 index cache; check permissions and free space in HF_LEROBOT_HOME",
                ) from exc
            return False
        except (ValueError, KeyError, TypeError, BadZipFile, EOFError):
            return False

        return True

    @staticmethod
    def load_file_sidecar(
        sidecar_path: str | Path, *, file_paths: Sequence[str] | None = None
    ) -> dict[str, VideoFileRecord]:
        """Load selected source records with shared, read-only file-backed arrays."""
        path, payload = mapped_sidecar(Path(sidecar_path).expanduser())
        selected = None if file_paths is None else set(file_paths)
        buffer = np.memmap(path, mode="r", dtype=np.uint8)
        records = {}
        for item in payload["files"]:
            if selected is not None and item["file_path"] not in selected:
                continue
            mp4 = Mp4Index.from_dict(item["mp4"], mapped_arrays(buffer, item))
            records[item["file_path"]] = VideoFileRecord(item["file_path"], int(item["file_size"]), mp4)
        return records

    def camera_id(self, camera_key: str) -> int:
        """Return the dense manifest index for a camera key."""
        return self._camera_to_id[camera_key]

    def lookup(self, episode_index: int, camera_key: str) -> EpisodeVideoSpan:
        """Return the indexed source span for an episode camera."""
        cam = self.camera_id(camera_key)
        return EpisodeVideoSpan(
            file_id=int(self.spans["file_id"][episode_index, cam]),
            mdat_offset=int(self.spans["mdat_offset"][episode_index, cam]),
            mdat_length=int(self.spans["mdat_length"][episode_index, cam]),
            first_pts=float(self.spans["first_pts"][episode_index, cam]),
            last_pts=float(self.spans["last_pts"][episode_index, cam]),
            frame_count=int(self.spans["frame_count"][episode_index, cam]),
            sample_lo=int(self.spans["sample_lo"][episode_index, cam]),
            sample_hi=int(self.spans["sample_hi"][episode_index, cam]),
            source_start_pts=float(self.spans["source_start_pts"][episode_index, cam]),
        )

    def file_lookup(self, file_id: int) -> VideoFileRecord:
        """Return a source record by dense file identifier."""
        return self.files[file_id]

    def mp4_index(self, episode_index: int, camera_key: str) -> Mp4Index:
        """Return the source MP4 index for an episode camera."""
        return self.files[self.lookup(episode_index, camera_key).file_id].mp4

    def sample_slice(self, episode_index: int, camera_key: str) -> Mp4SampleSlice:
        """Return the sample-table slice for an episode camera."""
        span = self.lookup(episode_index, camera_key)
        return 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,
        )

    def episode_byte_size(self, episode_index: int) -> int:
        """Exact synthesized video bytes retained while an episode is active."""
        if episode_index in self._episode_byte_sizes:
            return self._episode_byte_sizes[episode_index]
        return sum(
            synthesized_mp4_size(
                self.mp4_index(episode_index, camera_key),
                self.sample_slice(episode_index, camera_key),
            )
            for camera_key in self.video_keys
        )
