# 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

"""Pure episode-scoped Parquet reads for training-time dataset streaming."""

from __future__ import annotations

import posixpath
import threading
import time
from collections.abc import Sequence
from pathlib import Path
from typing import BinaryIO

import fsspec
import pyarrow as pa
import pyarrow.compute as pc
import pyarrow.parquet as pq

from lerobot.streaming.location import StorageLocation


class EpisodeParquetReader:
    """Read complete episodes with column projection from local or fsspec roots."""

    def __init__(
        self,
        data_root: str | Path,
        *,
        columns: Sequence[str],
        token: str | bool | None = None,
        max_retries: int = 4,
        retry_backoff_s: float = 0.05,
    ) -> None:
        """Configure projected episode reads from a local or fsspec root."""
        if not columns:
            raise ValueError("EpisodeParquetReader requires at least one projected column")
        self.columns = tuple(dict.fromkeys(columns))
        self._read_columns = (
            self.columns if "episode_index" in self.columns else (*self.columns, "episode_index")
        )
        data_root_str = str(data_root)
        if max_retries < 0:
            raise ValueError("max_retries must be non-negative")
        if retry_backoff_s < 0:
            raise ValueError("retry_backoff_s must be non-negative")
        self._max_retries = max_retries
        self._retry_backoff_s = retry_backoff_s
        location = StorageLocation.parse(data_root_str)
        self._open_lock = threading.Lock() if location.is_hf else None
        self._filesystem, self._root_path = fsspec.core.url_to_fs(
            data_root_str, **location.storage_options(token)
        )

    def read_episode(
        self,
        relative_path: str | Path,
        *,
        episode_index: int,
        expected_rows: int,
    ) -> pa.Table:
        """Read and validate one complete episode with column projection.

        Args:
            relative_path (`str | Path`):
                Parquet path relative to this reader's data root.
            episode_index (`int`):
                Episode identifier used for row-group pruning and row filtering.
            expected_rows (`int`):
                Positive episode length from the dataset metadata.

        Returns:
            `pyarrow.Table`: The selected episode's projected columns in source row order.

        Raises:
            ValueError: If required columns are absent or the returned episode is incomplete.

        Note:
            Missing or overlapping row-group statistics require broader projected reads before
            filtering. Only the validated complete episode is retained by the caller.
        """
        if expected_rows <= 0:
            raise ValueError(f"Episode {episode_index} must contain at least one row")

        path = posixpath.join(self._root_path.rstrip("/"), str(relative_path).lstrip("/"))
        with self._open_with_retry(path) as source:
            parquet = pq.ParquetFile(source)
            available = set(parquet.schema_arrow.names)
            missing = sorted(set(self._read_columns) - available)
            if missing:
                raise ValueError(f"Parquet file {relative_path} is missing projected columns: {missing}")
            table = parquet.read_row_groups(
                self._candidate_row_groups(parquet, episode_index), columns=list(self._read_columns)
            )

        # Statistics only exclude impossible groups. An episode can cross pure
        # and mixed groups, and groups without bounds must remain candidates.
        table = table.filter(pc.equal(table.column("episode_index"), episode_index))
        self._validate_complete_episode(table, episode_index, expected_rows, relative_path)
        if "episode_index" not in self.columns:
            table = table.drop_columns(["episode_index"])
        return table

    def _open_with_retry(self, path: str) -> BinaryIO:
        """Open a binary source, retrying stale Hub directory-cache misses."""
        for attempt in range(self._max_retries + 1):
            try:
                if self._open_lock is None:
                    return self._filesystem.open(path, "rb")
                with self._open_lock:
                    return self._filesystem.open(path, "rb")
            except FileNotFoundError:
                if attempt == self._max_retries:
                    raise
                self._filesystem.invalidate_cache(posixpath.dirname(path))
                time.sleep(self._retry_backoff_s * 2**attempt)
        raise RuntimeError("unreachable")

    @staticmethod
    def _candidate_row_groups(parquet: pq.ParquetFile, episode_index: int) -> list[int]:
        """Keep groups whose episode bounds overlap or whose statistics are missing."""
        episode_column = next(
            index
            for index in range(parquet.metadata.num_columns)
            if parquet.metadata.schema.column(index).path == "episode_index"
        )
        matches = []
        for row_group in range(parquet.metadata.num_row_groups):
            statistics = parquet.metadata.row_group(row_group).column(episode_column).statistics
            if (
                statistics is None
                or not statistics.has_min_max
                or int(statistics.min) <= episode_index <= int(statistics.max)
            ):
                matches.append(row_group)
        return matches

    @staticmethod
    def _validate_complete_episode(
        table: pa.Table,
        episode_index: int,
        expected_rows: int,
        relative_path: str | Path,
    ) -> None:
        """Reject missing or out-of-order episode rows; foreign rows were already filtered out."""
        actual_rows = len(table)
        if actual_rows != expected_rows:
            raise ValueError(
                f"Parquet episode {episode_index} in {relative_path}: "
                f"expected {expected_rows} rows, found {actual_rows}"
            )
        if "frame_index" in table.column_names:
            frame_indices = [int(value) for value in table.column("frame_index").to_pylist()]
            if frame_indices != list(range(expected_rows)):
                raise ValueError(
                    f"Parquet episode {episode_index} in {relative_path} has non-contiguous frame indices"
                )
