# 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

"""Revision-safe lifecycle for locally cached MP4 byte-index sidecars."""

from __future__ import annotations

import hashlib
import json
import logging
import re
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from uuid import uuid4

from filelock import FileLock, Timeout

from lerobot.streaming.location import StorageLocation
from lerobot.streaming.manifest import EpisodeVideoManifest
from lerobot.streaming.sidecar_utils import install_sidecar

SIDECAR_SCHEMA_VERSION = 3


class SidecarLockTimeoutError(TimeoutError):
    """Raised when another process does not finish a sidecar build in time."""


@dataclass(frozen=True)
class SidecarSpec:
    """Identity and source-file contract for an MP4 index sidecar."""

    repo_id: str
    revision: str
    data_root: str
    source_files: tuple[tuple[str, int | None], ...]
    schema_version: int = SIDECAR_SCHEMA_VERSION
    source_fingerprints: tuple[tuple[str, str], ...] = ()

    def __post_init__(self) -> None:
        """Validate and normalize the immutable source-file list."""
        if not self.repo_id:
            raise ValueError("repo_id must not be empty")
        if not self.revision:
            raise ValueError("revision must not be empty")
        normalized = tuple(
            sorted((str(path), None if size is None else int(size)) for path, size in self.source_files)
        )
        if any(not path or size is not None and size < 0 for path, size in normalized):
            raise ValueError("source file paths must be non-empty and sizes must be non-negative")
        object.__setattr__(self, "source_files", normalized)
        fingerprints = tuple(sorted(self.source_fingerprints))
        if fingerprints and {path for path, _value in fingerprints} != {path for path, _ in normalized}:
            raise ValueError("Source fingerprints must cover every source file")
        if any(not value for _path, value in fingerprints):
            raise ValueError("Source fingerprints must not be empty")
        object.__setattr__(self, "source_fingerprints", fingerprints)

    def to_dict(self) -> dict[str, object]:
        """Serialize the sidecar specification."""
        return {
            "schema_version": self.schema_version,
            "repo_id": self.repo_id,
            "revision": self.revision,
            "data_root": self.data_root,
            "source_files": [{"path": path, "size": size} for path, size in self.source_files],
            "source_fingerprints": dict(self.source_fingerprints),
        }

    @classmethod
    def from_dict(cls, data: Mapping[str, Any]) -> SidecarSpec:
        """Parse and validate a serialized sidecar specification."""
        source_files = data.get("source_files")
        if not isinstance(source_files, list):
            raise ValueError("MP4 sidecar source_files must be a list")
        parsed_files: list[tuple[str, int | None]] = []
        fingerprints = data.get("source_fingerprints", {})
        if not isinstance(fingerprints, dict) or any(
            not isinstance(path, str) or not isinstance(value, str) for path, value in fingerprints.items()
        ):
            raise ValueError("Invalid MP4 source fingerprints")
        for item in source_files:
            if not isinstance(item, dict) or not isinstance(item.get("path"), str):
                raise ValueError("Invalid MP4 sidecar source file entry")
            size = item.get("size")
            parsed_files.append((item["path"], None if size is None else int(size)))
        return cls(
            schema_version=int(data["schema_version"]),
            repo_id=str(data["repo_id"]),
            revision=str(data["revision"]),
            data_root=str(data["data_root"]),
            source_files=tuple(parsed_files),
            source_fingerprints=tuple(fingerprints.items()),
        )

    def with_source_files(self, source_files: tuple[tuple[str, int], ...]) -> SidecarSpec:
        """Return a copy with resolved source sizes."""
        return SidecarSpec(
            repo_id=self.repo_id,
            revision=self.revision,
            data_root=self.data_root,
            source_files=source_files,
            schema_version=self.schema_version,
            source_fingerprints=self.source_fingerprints,
        )

    def matches(self, candidate: SidecarSpec) -> bool:
        """Return whether a candidate satisfies this expected specification."""
        # Tags/branches may name the same immutable payload snapshot. Other roots still
        # require the metadata revision to match; source paths/sizes/fingerprints always do.
        pinned_repository = StorageLocation.parse(self.data_root).pinned_commit is not None
        if (
            self.schema_version != candidate.schema_version
            or self.repo_id != candidate.repo_id
            or (self.revision != candidate.revision and not pinned_repository)
            or self.data_root != candidate.data_root
            or self.source_fingerprints != candidate.source_fingerprints
        ):
            return False
        expected = dict(self.source_files)
        actual = dict(candidate.source_files)
        if expected.keys() != actual.keys():
            return False
        return all(size is None or actual[path] == size for path, size in expected.items())


SidecarBuilder = Callable[[Path, SidecarSpec], None]
SidecarDownloader = Callable[[Path, SidecarSpec], bool]


def sidecar_cache_path(cache_root: str | Path, spec: SidecarSpec) -> Path:
    """Derive a revision-keyed local cache path for a sidecar."""
    identity = json.dumps(
        {
            "schema_version": spec.schema_version,
            "repo_id": spec.repo_id,
            "revision": spec.revision,
            "data_root": spec.data_root,
            **({"source_fingerprints": dict(spec.source_fingerprints)} if spec.source_fingerprints else {}),
        },
        sort_keys=True,
        separators=(",", ":"),
    )
    digest = hashlib.sha256(identity.encode()).hexdigest()[:16]
    repo_slug = re.sub(r"[^A-Za-z0-9_.-]+", "--", spec.repo_id).strip("-") or "dataset"
    revision_slug = re.sub(r"[^A-Za-z0-9_.-]+", "-", spec.revision).strip("-")[:32] or "revision"
    return (
        Path(cache_root).expanduser() / repo_slug / f"mp4-v{spec.schema_version}-{revision_slug}-{digest}.npz"
    )


def ensure_mp4_sidecar(
    spec: SidecarSpec,
    cache_root: str | Path,
    *,
    build: SidecarBuilder,
    download: SidecarDownloader | None = None,
    lock_timeout_s: float = 30 * 60,
) -> Path:
    """Return a valid local sidecar, downloading or building it exactly once.

    This function never uploads. ``download`` and ``build`` must write only to the temporary path
    provided to them; a validated file becomes visible at the cache path through ``os.replace``.
    """
    destination = sidecar_cache_path(cache_root, spec)
    if EpisodeVideoManifest.validate_file_sidecar(destination, spec):
        return destination

    destination.parent.mkdir(parents=True, exist_ok=True)
    lock_path = destination.with_suffix(f"{destination.suffix}.lock")
    try:
        with FileLock(lock_path, timeout=lock_timeout_s):
            if EpisodeVideoManifest.validate_file_sidecar(destination, spec):
                return destination

            temporary = destination.parent / f".{destination.name}.{uuid4().hex}.tmp.npz"
            try:
                if download is not None:
                    logging.info("Looking for published MP4 sidecar for %s@%s", spec.repo_id, spec.revision)
                    if download(temporary, spec) and EpisodeVideoManifest.validate_file_sidecar(
                        temporary, spec
                    ):
                        install_sidecar(temporary, destination)
                        return destination
                    temporary.unlink(missing_ok=True)

                logging.info("Building MP4 sidecar for %s@%s", spec.repo_id, spec.revision)
                build(temporary, spec)
                if not EpisodeVideoManifest.validate_file_sidecar(temporary, spec):
                    raise ValueError("Built MP4 sidecar failed revision and source validation")
                install_sidecar(temporary, destination)
                return destination
            finally:
                temporary.unlink(missing_ok=True)
    except Timeout as exc:
        raise SidecarLockTimeoutError(
            f"Timed out waiting {lock_timeout_s:g}s for MP4 sidecar lock {lock_path}"
        ) from exc
