Skip to content

Commit de1bdc1

Browse files
committed
fix: stream record batches lazily in to_arrow_batch_reader()
1 parent 068aae5 commit de1bdc1

3 files changed

Lines changed: 198 additions & 2 deletions

File tree

‎pyiceberg/io/pyarrow.py‎

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1760,6 +1760,24 @@ def _read_all_delete_files(io: FileIO, tasks: Iterable[FileScanTask]) -> dict[st
17601760
return deletes_per_file
17611761

17621762

1763+
def _read_deletes_for_task(io: FileIO, task: FileScanTask) -> dict[str, list[ChunkedArray]]:
1764+
"""Read the delete files belonging to a single file scan task.
1765+
1766+
Unlike `_read_all_delete_files`, this reads nothing up front for tasks the
1767+
consumer may never reach, so streaming readers never pay for delete files
1768+
ahead of the batch being yielded.
1769+
"""
1770+
deletes_per_file: dict[str, list[ChunkedArray]] = {}
1771+
for delete_file in task.delete_files:
1772+
for file, arr in _read_deletes(io, delete_file).items():
1773+
if file in deletes_per_file:
1774+
deletes_per_file[file].append(arr)
1775+
else:
1776+
deletes_per_file[file] = [arr]
1777+
1778+
return deletes_per_file
1779+
1780+
17631781
class ArrowScan:
17641782
_table_metadata: TableMetadata
17651783
_io: FileIO
@@ -1891,6 +1909,42 @@ def batches_for_task(task: FileScanTask) -> list[pa.RecordBatch]:
18911909
# This break will also cancel all running tasks in the executor
18921910
break
18931911

1912+
def to_record_batches_lazy(self, tasks: Iterable[FileScanTask]) -> Iterator[pa.RecordBatch]:
1913+
"""Stream record batches one file scan task at a time, in the calling thread.
1914+
1915+
Unlike `to_record_batches`, this never fans work out to the executor and
1916+
never reads ahead of the consumer: each task's delete files are read only
1917+
when the consumer reaches that task, and each task's batches are yielded
1918+
directly instead of being collected into a per-task list first. Peak
1919+
memory therefore stays flat no matter how many files the scan covers.
1920+
1921+
This backs `to_arrow_batch_reader()`, which documents low-memory
1922+
streaming. Callers that want maximum throughput (`to_table()`,
1923+
`to_pandas()`) should keep using `to_record_batches`.
1924+
1925+
Args:
1926+
tasks: FileScanTasks representing the data files and delete files to read from.
1927+
1928+
Returns:
1929+
An Iterator of PyArrow RecordBatches, in task order.
1930+
Total number of rows will be capped if specified.
1931+
1932+
Raises:
1933+
ResolveError: When a required field cannot be found in the file
1934+
ValueError: When a field type in the file cannot be projected to the schema type
1935+
"""
1936+
total_row_count = 0
1937+
for task in tasks:
1938+
deletes_per_file = _read_deletes_for_task(self._io, task)
1939+
for batch in self._record_batches_from_scan_tasks_and_deletes([task], deletes_per_file):
1940+
current_batch_size = len(batch)
1941+
if self._limit is not None and total_row_count + current_batch_size >= self._limit:
1942+
yield batch.slice(0, self._limit - total_row_count)
1943+
return
1944+
else:
1945+
yield batch
1946+
total_row_count += current_batch_size
1947+
18941948
def _record_batches_from_scan_tasks_and_deletes(
18951949
self, tasks: Iterable[FileScanTask], deletes_per_file: dict[str, list[ChunkedArray]]
18961950
) -> Iterator[pa.RecordBatch]:

‎pyiceberg/table/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2381,7 +2381,7 @@ def _to_arrow_batch_reader_via_file_scan_tasks(
23812381
scan.case_sensitive,
23822382
scan.limit,
23832383
dictionary_columns=dictionary_columns,
2384-
).to_record_batches(tasks)
2384+
).to_record_batches_lazy(tasks)
23852385

23862386
if dictionary_columns:
23872387
# schema_to_pyarrow returns plain types, but ArrowScan yields dictionary-encoded

‎tests/io/test_pyarrow.py‎

Lines changed: 143 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,7 @@
6363
Or,
6464
)
6565
from pyiceberg.expressions.literals import literal
66-
from pyiceberg.io import S3_RETRY_STRATEGY_IMPL, S3_SSL_CA_CERT, InputStream, OutputStream, load_file_io
66+
from pyiceberg.io import S3_RETRY_STRATEGY_IMPL, S3_SSL_CA_CERT, FileIO, InputFile, InputStream, OutputFile, OutputStream, load_file_io
6767
from pyiceberg.io.pyarrow import (
6868
ICEBERG_SCHEMA,
6969
PYARROW_PARQUET_FIELD_ID_KEY,
@@ -5524,3 +5524,145 @@ def test_dictionary_columns_produces_dict_encoded_output(tmpdir: str) -> None:
55245524

55255525
# Values must be identical
55265526
assert result_plain.column("label").to_pylist() == result_dict.column("label").to_pylist()
5527+
5528+
5529+
class _CountingFileIO(FileIO):
5530+
"""FileIO decorator recording how many files are opened for reading."""
5531+
5532+
def __init__(self, inner: FileIO) -> None:
5533+
super().__init__()
5534+
self._inner = inner
5535+
self.files_opened_for_read = 0
5536+
5537+
def new_input(self, location: str) -> InputFile:
5538+
self.files_opened_for_read += 1
5539+
return self._inner.new_input(location)
5540+
5541+
def new_output(self, location: str) -> OutputFile:
5542+
return self._inner.new_output(location)
5543+
5544+
def delete(self, location: str | InputFile | OutputFile) -> None:
5545+
self._inner.delete(location)
5546+
5547+
5548+
class _MockScan:
5549+
"""Minimal scan stub for `_to_arrow_batch_reader_via_file_scan_tasks`."""
5550+
5551+
def __init__(self, table_metadata: TableMetadataV2, io: FileIO, limit: int | None = None) -> None:
5552+
self.table_metadata = table_metadata
5553+
self.io = io
5554+
self.row_filter = AlwaysTrue()
5555+
self.case_sensitive = True
5556+
self.limit = limit
5557+
5558+
5559+
def _write_lazy_reader_test_files(
5560+
tmpdir: str, num_files: int, rows_per_file: int
5561+
) -> tuple[Schema, TableMetadataV2, list[FileScanTask]]:
5562+
"""Write `num_files` parquet files with consecutive, non-overlapping ids."""
5563+
iceberg_schema = Schema(
5564+
NestedField(1, "id", IntegerType(), required=False),
5565+
)
5566+
arrow_schema = pa.schema([pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"})])
5567+
table_metadata = TableMetadataV2(
5568+
location=f"file://{tmpdir}",
5569+
last_column_id=1,
5570+
format_version=2,
5571+
schemas=[iceberg_schema],
5572+
partition_specs=[PartitionSpec()],
5573+
)
5574+
tasks = []
5575+
for file_idx in range(num_files):
5576+
start_id = file_idx * rows_per_file
5577+
arrow_table = pa.table(
5578+
[pa.array(range(start_id, start_id + rows_per_file), type=pa.int32())],
5579+
schema=arrow_schema,
5580+
)
5581+
data_file = _write_table_to_data_file(f"{tmpdir}/lazy_reader_{file_idx}.parquet", arrow_schema, arrow_table)
5582+
data_file.spec_id = 0
5583+
tasks.append(FileScanTask(data_file))
5584+
return iceberg_schema, table_metadata, tasks
5585+
5586+
5587+
def test_to_arrow_batch_reader_does_not_read_ahead(tmpdir: str) -> None:
5588+
"""Regression test for https://github.com/apache/iceberg-python/issues/2407.
5589+
5590+
`to_arrow_batch_reader()` documents low-memory streaming ("a RecordBatch is
5591+
read one at a time"), but the reader used to fan every file scan task out to
5592+
the executor and materialize each file's batches into a list, so taking a
5593+
single batch from the reader read every file in the scan.
5594+
"""
5595+
from pyiceberg.table import _to_arrow_batch_reader_via_file_scan_tasks
5596+
5597+
num_files = 4
5598+
iceberg_schema, table_metadata, tasks = _write_lazy_reader_test_files(tmpdir, num_files, rows_per_file=5000)
5599+
io = _CountingFileIO(PyArrowFileIO())
5600+
5601+
reader = _to_arrow_batch_reader_via_file_scan_tasks(_MockScan(table_metadata, io), iceberg_schema, tasks)
5602+
5603+
first_batch = next(reader)
5604+
assert first_batch.num_rows > 0
5605+
assert io.files_opened_for_read == 1, (
5606+
f"expected exactly 1 file opened after consuming the first batch, got {io.files_opened_for_read}"
5607+
)
5608+
5609+
5610+
def test_to_record_batches_lazy_matches_eager(tmpdir: str) -> None:
5611+
"""The lazy streaming path must return the same rows, in the same order, as the threaded path."""
5612+
num_files = 3
5613+
rows_per_file = 2500
5614+
total_rows = num_files * rows_per_file
5615+
iceberg_schema, table_metadata, tasks = _write_lazy_reader_test_files(tmpdir, num_files, rows_per_file)
5616+
expected_ids = list(range(total_rows))
5617+
5618+
for limit in (None, 0, 1, 100, total_rows, total_rows + 10):
5619+
eager_scan = ArrowScan(table_metadata, PyArrowFileIO(), iceberg_schema, AlwaysTrue(), True, limit)
5620+
lazy_scan = ArrowScan(table_metadata, PyArrowFileIO(), iceberg_schema, AlwaysTrue(), True, limit)
5621+
5622+
eager_ids = [row for batch in eager_scan.to_record_batches(tasks) for row in batch.column("id").to_pylist()]
5623+
lazy_ids = [row for batch in lazy_scan.to_record_batches_lazy(tasks) for row in batch.column("id").to_pylist()]
5624+
5625+
assert lazy_ids == eager_ids, f"limit={limit}: lazy path diverged from the threaded path"
5626+
assert lazy_ids == expected_ids[: len(lazy_ids)], f"limit={limit}: unexpected row contents or ordering"
5627+
if limit is None or limit >= total_rows:
5628+
assert len(lazy_ids) == total_rows
5629+
else:
5630+
assert len(lazy_ids) == limit
5631+
5632+
5633+
def test_to_record_batches_lazy_applies_positional_deletes(tmpdir: str) -> None:
5634+
"""The lazy path must apply per-task positional deletes exactly like the eager path."""
5635+
from pyiceberg.table import _to_arrow_batch_reader_via_file_scan_tasks
5636+
5637+
iceberg_schema = Schema(
5638+
NestedField(1, "id", IntegerType(), required=False),
5639+
)
5640+
arrow_schema = pa.schema([pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"})])
5641+
table_metadata = TableMetadataV2(
5642+
location=f"file://{tmpdir}",
5643+
last_column_id=1,
5644+
format_version=2,
5645+
schemas=[iceberg_schema],
5646+
partition_specs=[PartitionSpec()],
5647+
)
5648+
5649+
data_file = _write_table_to_data_file(
5650+
f"{tmpdir}/lazy_reader_deletes.parquet",
5651+
arrow_schema,
5652+
pa.table([pa.array([1, 2, 3, 4], type=pa.int32())], schema=arrow_schema),
5653+
)
5654+
data_file.spec_id = 0
5655+
5656+
# Positional delete of row position 2 (value 3)
5657+
deletes_path = f"{tmpdir}/lazy_reader_deletes_pos.parquet"
5658+
pq.write_table(pa.table({"file_path": [data_file.file_path], "pos": [2]}), deletes_path)
5659+
delete_file = DataFile.from_args(
5660+
content=DataFileContent.POSITION_DELETES, file_path=deletes_path, file_format=FileFormat.PARQUET
5661+
)
5662+
tasks = [FileScanTask(data_file, delete_files={delete_file})]
5663+
5664+
result = _to_arrow_batch_reader_via_file_scan_tasks(
5665+
_MockScan(table_metadata, PyArrowFileIO()), iceberg_schema, tasks
5666+
).read_all()
5667+
5668+
assert result.column("id").to_pylist() == [1, 2, 4]

0 commit comments

Comments
 (0)