diff --git a/mkdocs/docs/api.md b/mkdocs/docs/api.md index 786223d47e..8a72957880 100644 --- a/mkdocs/docs/api.md +++ b/mkdocs/docs/api.md @@ -367,7 +367,7 @@ for buf in tbl.scan().to_arrow_batch_reader(): ### Streaming writes from a `RecordBatchReader` -`tbl.append()` and `tbl.overwrite()` also accept a `pyarrow.RecordBatchReader` directly, which lets you write datasets that don't fit in memory without materialising them as a `pa.Table` first. PyIceberg consumes the reader once and microbatches it into Parquet files of approximately `write.target-file-size-bytes` (default 512 MiB), keeping memory usage bounded by the target size. All files are committed in a single snapshot. +`tbl.append()` and `tbl.overwrite()` also accept a `pyarrow.RecordBatchReader` directly, which lets you write datasets that don't fit in memory without materialising them as a `pa.Table` first. PyIceberg consumes the reader once and writes Parquet files one at a time, rolling to a new file when the file on disk reaches `write.target-file-size-bytes` (default 512 MiB). Memory usage is bounded by one row group (the smaller of `write.parquet.row-group-limit` rows and `write.parquet.row-group-size-bytes`, uncompressed) plus one input batch and the writer's buffers. All files are committed in a single snapshot. ```python reader = pa.RecordBatchReader.from_batches(schema, batch_iter) diff --git a/mkdocs/docs/configuration.md b/mkdocs/docs/configuration.md index 18f9db973b..e68146c47a 100644 --- a/mkdocs/docs/configuration.md +++ b/mkdocs/docs/configuration.md @@ -82,6 +82,7 @@ Iceberg tables support table properties to configure table behavior. | `write.parquet.compression-codec` | `{uncompressed,zstd,gzip,snappy}` | zstd | Sets the Parquet compression coddec. | | `write.parquet.compression-level` | Integer | null | Parquet compression level for the codec. If not set, it is up to PyIceberg | | `write.parquet.row-group-limit` | Number of rows | 1048576 | The upper bound of the number of entries within a single row group | +| `write.parquet.row-group-size-bytes` | Size in bytes | 128MB | Streaming `RecordBatchReader` writes only: flush buffered rows into a row group when their uncompressed Arrow size reaches this value. `pa.Table` writes ignore it. Unlike Java, the size is uncompressed Arrow bytes. | | `write.parquet.page-size-bytes` | Size in bytes | 1MB | Set a target threshold for the approximate encoded size of data pages within a column chunk | | `write.parquet.page-row-limit` | Number of rows | 20000 | Set a target threshold for the maximum number of rows within a column chunk | | `write.parquet.dict-size-bytes` | Size in bytes | 2MB | Set the dictionary page size limit per row group | diff --git a/pyiceberg/io/fileformat.py b/pyiceberg/io/fileformat.py index 65c0cbeff9..79c321424a 100644 --- a/pyiceberg/io/fileformat.py +++ b/pyiceberg/io/fileformat.py @@ -122,6 +122,10 @@ def write(self, table: pa.Table) -> None: def close(self) -> DataFileStatistics: """Finalize the file and return statistics.""" + @abstractmethod + def length(self) -> int: + """Return the estimated size of the file in bytes, including rows buffered but not yet written.""" + def result(self) -> DataFileStatistics: """Return statistics from a previous close() call.""" if self._result is None: diff --git a/pyiceberg/io/pyarrow.py b/pyiceberg/io/pyarrow.py index 1e59107da3..55a5a90e75 100644 --- a/pyiceberg/io/pyarrow.py +++ b/pyiceberg/io/pyarrow.py @@ -148,7 +148,7 @@ ) from pyiceberg.table import DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE, TableProperties from pyiceberg.table.deletion_vector import deletion_vectors_from_puffin_file -from pyiceberg.table.locations import load_location_provider +from pyiceberg.table.locations import LocationProvider, load_location_provider from pyiceberg.table.metadata import TableMetadata from pyiceberg.table.name_mapping import NameMapping, apply_name_mapping from pyiceberg.table.puffin import PuffinFile @@ -2657,12 +2657,23 @@ def __init__(self, output_file: OutputFile, file_schema: Schema, properties: Pro self._properties = properties self._writer: pq.ParquetWriter | None = None self._fos: OutputStream | None = None + self._pending: list[pa.Table] = [] + self._pending_rows = 0 + self._pending_bytes = 0 + self._compression_ratio = 1.0 + self._closed = False + self._length = 0 self._parquet_writer_kwargs = _get_parquet_writer_kwargs(properties) - self._row_group_size = property_as_int( + self._row_group_size: int = property_as_int( # type: ignore # The property is set with non-None value. properties=properties, property_name=TableProperties.PARQUET_ROW_GROUP_LIMIT, default=TableProperties.PARQUET_ROW_GROUP_LIMIT_DEFAULT, ) + self._row_group_bytes: int = property_as_int( # type: ignore # The property is set with non-None value. + properties=properties, + property_name=TableProperties.PARQUET_ROW_GROUP_SIZE_BYTES, + default=TableProperties.PARQUET_ROW_GROUP_SIZE_BYTES_DEFAULT, + ) def write(self, table: pa.Table) -> None: if self._writer is None: @@ -2678,15 +2689,46 @@ def write(self, table: pa.Table) -> None: fos.close() raise self._fos = fos - self._writer.write(table, row_group_size=self._row_group_size) + self._pending.append(table) + self._pending_rows += table.num_rows + self._pending_bytes += table.nbytes + if self._pending_bytes >= self._row_group_bytes: + self._flush(pa.concat_tables(self._pending)) + self._pending = [] + self._pending_rows = 0 + self._pending_bytes = 0 + elif self._pending_rows >= self._row_group_size: + pending = pa.concat_tables(self._pending) + full_rows = pending.num_rows - pending.num_rows % self._row_group_size + self._flush(pending.slice(0, full_rows)) + tail = pending.slice(full_rows) + self._pending = [tail] if tail.num_rows else [] + self._pending_rows = tail.num_rows + self._pending_bytes = tail.nbytes + + def _flush(self, table: pa.Table) -> None: + if self._writer is None or self._fos is None: + raise ValueError("Cannot flush a writer that was never written to") + before = self._fos.tell() + self._writer.write_table(table, row_group_size=self._row_group_size) + if table.nbytes > 0: + self._compression_ratio = min(1.0, (self._fos.tell() - before) / table.nbytes) def close(self) -> DataFileStatistics: if self._result is not None: return self._result if self._writer is None or self._fos is None: raise ValueError("Cannot close a writer that was never written to") + self._length = self.length() + self._closed = True with self._fos: - self._writer.close() + try: + if self._pending: + self._writer.write_table(pa.concat_tables(self._pending), row_group_size=self._row_group_size) + self._pending = [] + finally: + self._writer.close() + self._length = self._fos.tell() self._result = data_file_statistics_from_parquet_metadata( parquet_metadata=self._writer.writer.metadata, stats_columns=compute_statistics_plan(self._file_schema, self._properties), @@ -2694,6 +2736,18 @@ def close(self) -> DataFileStatistics: ) return self._result + def length(self) -> int: + """Return the bytes written so far plus the estimated on-disk size of the rows not yet flushed. + + Buffered rows count at their in-memory size scaled by the compression ratio observed on + the last flush. The ratio is 1.0 until the first flush in each file, so a target smaller + than one row group gets no correction. A file that rolls on this value can land above the + target when later rows compress worse than the last flushed row group. + """ + if self._closed or self._fos is None: + return self._length + return self._fos.tell() + int(self._pending_bytes * self._compression_ratio) + class ParquetFormatModel(FileFormatModel): """Format model for Apache Parquet.""" @@ -2721,25 +2775,50 @@ def add_field_metadata(self, field: NestedField, metadata: dict[bytes, bytes], i FileFormatFactory.register(ParquetFormatModel()) -def write_file(io: FileIO, table_metadata: TableMetadata, tasks: Iterator[WriteTask]) -> Iterator[DataFile]: - from pyiceberg.table import DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE, TableProperties +def _build_data_file( + table_metadata: TableMetadata, + file_format: FileFormat, + file_path: str, + output_file: OutputFile, + partition: Record, + statistics: DataFileStatistics, +) -> DataFile: + return DataFile.from_args( + content=DataFileContent.DATA, + file_path=file_path, + file_format=file_format, + partition=partition, + file_size_in_bytes=len(output_file), + # After this has been fixed: + # https://github.com/apache/iceberg-python/issues/271 + # sort_order_id=task.sort_order_id, + sort_order_id=None, + # Just copy these from the table for now + spec_id=table_metadata.default_spec_id, + equality_ids=None, + key_metadata=None, + **statistics.to_serialized_dict(), + ) + +def _resolve_write_setup(table_metadata: TableMetadata) -> tuple[FileFormat, FileFormatModel, LocationProvider, Schema]: file_format = FileFormat( - table_metadata.properties.get( - TableProperties.WRITE_FILE_FORMAT, - TableProperties.WRITE_FILE_FORMAT_DEFAULT, - ) + table_metadata.properties.get(TableProperties.WRITE_FILE_FORMAT, TableProperties.WRITE_FILE_FORMAT_DEFAULT) ) format_model = FileFormatFactory.get(file_format) location_provider = load_location_provider(table_location=table_metadata.location, table_properties=table_metadata.properties) + table_schema = table_metadata.schema() + if (sanitized_schema := sanitize_column_names(table_schema)) != table_schema: + file_schema = sanitized_schema + else: + file_schema = table_schema + return file_format, format_model, location_provider, file_schema - def write_data_file(task: WriteTask) -> DataFile: - table_schema = table_metadata.schema() - if (sanitized_schema := sanitize_column_names(table_schema)) != table_schema: - file_schema = sanitized_schema - else: - file_schema = table_schema +def write_file(io: FileIO, table_metadata: TableMetadata, tasks: Iterator[WriteTask]) -> Iterator[DataFile]: + file_format, format_model, location_provider, file_schema = _resolve_write_setup(table_metadata) + + def write_data_file(task: WriteTask) -> DataFile: downcast_ns_timestamp_to_us = Config().get_bool(DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE) or False batches = [ _to_requested_schema( @@ -2761,23 +2840,14 @@ def write_data_file(task: WriteTask) -> DataFile: writer = format_model.create_writer(fo, file_schema, table_metadata.properties) with writer: writer.write(arrow_table) - statistics = writer.result() - return DataFile.from_args( - content=DataFileContent.DATA, - file_path=file_path, + return _build_data_file( + table_metadata=table_metadata, file_format=file_format, + file_path=file_path, + output_file=fo, partition=task.partition_key.partition if task.partition_key else Record(), - file_size_in_bytes=len(fo), - # After this has been fixed: - # https://github.com/apache/iceberg-python/issues/271 - # sort_order_id=task.sort_order_id, - sort_order_id=None, - # Just copy these from the table for now - spec_id=table_metadata.default_spec_id, - equality_ids=None, - key_metadata=None, - **statistics.to_serialized_dict(), + statistics=writer.result(), ) executor = ExecutorFactory.get_or_create() @@ -2795,9 +2865,8 @@ def bin_pack_arrow_table(tbl: pa.Table, target_file_size: int) -> Iterator[list[ bytes. The resulting Parquet file after compression (zstd by default, plus dictionary/RLE encoding) is typically 3-10× smaller than ``target_file_size``. This is a coarse proxy for the spec-defined - ``write.target-file-size-bytes`` and will be tightened to true on-disk - bytes once the writer is switched to a rolling-``ParquetWriter`` with - ``OutputStream.tell()`` (#2998). + ``write.target-file-size-bytes``. The ``pa.RecordBatchReader`` write + path measures on-disk bytes instead. """ from pyiceberg.utils.bin_packing import PackingIterator @@ -2832,9 +2901,8 @@ def bin_pack_record_batches(batches: Iterable[pa.RecordBatch], target_file_size: bytes (``RecordBatch.nbytes``), not compressed on-disk Parquet bytes. The resulting Parquet file after compression is typically 3-10× smaller than ``target_file_size``. Matches the existing - :func:`bin_pack_arrow_table` semantics; both will be tightened to true - on-disk bytes once the writer is switched to a rolling- - ``ParquetWriter`` with ``OutputStream.tell()`` (#2998). + :func:`bin_pack_arrow_table` semantics. The ``pa.RecordBatchReader`` + write path no longer uses this function; it measures on-disk bytes. """ buffer: list[pa.RecordBatch] = [] buffer_bytes = 0 @@ -2929,7 +2997,6 @@ def _get_parquet_writer_kwargs(table_properties: Properties) -> dict[str, Any]: from pyiceberg.table import TableProperties for key_pattern in [ - TableProperties.PARQUET_ROW_GROUP_SIZE_BYTES, TableProperties.PARQUET_BLOOM_FILTER_MAX_BYTES, f"{TableProperties.PARQUET_BLOOM_FILTER_COLUMN_ENABLED_PREFIX}.*", ]: @@ -2978,8 +3045,8 @@ def _dataframe_to_data_files( For a ``pa.Table`` the data is materialised in memory and bin-packed into target-sized files (with partition splitting if the table is partitioned). - For a ``pa.RecordBatchReader`` batches are streamed and microbatched into - target-sized files using bounded memory (see :func:`bin_pack_record_batches`). + For a ``pa.RecordBatchReader`` batches are streamed into one file at a time, + which rolls when its on-disk size reaches ``write.target-file-size-bytes``. Streaming writes are currently only supported on unpartitioned tables; partitioned support is tracked in https://github.com/apache/iceberg-python/issues/2152. @@ -3012,14 +3079,38 @@ def _dataframe_to_data_files( "Materialise the reader as a pa.Table first, or follow " "https://github.com/apache/iceberg-python/issues/2152 for partitioned streaming support." ) - yield from write_file( - io=io, - table_metadata=table_metadata, - tasks=( - WriteTask(write_uuid=write_uuid, task_id=next(counter), record_batches=batches, schema=task_schema) - for batches in bin_pack_record_batches(df, target_file_size) - ), - ) + file_format, format_model, location_provider, file_schema = _resolve_write_setup(table_metadata) + + batches = iter(df) + for batch in batches: + task = WriteTask(write_uuid=write_uuid, task_id=next(counter), record_batches=[], schema=task_schema) + file_path = location_provider.new_data_location( + data_file_name=task.generate_data_file_filename(format_model.file_extension()), + partition_key=None, + ) + fo = io.new_output(file_path) + writer = format_model.create_writer(fo, file_schema, table_metadata.properties) + with writer: + current: pa.RecordBatch | None = batch + while current is not None: + projected = _to_requested_schema( + requested_schema=file_schema, + file_schema=task_schema, + batch=current, + downcast_ns_timestamp_to_us=downcast_ns_timestamp_to_us, + include_field_ids=True, + format_model=format_model, + ) + writer.write(pa.Table.from_batches([projected])) + current = next(batches, None) if writer.length() < target_file_size else None + yield _build_data_file( + table_metadata=table_metadata, + file_format=file_format, + file_path=file_path, + output_file=fo, + partition=Record(), + statistics=writer.result(), + ) return if table_metadata.spec().is_unpartitioned(): diff --git a/pyiceberg/table/__init__.py b/pyiceberg/table/__init__.py index 2c5c26800c..2a01c33bbb 100644 --- a/pyiceberg/table/__init__.py +++ b/pyiceberg/table/__init__.py @@ -496,8 +496,8 @@ def append( Shorthand API for appending PyArrow data to a table transaction. Accepts either a fully materialised ``pa.Table`` or a streaming - ``pa.RecordBatchReader``. Streaming is microbatched by - ``write.target-file-size-bytes`` so memory stays bounded; the reader is + ``pa.RecordBatchReader``. Streaming writes one file at a time and rolls + to a new file at ``write.target-file-size-bytes``; the reader is consumed once and cannot be reused. Streaming writes are currently only supported on unpartitioned tables; @@ -523,13 +523,12 @@ def append( in storage that are not referenced by any snapshot. Clean these up with expire/orphan-file maintenance jobs. - ``write.target-file-size-bytes`` is currently interpreted as - uncompressed in-memory Arrow bytes (the bin-packing weight) rather - than compressed on-disk Parquet bytes. The resulting files are - typically 3-10× smaller than the property suggests after - compression. This matches the existing ``pa.Table`` write path and - will be tightened once the writer is switched to a - rolling-``ParquetWriter`` with ``OutputStream.tell()`` (#2998). + For a ``pa.RecordBatchReader``, ``write.target-file-size-bytes`` + is measured in on-disk bytes through the format writer. A file + rolls when it reaches the target. Memory is bounded by one row + group (the smaller of ``write.parquet.row-group-limit`` rows and + ``write.parquet.row-group-size-bytes``, uncompressed) plus one + input batch and the writer's buffers. Args: df: An Arrow Table or a RecordBatchReader of records to append. @@ -652,8 +651,8 @@ def overwrite( Shorthand for adding a table overwrite with a PyArrow table or RecordBatchReader to the transaction. Accepts either a fully materialised ``pa.Table`` or a streaming - ``pa.RecordBatchReader``. Streaming is microbatched by - ``write.target-file-size-bytes`` so memory stays bounded; the reader is + ``pa.RecordBatchReader``. Streaming writes one file at a time and rolls + to a new file at ``write.target-file-size-bytes``; the reader is consumed once and cannot be reused. Streaming writes are currently only supported on unpartitioned tables; @@ -679,13 +678,12 @@ def overwrite( in storage that are not referenced by any snapshot. Clean these up with expire/orphan-file maintenance jobs. - ``write.target-file-size-bytes`` is currently interpreted as - uncompressed in-memory Arrow bytes (the bin-packing weight) rather - than compressed on-disk Parquet bytes. The resulting files are - typically 3-10× smaller than the property suggests after - compression. This matches the existing ``pa.Table`` write path and - will be tightened once the writer is switched to a - rolling-``ParquetWriter`` with ``OutputStream.tell()`` (#2998). + For a ``pa.RecordBatchReader``, ``write.target-file-size-bytes`` + is measured in on-disk bytes through the format writer. A file + rolls when it reaches the target. Memory is bounded by one row + group (the smaller of ``write.parquet.row-group-limit`` rows and + ``write.parquet.row-group-size-bytes``, uncompressed) plus one + input batch and the writer's buffers. An overwrite may produce zero or more snapshots based on the operation: diff --git a/tests/catalog/test_catalog_behaviors.py b/tests/catalog/test_catalog_behaviors.py index b859e2d541..cffe4af7e5 100644 --- a/tests/catalog/test_catalog_behaviors.py +++ b/tests/catalog/test_catalog_behaviors.py @@ -1215,8 +1215,8 @@ def test_drop_namespace_raises_error_when_namespace_not_empty( # RecordBatchReader streaming append/overwrite tests # -# Streaming writes accept a pa.RecordBatchReader and microbatch it into target-sized -# Parquet files instead of materialising the full Arrow Table in memory. Tracks +# Streaming writes accept a pa.RecordBatchReader and write it into Parquet files that +# roll at the target size instead of materialising the full Arrow Table in memory. Tracks # https://github.com/apache/iceberg-python/issues/2152. @@ -1248,7 +1248,7 @@ def test_append_record_batch_reader(catalog: Catalog) -> None: def test_append_record_batch_reader_microbatched(catalog: Catalog) -> None: """A reader bigger than the per-file target produces multiple Parquet files - in a single snapshot — verifying the byte-budget microbatching path.""" + in a single snapshot, which verifies that the streaming writer rolls files.""" catalog.create_namespace("default") identifier = f"default.append_record_batch_reader_microbatch_{catalog.name}" reader, total_rows = _simple_record_batch_reader(num_batches=8) diff --git a/tests/integration/test_writes/test_writes.py b/tests/integration/test_writes/test_writes.py index 30fdd76ab7..58b3a05fc5 100644 --- a/tests/integration/test_writes/test_writes.py +++ b/tests/integration/test_writes/test_writes.py @@ -694,7 +694,6 @@ def test_write_parquet_other_properties( @pytest.mark.parametrize( "properties", [ - {"write.parquet.row-group-size-bytes": "42"}, {"write.parquet.bloom-filter-enabled.column.bool": "42"}, {"write.parquet.bloom-filter-max-bytes": "42"}, ], diff --git a/tests/io/test_fileformat.py b/tests/io/test_fileformat.py index d5d487fa6d..cd9cf39a32 100644 --- a/tests/io/test_fileformat.py +++ b/tests/io/test_fileformat.py @@ -74,6 +74,9 @@ def write(self, table: Any) -> None: def close(self) -> DataFileStatistics: raise NotImplementedError + def length(self) -> int: + return 0 + writer = _DummyWriter() with pytest.raises(RuntimeError, match="Writer has not been closed yet"): writer.result() diff --git a/tests/io/test_format_writers.py b/tests/io/test_format_writers.py index 86ae65156b..bd6acede2f 100644 --- a/tests/io/test_format_writers.py +++ b/tests/io/test_format_writers.py @@ -27,7 +27,7 @@ from pyiceberg.io.pyarrow import PyArrowFileIO from pyiceberg.manifest import FileFormat from pyiceberg.schema import Schema -from pyiceberg.types import LongType, NestedField +from pyiceberg.types import LongType, NestedField, StringType @pytest.fixture(params=FileFormatFactory.available_formats(), ids=lambda f: f.name.lower()) @@ -96,6 +96,48 @@ def test_close_is_idempotent( assert stats1 is stats2 +def test_length_tracks_bytes_written( + format_model: FileFormatModel, table_schema_simple: Schema, arrow_table_simple: pa.Table, tmp_path: Path +) -> None: + """length() is 0 before the first write, grows after a write, and equals the file size after close.""" + output_file = PyArrowFileIO().new_output(str(tmp_path / f"test.{format_model.file_extension()}")) + writer = format_model.create_writer(output_file, table_schema_simple, {}) + assert writer.length() == 0 + writer.write(arrow_table_simple) + first = writer.length() + assert first > 0 + writer.write(arrow_table_simple) + assert writer.length() > first + writer.close() + assert writer.length() == len(output_file) + + +def test_parquet_length_counts_pending_bytes_before_first_flush(tmp_path: Path) -> None: + """Before the first flush, length() counts buffered rows at their Arrow size.""" + schema = Schema(NestedField(1, "payload", StringType(), required=False)) + writer = FileFormatFactory.get(FileFormat.PARQUET).create_writer( + PyArrowFileIO().new_output(str(tmp_path / "test.parquet")), schema, {} + ) + table = pa.table({"payload": ["constant-payload"] * 100}) + writer.write(table) + assert writer.length() == writer._fos.tell() + table.nbytes # type: ignore[attr-defined] + writer.close() + + +def test_parquet_length_scales_pending_bytes_by_last_flush_ratio(tmp_path: Path) -> None: + """After a flush, length() scales buffered rows by the compression ratio of that flush.""" + schema = Schema(NestedField(1, "payload", StringType(), required=False)) + writer = FileFormatFactory.get(FileFormat.PARQUET).create_writer( + PyArrowFileIO().new_output(str(tmp_path / "test.parquet")), schema, {"write.parquet.row-group-limit": "1000"} + ) + writer.write(pa.table({"payload": ["constant-payload"] * 1000})) + table = pa.table({"payload": ["constant-payload"] * 500}) + writer.write(table) + tell = writer._fos.tell() # type: ignore[attr-defined] + assert tell <= writer.length() < tell + table.nbytes + writer.close() + + def test_close_without_write_raises(format_model: FileFormatModel, table_schema_simple: Schema, tmp_path: Path) -> None: """Closing a writer that was never written to raises ValueError.""" file_path = str(tmp_path / f"test.{format_model.file_extension()}") @@ -125,6 +167,31 @@ def test_parquet_writer_closes_output_stream_on_construction_failure( assert writer._fos is None +def test_parquet_writer_failed_flush_keeps_original_error( + table_schema_simple: Schema, + arrow_table_simple: pa.Table, + tmp_path: Path, +) -> None: + """On an error exit, a pending row group that fails to flush closes the stream and does not hide the error.""" + from pyiceberg.io.pyarrow import ParquetFormatModel + + def fail() -> None: + raise ValueError("original error") + + writer = ParquetFormatModel().create_writer( + PyArrowFileIO().new_output(str(tmp_path / "test.parquet")), table_schema_simple, {} + ) + with pytest.raises(ValueError, match="original error"): + with writer: + writer.write(arrow_table_simple) + writer.write(pa.table({"other": [1]})) + fail() + + assert writer._fos is not None + assert writer._fos.closed # type: ignore[attr-defined] + assert writer.length() > 0 + + def test_parquet_format_model_adds_field_id_metadata() -> None: """ParquetFormatModel.add_field_metadata writes the Parquet field-id key when requested.""" from pyiceberg.io.pyarrow import PYARROW_PARQUET_FIELD_ID_KEY, ParquetFormatModel diff --git a/tests/io/test_pyarrow.py b/tests/io/test_pyarrow.py index 892d8e54eb..05edd84134 100644 --- a/tests/io/test_pyarrow.py +++ b/tests/io/test_pyarrow.py @@ -16,9 +16,11 @@ # under the License. # pylint: disable=protected-access,unused-argument,redefined-outer-name import logging +import math import os import sys import tempfile +import time import uuid import warnings from collections.abc import Iterator @@ -73,7 +75,9 @@ StatsAggregator, _check_pyarrow_schema_compatible, _ConvertToArrowSchema, + _dataframe_to_data_files, _determine_partitions, + _get_parquet_writer_kwargs, _primitive_to_physical, _read_deletes, _task_to_record_batches, @@ -3141,6 +3145,227 @@ def test_write_file_parquet_round_trip(tmp_path: Path, table_schema_simple: Sche assert result.column_names == ["foo", "bar", "baz"] +_STREAM_SCHEMA = Schema( + NestedField(field_id=1, name="id", field_type=LongType(), required=False), + NestedField(field_id=2, name="payload", field_type=StringType(), required=False), +) +_STREAM_ROWS_PER_BATCH = 1000 + + +def _stream_batches(num_batches: int, consumed: list[int] | None = None) -> Iterator[pa.RecordBatch]: + for i in range(num_batches): + if consumed is not None: + consumed.append(i) + yield pa.RecordBatch.from_pydict( + { + "id": pa.array(range(i * _STREAM_ROWS_PER_BATCH, (i + 1) * _STREAM_ROWS_PER_BATCH), type=pa.int64()), + "payload": [uuid4().hex for _ in range(_STREAM_ROWS_PER_BATCH)], + } + ) + + +def _stream_reader(batches: Iterator[pa.RecordBatch]) -> pa.RecordBatchReader: + return pa.RecordBatchReader.from_batches(pa.schema([("id", pa.int64()), ("payload", pa.string())]), batches) + + +def _stream_table_metadata(tmp_path: Path, target_file_size: int, row_group_limit: int | None = None) -> TableMetadataV2: + properties = {TableProperties.WRITE_TARGET_FILE_SIZE_BYTES: str(target_file_size)} + if row_group_limit is not None: + properties[TableProperties.PARQUET_ROW_GROUP_LIMIT] = str(row_group_limit) + return TableMetadataV2( + location=f"file://{tmp_path}", + last_column_id=2, + format_version=2, + schemas=[_STREAM_SCHEMA], + partition_specs=[PartitionSpec()], + properties=properties, + ) + + +def test_dataframe_to_data_files_reader_rolls_at_on_disk_target(tmp_path: Path) -> None: + target = 64 * 1024 + table_metadata = _stream_table_metadata(tmp_path, target, row_group_limit=_STREAM_ROWS_PER_BATCH) + data_files = list(_dataframe_to_data_files(table_metadata, _stream_reader(_stream_batches(20)), PyArrowFileIO())) + + assert len(data_files) > 1 + for data_file in data_files[:-1]: + assert data_file.file_size_in_bytes >= target + for data_file in data_files: + assert data_file.file_size_in_bytes == os.path.getsize(data_file.file_path.removeprefix("file://")) + assert sum(data_file.record_count for data_file in data_files) == 20 * _STREAM_ROWS_PER_BATCH + assert len({data_file.file_path for data_file in data_files}) == len(data_files) + + +def test_dataframe_to_data_files_reader_rolls_near_target_for_compressible_data(tmp_path: Path) -> None: + rows_per_batch = 1000 + num_batches = 600 + + def compressible_batches() -> Iterator[pa.RecordBatch]: + for i in range(num_batches): + yield pa.RecordBatch.from_pydict( + { + "id": pa.array(range(i * rows_per_batch, (i + 1) * rows_per_batch), type=pa.int64()), + "payload": ["payload-" * 8] * rows_per_batch, + } + ) + + target = 256 * 1024 + table_metadata = _stream_table_metadata(tmp_path, target, row_group_limit=2 * rows_per_batch) + data_files = list(_dataframe_to_data_files(table_metadata, _stream_reader(compressible_batches()), PyArrowFileIO())) + + assert len(data_files) > 3 + for data_file in data_files[:-1]: + assert 0.9 * target <= data_file.file_size_in_bytes <= 1.25 * target + assert sum(data_file.record_count for data_file in data_files) == num_batches * rows_per_batch + + +def test_dataframe_to_data_files_reader_large_target_writes_one_file(tmp_path: Path) -> None: + data_files = list( + _dataframe_to_data_files( + _stream_table_metadata(tmp_path, 512 * 1024 * 1024), _stream_reader(_stream_batches(20)), PyArrowFileIO() + ) + ) + + assert len(data_files) == 1 + assert data_files[0].record_count == 20 * _STREAM_ROWS_PER_BATCH + + +def test_dataframe_to_data_files_reader_packs_small_batches_into_row_groups(tmp_path: Path) -> None: + def small_batches() -> Iterator[pa.RecordBatch]: + for i in range(500): + yield pa.RecordBatch.from_pydict( + {"id": pa.array(range(i * 10, (i + 1) * 10), type=pa.int64()), "payload": ["x"] * 10} + ) + + table_metadata = _stream_table_metadata(tmp_path, 512 * 1024 * 1024, row_group_limit=1200) + data_files = list(_dataframe_to_data_files(table_metadata, _stream_reader(small_batches()), PyArrowFileIO())) + + assert len(data_files) == 1 + metadata = pq.read_metadata(data_files[0].file_path.removeprefix("file://")) + assert metadata.num_rows == 5000 + assert metadata.num_row_groups == math.ceil(5000 / 1200) + assert [metadata.row_group(i).num_rows for i in range(metadata.num_row_groups)] == [1200, 1200, 1200, 1200, 200] + + +def test_dataframe_to_data_files_reader_many_one_row_batches_is_linear(tmp_path: Path) -> None: + table_metadata = _stream_table_metadata(tmp_path, 512 * 1024 * 1024) + reader = _stream_reader(pa.RecordBatch.from_pydict({"id": [i], "payload": ["x"]}) for i in range(5000)) + + start = time.perf_counter() + data_files = list(_dataframe_to_data_files(table_metadata, reader, PyArrowFileIO())) + elapsed = time.perf_counter() - start + + assert len(data_files) == 1 + metadata = pq.read_metadata(data_files[0].file_path.removeprefix("file://")) + assert metadata.num_rows == 5000 + assert metadata.num_row_groups == 1 + assert elapsed < 5 + + +def test_dataframe_to_data_files_reader_flushes_row_groups_on_bytes(tmp_path: Path) -> None: + schema = Schema( + NestedField(field_id=1, name="id", field_type=LongType(), required=False), + NestedField(field_id=2, name="value", field_type=DoubleType(), required=False), + ) + batches = [ + pa.RecordBatch.from_pydict( + {"id": pa.array(range(i * 1000, (i + 1) * 1000), type=pa.int64()), "value": pa.array([0.5] * 1000)} + ) + for i in range(30) + ] + row_group_bytes = 40_000 + batches_per_row_group = math.ceil(row_group_bytes / batches[0].nbytes) + table_metadata = TableMetadataV2( + location=f"file://{tmp_path}", + last_column_id=2, + format_version=2, + schemas=[schema], + partition_specs=[PartitionSpec()], + properties={ + TableProperties.PARQUET_ROW_GROUP_SIZE_BYTES: str(row_group_bytes), + TableProperties.PARQUET_ROW_GROUP_LIMIT: "1000000", + }, + ) + reader = pa.RecordBatchReader.from_batches(batches[0].schema, iter(batches)) + + data_files = list(_dataframe_to_data_files(table_metadata, reader, PyArrowFileIO())) + + assert len(data_files) == 1 + metadata = pq.read_metadata(data_files[0].file_path.removeprefix("file://")) + assert batches[0].nbytes == 16_000 + assert metadata.num_row_groups == math.ceil(30 / batches_per_row_group) == 10 + assert {metadata.row_group(i).num_rows for i in range(metadata.num_row_groups)} == {batches_per_row_group * 1000} + + +@pytest.mark.parametrize("row_group_limit, row_group_bytes", [(4, None), (10, None), (100, None), (4, 64), (10, 64), (100, 64)]) +def test_write_file_row_groups_match_single_parquet_write( + tmp_path: Path, table_schema_simple: Schema, row_group_limit: int, row_group_bytes: int | None +) -> None: + properties = {TableProperties.PARQUET_ROW_GROUP_LIMIT: str(row_group_limit)} + if row_group_bytes is not None: + properties[TableProperties.PARQUET_ROW_GROUP_SIZE_BYTES] = str(row_group_bytes) + table_metadata, task = _simple_write_task_and_metadata(tmp_path, table_schema_simple, properties) + arrow_data = pa.Table.from_batches(task.record_batches * 7) + if row_group_bytes is not None: + assert arrow_data.nbytes >= row_group_bytes + task = WriteTask(write_uuid=task.write_uuid, task_id=0, record_batches=arrow_data.to_batches(), schema=task.schema) + + data_file = next(iter(write_file(io=PyArrowFileIO(), table_metadata=table_metadata, tasks=iter([task])))) + new_path = data_file.file_path.removeprefix("file://") + written = pq.read_table(new_path) + + old_path = str(tmp_path / "single_write.parquet") + with pq.ParquetWriter( + old_path, schema=written.schema, store_decimal_as_integer=True, **_get_parquet_writer_kwargs(table_metadata.properties) + ) as writer: + writer.write(written, row_group_size=row_group_limit) + + new_metadata, old_metadata = pq.read_metadata(new_path), pq.read_metadata(old_path) + assert new_metadata.num_row_groups == old_metadata.num_row_groups == math.ceil(21 / row_group_limit) + assert [new_metadata.row_group(i).num_rows for i in range(new_metadata.num_row_groups)] == [ + old_metadata.row_group(i).num_rows for i in range(old_metadata.num_row_groups) + ] + with open(new_path, "rb") as new_file, open(old_path, "rb") as old_file: + assert new_file.read() == old_file.read() + + +def test_dataframe_to_data_files_reader_yields_file_before_reading_rest(tmp_path: Path) -> None: + consumed: list[int] = [] + data_files = _dataframe_to_data_files( + _stream_table_metadata(tmp_path, 64 * 1024), _stream_reader(_stream_batches(20, consumed)), PyArrowFileIO() + ) + + first = next(iter(data_files)) + + assert len(consumed) * _STREAM_ROWS_PER_BATCH == first.record_count + assert len(consumed) < 20 + + +def test_dataframe_to_data_files_reader_error_closes_stream(tmp_path: Path) -> None: + def failing_batches() -> Iterator[pa.RecordBatch]: + yield from _stream_batches(3) + raise ValueError("reader failed") + + streams: list[Any] = [] + original_create = PyArrowFile.create + + def tracking_create(self: PyArrowFile, overwrite: bool = False) -> OutputStream: + stream = original_create(self, overwrite=overwrite) + streams.append(stream) + return stream + + with patch.object(PyArrowFile, "create", tracking_create): + with pytest.raises(ValueError, match="reader failed"): + list( + _dataframe_to_data_files( + _stream_table_metadata(tmp_path, 512 * 1024 * 1024), _stream_reader(failing_batches()), PyArrowFileIO() + ) + ) + + assert len(streams) == 1 + assert streams[0].closed + + def test__to_requested_schema_timestamps( arrow_table_schema_with_all_timestamp_precisions: pa.Schema, arrow_table_with_all_timestamp_precisions: pa.Table,