Skip to content
153 changes: 142 additions & 11 deletions python/ray/data/_internal/datasource/iceberg_datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,15 +11,17 @@
import pyarrow as pa
from packaging import version

from ray.data._internal.arrow_block import _BATCH_SIZE_PRESERVING_STUB_COL_NAME
from ray.data._internal.planner.plan_expression.expression_visitors import _ExprVisitor
from ray.data._internal.util import _check_import
from ray.data.block import Block, BlockMetadata
from ray.data.block import Block, BlockAccessor, BlockMetadata
from ray.data.datasource.datasource import Datasource, ReadTask
from ray.data.expressions import (
AliasExpr,
BinaryExpr,
ColumnExpr,
DownloadExpr,
Expr,
LiteralExpr,
MonotonicallyIncreasingIdExpr,
Operation,
Expand Down Expand Up @@ -84,6 +86,44 @@

logger = logging.getLogger(__name__)

# Field ID for the fabricated stub field, see ``_get_empty_projection_schema``.
# Iceberg matches columns by ID, so this must not be a real column's: colliding
# with a non-boolean column raises ``ResolveError``, and with a boolean column
# reads that column's data instead (the row count is still right, but the read
# is no longer free). Catalogs assign IDs sequentially from 1, so a value this
# large will not collide; Iceberg reserves 2147483546 and above for metadata
# columns, so stay below that.
_STUB_FIELD_ID = 2147483000


def _get_empty_projection_schema() -> "Schema":
"""Return a projected schema that reads no data but keeps the row count.

PyIceberg rebuilds its output to match the projected schema, and a table
rebuilt from zero fields reports zero rows however many the scan matched --
so an empty projection, which ``Dataset.count()`` legitimately requests,
would silently count zero. Project a single optional field that no data file
has instead: Iceberg fills a missing optional field with nulls, the same way
it reads a file written before a column was added, so the row count arrives
in a column that costs nothing to read.

The field is named after Ray's stub column so the resulting block matches
what every other reader produces for an empty projection. ``boolean`` is
used because Iceberg gained a null-valued type only after 0.9.0;
``_get_read_task`` retypes the column to ``null`` on the way out.
"""
from pyiceberg.schema import Schema
from pyiceberg.types import BooleanType, NestedField

return Schema(
NestedField(
field_id=_STUB_FIELD_ID,
name=_BATCH_SIZE_PRESERVING_STUB_COL_NAME,
field_type=BooleanType(),
required=False,
)
)


class _IcebergExpressionVisitor(
_ExprVisitor["BooleanExpression | UnboundTerm | Literal"]
Expand Down Expand Up @@ -210,6 +250,7 @@ def _get_read_task(
case_sensitive: bool,
limit: Optional[int],
schema: "Schema",
empty_projection: bool = False,
) -> Iterable[Block]:
# Determine the PyIceberg version to handle backward compatibility
import pyiceberg
Expand Down Expand Up @@ -255,7 +296,14 @@ def _generate_tables() -> Iterable[pa.Table]:
)
yield table

yield from _generate_tables()
for table in _generate_tables():
if empty_projection:
# The scan read a boolean-typed stub, see
# ``_get_empty_projection_schema``. Re-derive it through Ray's own
# empty projection so the column is ``null``-typed, matching the
# schema we report and every other reader's stub.
table = BlockAccessor.for_block(table).select([])
yield table
Comment on lines +299 to +306

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why are we reading all tables for the row count? Why not read from the manifest http://iceberg.apache.org/spec/#point-in-time-reads-time-travel



@DeveloperAPI
Expand Down Expand Up @@ -297,8 +345,13 @@ def __init__(
_check_import(self, module="pyiceberg", package="pyiceberg")
from pyiceberg.expressions import AlwaysTrue

self._scan_kwargs = scan_kwargs if scan_kwargs is not None else {}
self._catalog_kwargs = catalog_kwargs if catalog_kwargs is not None else {}
# Copy both dicts: we pop ``name`` from the catalog kwargs and write
# ``snapshot_id`` into the scan kwargs, and doing that in place would
# break a second read that reuses the caller's dict.
self._scan_kwargs = dict(scan_kwargs) if scan_kwargs is not None else {}
self._catalog_kwargs = (
dict(catalog_kwargs) if catalog_kwargs is not None else {}
)

if "name" in self._catalog_kwargs:
self._catalog_name = self._catalog_kwargs.pop("name")
Expand Down Expand Up @@ -344,10 +397,39 @@ def plan_files(self) -> List["FileScanTask"]:
# Calculate and cache the plan_files if they don't already exist
if self._plan_files is None:
data_scan = self._get_data_scan()
self._plan_files = data_scan.plan_files()
# Annotated ``Iterable[FileScanTask]``; a list today, but this cache
# is iterated more than once, so don't depend on that.
self._plan_files = list(data_scan.plan_files())

return self._plan_files

def apply_predicate(self, predicate_expr: Expr) -> "IcebergDatasource":
"""Push a predicate down, discarding plan files planned without it.

The base implementation shallow-copies, so the clone would inherit a
cache planned before this predicate existed: all-``AlwaysTrue``
residuals, which ``get_read_tasks`` turns into an unfiltered row count
for a filtered read.
"""
return self._invalidate_plan_files(super().apply_predicate(predicate_expr))

def apply_projection(
self, projection_map: Optional[Dict[str, str]]
) -> "IcebergDatasource":
"""Push a projection down, discarding plan files planned without it.

The projected columns are part of the scan, so a cache built for a
different projection does not belong to the clone. See ``apply_predicate``.
"""
return self._invalidate_plan_files(super().apply_projection(projection_map))

def _invalidate_plan_files(self, clone: "IcebergDatasource") -> "IcebergDatasource":
"""Drop ``clone``'s inherited plan-file cache, so it re-plans its scan."""
# ``self`` signals "no pushdown applied": no clone, cache still valid.
if clone is not self:
clone._plan_files = None
return clone
Comment on lines +426 to +431

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not sure I understand why we're invalidating the property here.


def _get_combined_filter(self) -> "BooleanExpression":
"""Get the combined filter including both row_filter and pushed-down predicates."""
combined_filter = self._row_filter
Expand All @@ -367,13 +449,22 @@ def _get_combined_filter(self) -> "BooleanExpression":

return combined_filter

def _is_empty_projection(self) -> bool:
"""Whether the pushed-down projection selects no columns at all."""
return self._get_data_columns() == []

def _get_data_scan(self) -> "DataScan":
# Get the combined filter
combined_filter = self._get_combined_filter()

# Convert back to tuple for PyIceberg API (None -> ("*",))
# Convert back to tuple for PyIceberg API (None -> ("*",)). An empty
# projection also scans everything: it selects its own columns through
# the projected schema instead, see ``_get_empty_projection_schema``.
data_columns = self._get_data_columns()
selected_fields = ("*",) if data_columns is None else tuple(data_columns)
if not data_columns:
selected_fields = ("*",)
else:
selected_fields = tuple(data_columns)

data_scan = self.table.scan(
row_filter=combined_filter,
Expand Down Expand Up @@ -433,6 +524,7 @@ def get_read_tasks(
per_task_row_limit: Optional[int] = None,
data_context: Optional["DataContext"] = None,
) -> List[ReadTask]:
from pyiceberg.expressions import AlwaysTrue
from pyiceberg.io import pyarrow as pyi_pa_io
from pyiceberg.manifest import DataFileContent

Expand All @@ -447,10 +539,21 @@ def get_read_tasks(
# Get the arrow schema, to set in the metadata
pya_schema = pyi_pa_io.schema_to_pyarrow(projected_schema)

# An empty projection reads a fabricated stub column instead of none at
# all, see ``_get_empty_projection_schema``. Declare the placeholder so
# the reported schema matches the blocks; it is hidden from the
# user-visible schema.
empty_projection = self._is_empty_projection()
if empty_projection:
projected_schema = _get_empty_projection_schema()
pya_schema = pa.schema(
[pa.field(_BATCH_SIZE_PRESERVING_STUB_COL_NAME, pa.null())]
)

# Set the n_chunks to the min of the number of plan files and the actual
# requested n_chunks, so that there are no empty tasks
if parallelism > len(list(plan_files)):
parallelism = len(list(plan_files))
if parallelism > len(plan_files):
parallelism = len(plan_files)
logger.warning(
f"Reducing the parallelism to {parallelism}, as that is the number of files"
)
Expand All @@ -467,6 +570,26 @@ def get_read_tasks(
case_sensitive = self._scan_kwargs.get("case_sensitive", True)
limit = self._scan_kwargs.get("limit")

# Manifests record how many rows each file holds, which is only still an
# exact count if every surviving row matches the filter. PyIceberg puts
# whatever the filter leaves after file pruning in
# ``FileScanTask.residual``, so ``AlwaysTrue`` means nothing is left to
# check and the counts hold. Same test PyIceberg's own
# ``DataScan.count()`` uses. A ``limit`` stops the read early, so no
# count taken from the manifests describes what comes back.
#
# The subtraction below only accounts for position deletes, but an
# equality delete cannot reach here: ``plan_files`` rejects them in both
# of its paths, locally in ``_plan_files_local`` and over REST in
# ``FileScanTask.from_rest_response``.
counts_are_exact = limit is None and (
isinstance(row_filter, AlwaysTrue)
or all(
isinstance(getattr(task, "residual", None), AlwaysTrue)
for task in plan_files
)
)
Comment thread
cursor[bot] marked this conversation as resolved.
Comment thread
cursor[bot] marked this conversation as resolved.

get_read_task = partial(
_get_read_task,
table_io=table_io,
Expand All @@ -475,6 +598,7 @@ def get_read_tasks(
case_sensitive=case_sensitive,
limit=limit,
schema=projected_schema,
empty_projection=empty_projection,
)

read_tasks = []
Expand All @@ -495,9 +619,16 @@ def get_read_tasks(
for delete in unique_deletes
if delete.content == DataFileContent.POSITION_DELETES
)
# ``None`` means "unknown", which makes ``Dataset.count()`` do the
# read instead of trusting these manifest counts.
num_rows = None
if counts_are_exact:
num_rows = (
sum(task.file.record_count for task in chunk_tasks)
- position_delete_count
)
metadata = BlockMetadata(
num_rows=sum(task.file.record_count for task in chunk_tasks)
- position_delete_count,
num_rows=num_rows,
size_bytes=sum(task.file.file_size_in_bytes for task in chunk_tasks),
input_files=[task.file.file_path for task in chunk_tasks],
exec_stats=None,
Expand Down
Loading
Loading