# 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

"""Shared, read-only local representation of compressed MP4 sidecar arrays."""

from __future__ import annotations

import hashlib
import json
import os
import tempfile
from collections import deque
from collections.abc import Generator
from concurrent.futures import Future, ThreadPoolExecutor
from contextlib import closing
from pathlib import Path
from typing import Any, BinaryIO
from zipfile import BadZipFile, ZipFile

import numpy as np
from filelock import FileLock
from huggingface_hub.constants import HF_HOME
from numpy.typing import NDArray

from lerobot.streaming.mp4 import Mp4Index

ARRAY_NAMES = (
    "sample_pts",
    "sample_durations",
    "sample_composition_offsets",
    "sample_sizes",
    "sample_offsets",
    "sync_samples",
)
_MAGIC = b"LRIDX001"


def _content_digest(source: BinaryIO) -> str:
    """Digest the ZIP directory's member names, sizes and CRC-32s, leaving the position unchanged.

    Coarse filesystem timestamps and inode reuse can give a sidecar that is rewritten quickly the
    same stat identity; its member checksums change with the content. Only the central directory
    is read, never the compressed arrays.
    """
    position = source.tell()
    try:
        with ZipFile(source) as archive:
            members = sorted((info.filename, info.file_size, info.CRC) for info in archive.infolist())
    finally:
        source.seek(position)
    return hashlib.sha256(repr(members).encode()).hexdigest()


def _signature(stat: os.stat_result, content: str) -> tuple[int | str, ...]:
    """Identify a local sidecar generation from its file identity and its content digest."""
    return stat.st_dev, stat.st_ino, stat.st_size, stat.st_mtime_ns, stat.st_ctime_ns, content


def _cache_path(path: Path, signature: tuple[int | str, ...] | None = None) -> Path:
    """Choose a generation-specific path for the derived read-only index."""
    if signature is None:
        with path.open("rb") as source:
            signature = _signature(os.fstat(source.fileno()), _content_digest(source))
    digest = hashlib.sha256(repr(signature).encode()).hexdigest()[:24]
    cache_root = Path(os.environ.get("HF_LEROBOT_HOME", str(Path(HF_HOME) / "lerobot"))).expanduser()
    return cache_root / "streaming-indexes" / f"mp4-{digest}.bin"


def _read_metadata(path: Path) -> dict[str, Any]:
    """Read the index footer and validate array bounds before mapping."""
    with path.open("rb") as source:
        size = os.fstat(source.fileno()).st_size
        source.seek(-16, os.SEEK_END)
        length = int.from_bytes(source.read(8), "little")
        if source.read(8) != _MAGIC or not 0 < length <= size - 16:
            raise ValueError("Invalid mapped MP4 index footer")
        source.seek(size - 16 - length)
        payload = json.loads(source.read(length))
    if (
        not isinstance(payload, dict)
        or payload.get("version") != 3
        or not isinstance(payload.get("sidecar"), dict)
        or not isinstance(payload.get("files"), list)
    ):
        raise ValueError("Invalid mapped MP4 index schema")
    for item in payload["files"]:
        if not isinstance(item, dict) or not isinstance(item.get("arrays"), dict):
            raise ValueError("Invalid mapped MP4 index record")
        for name in ARRAY_NAMES:
            offset, count, dtype_string = item["arrays"][name]
            dtype = np.dtype(dtype_string)
            if (
                type(offset) is not int
                or type(count) is not int
                or dtype.kind not in "iuf"
                or count < 0
                or offset < 0
                or offset % dtype.alignment
                or offset + count * dtype.itemsize > size - 16 - length
            ):
                raise ValueError("Invalid mapped MP4 index array")
        counts = {item["arrays"][name][1] for name in ARRAY_NAMES if name != "sync_samples"}
        if len(counts) != 1:
            raise ValueError("Inconsistent MP4 sample counts")
    return payload


def sidecar_payload(path: Path) -> dict[str, Any]:
    """Read identity without decompressing arrays or preparing a derived cache."""
    try:
        return _read_metadata(_cache_path(path))
    except (OSError, ValueError, KeyError, TypeError, BadZipFile):
        pass
    with np.load(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 {path}")
    return payload


def validate_source_arrays(path: Path, payload: dict[str, Any]) -> None:
    """Check temporary NPZ contents without retaining arrays or creating a cache."""
    with ZipFile(path) as archive:
        for file_index, item in enumerate(payload["files"]):
            _read_arrays(archive, file_index, item)


def _read_arrays(archive: ZipFile, file_index: int, item: dict[str, Any]) -> dict[str, NDArray[np.generic]]:
    """Read numeric members directly, without NpzFile's linear filename searches."""
    arrays = {}
    for name in ARRAY_NAMES:
        # Older Python ZIP readers seek relative to the shared descriptor while
        # opening ZIP64 headers. Hold their own source lock across that sequence
        # so other headers or payload reads cannot move it; decompress unlocked.
        with archive._lock:  # type: ignore[attr-defined]  # CPython ZipFile's shared seek lock.
            member = archive.open(f"{file_index}/{name}.npy")
        with member:
            arrays[name] = np.lib.format.read_array(member, allow_pickle=False)
    _validate_arrays(arrays)
    Mp4Index.from_dict(item["mp4"], arrays)
    return arrays


def _iter_arrays(
    archive: ZipFile,
    files: list[dict[str, Any]],
    *,
    workers: int,
    max_pending_bytes: int = 256 * 1024**2,
) -> Generator[dict[str, NDArray[np.generic]]]:
    """Decompress in parallel with ordered, count- and byte-bounded read-ahead.

    ZIP uncompressed member sizes bound admitted array bytes. An oversized record
    runs alone. ZipFile serializes shared source seeks, but decompression overlaps;
    sharing its directory avoids one large member table per worker process.
    """
    if workers < 1 or max_pending_bytes < 1:
        raise ValueError("workers and max_pending_bytes must be positive")
    if workers == 1:
        for index, item in enumerate(files):
            yield _read_arrays(archive, index, item)
        return
    pending: deque[tuple[Future[dict[str, NDArray[np.generic]]], int]] = deque()
    pending_bytes = 0
    next_index = 0
    with ThreadPoolExecutor(max_workers=workers, thread_name_prefix="sidecar-index") as executor:
        try:
            while next_index < len(files) or pending:
                while next_index < len(files) and len(pending) < workers:
                    size = sum(archive.getinfo(f"{next_index}/{name}.npy").file_size for name in ARRAY_NAMES)
                    if pending and pending_bytes + size > max_pending_bytes:
                        break
                    pending.append(
                        (executor.submit(_read_arrays, archive, next_index, files[next_index]), size)
                    )
                    pending_bytes += size
                    next_index += 1
                future, size = pending.popleft()
                yield future.result()
                del future
                pending_bytes -= size
        finally:
            for future, _ in pending:
                future.cancel()


def _validate_arrays(arrays: dict[str, NDArray[np.generic]]) -> None:
    """Check numeric one-dimensional arrays and consistent sample counts."""
    for name, array in arrays.items():
        if array.ndim != 1 or array.dtype.kind not in "iuf":
            raise ValueError(f"Invalid MP4 sample array: {name}")
    if len({len(array) for name, array in arrays.items() if name != "sync_samples"}) != 1:
        raise ValueError("Inconsistent MP4 sample counts")


def mapped_sidecar(path: Path, *, workers: int = 4) -> tuple[Path, dict[str, Any]]:
    """Convert once under a lock, then reuse an immutable memory-mappable index.

    Only index metadata is cached, never video payloads. Each source generation gets
    a different file: replacing a sidecar cannot invalidate live array views.
    """
    if workers < 1:
        raise ValueError("workers must be positive")
    # Shared filesystems can retain stale pathname attributes until the file is
    # opened. Identify and read the same descriptor, including across the lock wait.
    with path.open("rb") as source:
        content = _content_digest(source)
        signature = _signature(os.fstat(source.fileno()), content)
        destination = _cache_path(path, signature)

        def cached() -> dict[str, Any] | None:
            """Return valid cached metadata or signal that conversion is needed."""
            try:
                return _read_metadata(destination)
            except (OSError, ValueError, KeyError, TypeError):
                return None

        payload = cached()
        if payload is not None:
            return destination, payload
        destination.parent.mkdir(parents=True, exist_ok=True)
        with FileLock(str(destination) + ".lock", timeout=30 * 60):
            payload = cached()
            if payload is not None:
                return destination, payload
            temporary: Path | None = None
            try:
                with np.load(source, allow_pickle=False) as data:
                    # Workers share this descriptor, so re-check only its stat identity here.
                    if _signature(os.fstat(source.fileno()), content) != signature:
                        raise OSError("MP4 sidecar changed before index conversion")
                    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 {path}")
                    with tempfile.NamedTemporaryFile(
                        dir=destination.parent, suffix=".index.tmp", delete=False
                    ) as out:
                        temporary = Path(out.name)
                        with closing(_iter_arrays(data.zip, payload["files"], workers=workers)) as records:
                            items = iter(payload["files"])
                            for arrays in records:
                                item = next(items)
                                item["arrays"] = {}
                                for name, array in arrays.items():
                                    out.write(b"\0" * (-out.tell() % 8))
                                    item["arrays"][name] = [out.tell(), array.size, array.dtype.str]
                                    out.write(array.tobytes())
                                # zip/enumerate retain a tuple with the previous arrays.
                                del arrays, array
                        if _signature(os.fstat(source.fileno()), content) != signature:
                            raise OSError("MP4 sidecar changed during index conversion")
                        metadata = json.dumps(payload, separators=(",", ":")).encode()
                        out.write(metadata)
                        out.write(len(metadata).to_bytes(8, "little"))
                        out.write(_MAGIC)
                        out.flush()
                        os.fsync(out.fileno())
                os.replace(temporary, destination)
            finally:
                if temporary is not None:
                    temporary.unlink(missing_ok=True)
        return destination, payload


def install_sidecar(source: Path, destination: Path) -> None:
    """Atomically install a validated sidecar and retain its prepared mapped index.

    The caller holds the sidecar lock. Renaming changes inode ctime, so transfer
    the derived file to the post-rename generation instead of decompressing again.
    """
    prepared, _ = mapped_sidecar(source)
    with source.open("rb") as handle:
        content = _content_digest(handle)
        before = os.fstat(handle.fileno())
        if _cache_path(source, _signature(before, content)) != prepared:
            raise OSError("MP4 sidecar changed before installation")
        os.replace(source, destination)
        after = os.fstat(handle.fileno())
        if (before.st_size, before.st_mtime_ns) != (after.st_size, after.st_mtime_ns):
            raise OSError("MP4 sidecar changed during installation")
        installed = _cache_path(destination, _signature(after, content))
        if installed != prepared:
            with FileLock(str(installed) + ".lock", timeout=30 * 60):
                os.replace(prepared, installed)


def mapped_arrays(buffer: np.memmap[Any, Any], item: dict[str, Any]) -> dict[str, NDArray[np.generic]]:
    """Return views whose base retains the single shared read-only mapping."""
    return {
        name: np.ndarray((count,), dtype=np.dtype(dtype), buffer=buffer, offset=offset)
        for name, (offset, count, dtype) in item["arrays"].items()
    }
