From 83d0022890942f636652618919bd8626b98d1a20 Mon Sep 17 00:00:00 2001 From: "xiaohongbo.xhb" Date: Wed, 16 Sep 2026 02:52:24 -0700 Subject: [PATCH] [python] Reduce memory usage for Parquet row-id updates --- paimon-python/pypaimon/read/table_read.py | 6 + .../tests/map_shared_shredding_write_test.py | 44 +++ .../table_update_by_row_id_chunked_test.py | 112 +++++++- .../tests/table_upsert_by_key_test.py | 250 ++++++++++++++++++ paimon-python/pypaimon/write/table_update.py | 1 + .../pypaimon/write/table_update_by_row_id.py | 227 +++++++++++++++- .../pypaimon/write/writer/data_writer.py | 43 ++- .../write/writer/single_file_writer.py | 106 ++++++++ 8 files changed, 767 insertions(+), 22 deletions(-) create mode 100644 paimon-python/pypaimon/write/writer/single_file_writer.py diff --git a/paimon-python/pypaimon/read/table_read.py b/paimon-python/pypaimon/read/table_read.py index 2b0d5a4ea3c3..5b8407874ea1 100644 --- a/paimon-python/pypaimon/read/table_read.py +++ b/paimon-python/pypaimon/read/table_read.py @@ -196,6 +196,12 @@ def _to_managed_arrow_batch_reader( self, splits: List[Split], blob_parallelism: Optional[int] = None): + """Return a closeable batch reader supporting context management. + + Newer PyArrow versions use ``RecordBatchReader.from_stream``. Older + versions fall back to ``_ClosableArrowBatchReader``, which closes both + the batch iterator and its underlying reader. + """ reader, batch_iterator = self._new_arrow_batch_reader( splits, blob_parallelism) if (_RECORD_BATCH_READER_FROM_STREAM is not None diff --git a/paimon-python/pypaimon/tests/map_shared_shredding_write_test.py b/paimon-python/pypaimon/tests/map_shared_shredding_write_test.py index eba20b427037..fe82a6d2ed3e 100644 --- a/paimon-python/pypaimon/tests/map_shared_shredding_write_test.py +++ b/paimon-python/pypaimon/tests/map_shared_shredding_write_test.py @@ -92,6 +92,50 @@ def test_write_and_read_parquet(self): result.column("metrics_overflow").to_pylist(), ) + def test_row_id_update_preserves_shared_shredding(self): + table = self._create_table('parquet', 2, { + 'data-evolution.enabled': 'true', + 'row-tracking.enabled': 'true', + }) + self._write(table, pa.Table.from_pydict({ + 'id': [1, 2], + 'metrics': [[('hot', 1)], [('warm', 2)]], + }, schema=self.arrow_schema)) + + read_builder = table.new_read_builder().with_projection(['id', '_ROW_ID']) + rows = read_builder.new_read().to_arrow( + read_builder.new_scan().plan().splits()) + row_ids = dict(zip(rows.column('id').to_pylist(), + rows.column('_ROW_ID').to_pylist())) + row_id = row_ids[2] + update = pa.Table.from_pydict({ + '_ROW_ID': [row_id], + 'metrics': [[('hot', 99), ('new', 7)]], + }, schema=pa.schema([ + ('_ROW_ID', pa.int64()), + self.arrow_schema.field('metrics'), + ])) + + builder = table.new_batch_write_builder() + messages = builder.new_update().with_update_type( + ['metrics']).update_by_arrow_with_row_id(update) + self.assertEqual(1, len(messages[0].new_files)) + overlay = messages[0].new_files[0] + self.assertTrue(is_shared_shredding( + pq.read_schema(overlay.file_path).field('metrics'))) + commit = builder.new_commit() + commit.commit(messages) + commit.close() + + read_builder = table.new_read_builder().with_projection(['id', 'metrics']) + result = read_builder.new_read().to_arrow( + read_builder.new_scan().plan().splits()) + self.assertEqual( + {1: [('hot', 1)], 2: [('hot', 99), ('new', 7)]}, + dict(zip(result.column('id').to_pylist(), + result.column('metrics').to_pylist())), + ) + def test_reject_orc(self): with self.assertRaisesRegex( ValueError, diff --git a/paimon-python/pypaimon/tests/table_update_by_row_id_chunked_test.py b/paimon-python/pypaimon/tests/table_update_by_row_id_chunked_test.py index c0902c5421ba..3ce93701ee4a 100644 --- a/paimon-python/pypaimon/tests/table_update_by_row_id_chunked_test.py +++ b/paimon-python/pypaimon/tests/table_update_by_row_id_chunked_test.py @@ -21,13 +21,76 @@ import pyarrow as pa import pyarrow.compute as pc +import pyarrow.parquet as pq +from pypaimon.common.options.core_options import ChangelogProducer from pypaimon.table.special_fields import SpecialFields -from pypaimon.write.table_update_by_row_id import TableUpdateByRowId +from pypaimon.write.table_update_by_row_id import ( + TableUpdateByRowId, + _RowIdUpdateFileWriter, +) +from pypaimon.write.writer.single_file_writer import SingleFileWriter class TableUpdateByRowIdChunkedTest(unittest.TestCase): + def test_streaming_overlay_rejects_sidecar_before_output(self): + table = mock.Mock() + table.is_primary_key_table = False + table.fields = [] + options = table.options + options.file_format.return_value = 'parquet' + options.variant_shredding_enabled.return_value = False + options.data_evolution_row_sidecar_enabled.return_value = False + options.with_vector_format.return_value = False + options.changelog_producer.return_value = ChangelogProducer.NONE + self.assertTrue(_RowIdUpdateFileWriter.supports_table(table)) + + options.data_evolution_row_sidecar_enabled.return_value = True + with self.assertRaisesRegex(ValueError, 'requires plain Parquet'): + _RowIdUpdateFileWriter(table, (), ['id']) + table.file_io.new_output_stream.assert_not_called() + + def test_streaming_overlay_rejects_shared_shredding_before_output(self): + table = mock.Mock() + table.is_primary_key_table = False + field = mock.Mock() + field.name = 'metrics' + table.fields = [field] + options = table.options + options.file_format.return_value = 'parquet' + options.variant_shredding_enabled.return_value = False + options.data_evolution_row_sidecar_enabled.return_value = False + options.with_vector_format.return_value = False + options.changelog_producer.return_value = ChangelogProducer.NONE + options.map_storage_layout.return_value = 'shared-shredding' + + with self.assertRaisesRegex(ValueError, 'requires plain Parquet'): + _RowIdUpdateFileWriter(table, (), ['metrics']) + table.file_io.new_output_stream.assert_not_called() + + def test_single_file_writer_rejects_other_formats_before_open(self): + file_io = mock.Mock() + with self.assertRaisesRegex(NotImplementedError, 'only supports Parquet'): + SingleFileWriter(file_io, 'unused.orc', pa.schema([]), 'orc', + 'zstd', 1, [], mock.Mock()) + file_io.new_output_stream.assert_not_called() + + def test_single_file_writer_failed_open_only_deletes_owned_file(self): + file_io = mock.Mock() + file_io.new_output_stream.side_effect = FileExistsError('already exists') + with self.assertRaises(FileExistsError): + SingleFileWriter(file_io, 'existing.parquet', pa.schema([]), 'parquet', + 'zstd', 1, [], mock.Mock()) + file_io.delete_quietly.assert_not_called() + + file_io.new_output_stream.side_effect = None + with mock.patch.object(pq, 'ParquetWriter', side_effect=OSError('failed to open')): + with self.assertRaisesRegex(OSError, 'failed to open'): + SingleFileWriter(file_io, 'new.parquet', pa.schema([]), 'parquet', + 'zstd', 1, [], mock.Mock()) + file_io.delete_quietly.assert_called_once_with('new.parquet') + @staticmethod def _updater(): updater = TableUpdateByRowId.__new__(TableUpdateByRowId) @@ -105,6 +168,53 @@ def test_update_position_outside_column_range_raises(self): self._updater()._merge_update_with_original( original, updates, ["payload"], first_row_id=0) + def test_streaming_update_rejects_non_contiguous_original_row_ids(self): + updates = pa.table({ + SpecialFields.ROW_ID.name: pa.array([0], type=pa.int64()), + "payload": pa.array([3]), + }) + + for row_ids in ([0, 2], [0, 2, 1]): + with self.subTest(row_ids=row_ids): + batch = pa.RecordBatch.from_pydict({ + SpecialFields.ROW_ID.name: + pa.array(row_ids, type=pa.int64()), + "payload": pa.array(row_ids), + }) + managed_reader = mock.MagicMock() + managed_reader.__enter__.return_value = iter([batch]) + table_read = mock.Mock() + table_read._to_managed_arrow_batch_reader.return_value = ( + managed_reader) + updater = self._updater() + updater._original_file_read = mock.Mock( + return_value=(table_read, mock.sentinel.split)) + + with self.assertRaisesRegex( + ValueError, "not contiguous at row ID 0"): + list(updater._merged_batches(0, updates, ["payload"])) + + def test_streaming_update_rejects_row_ids_before_original_group(self): + updates = pa.table({ + SpecialFields.ROW_ID.name: pa.array([9], type=pa.int64()), + "payload": pa.array([3]), + }) + batch = pa.RecordBatch.from_pydict({ + SpecialFields.ROW_ID.name: pa.array([10, 11], type=pa.int64()), + "payload": pa.array([1, 2]), + }) + managed_reader = mock.MagicMock() + managed_reader.__enter__.return_value = iter([batch]) + table_read = mock.Mock() + table_read._to_managed_arrow_batch_reader.return_value = managed_reader + updater = self._updater() + updater._original_file_read = mock.Mock( + return_value=(table_read, mock.sentinel.split)) + + with self.assertRaisesRegex( + ValueError, "precede the original file group"): + list(updater._merged_batches(10, updates, ["payload"])) + def test_total_offsets_over_int32_remain_in_separate_chunks(self): child_length = 1_100_000_000 large_chunk = self._list_chunk([0, child_length]) diff --git a/paimon-python/pypaimon/tests/table_upsert_by_key_test.py b/paimon-python/pypaimon/tests/table_upsert_by_key_test.py index 7f06fcab4299..564d61f3ed51 100644 --- a/paimon-python/pypaimon/tests/table_upsert_by_key_test.py +++ b/paimon-python/pypaimon/tests/table_upsert_by_key_test.py @@ -16,12 +16,16 @@ # under the License. import os +from contextlib import contextmanager import unittest from unittest import mock import pyarrow as pa +import pyarrow.parquet as pq +from pypaimon.read.table_read import TableRead from pypaimon.table.special_fields import SpecialFields +from pypaimon.write.table_update_by_row_id import _RowIdUpdateFileWriter from pypaimon.tests.data_evolution_test_helpers import ( BatchModeMixin, DataEvolutionTestBase, @@ -61,6 +65,197 @@ def _apply_upsert(self, table_update, data, upsert_keys, cid): def _apply_upsert_rows(self, table_update, rows, upsert_keys, cid): raise NotImplementedError + def test_upsert_row_groups_do_not_follow_read_batches(self): + schema = pa.schema([('id', pa.int32()), ('score', pa.int32())]) + for read_size in (73, 1024): + with self.subTest(read_size=read_size): + table = self._create_table(pa_schema=schema, options={ + **self.table_options, 'read.batch-size': str(read_size)}) + original = pa.Table.from_pydict( + {'id': list(range(5000)), 'score': list(range(5000))}, schema=schema) + self._write_arrow(table, original) + updates = pa.Table.from_pydict({'id': [13], 'score': [-1]}, schema=schema) + messages = self._upsert(table, updates, ['id'], ['score']) + files = [f for message in messages for f in message.new_files] + self.assertEqual(len(files), 1) + metadata = pq.read_metadata(files[0].file_path) + self.assertEqual(metadata.num_row_groups, 1) + self.assertEqual(metadata.row_group(0).num_rows, 5000) + scores = self._read_all(table).to_pydict()['score'] + expected = list(range(5000)) + expected[13] = -1 + self.assertEqual(scores, expected) + + def test_row_group_byte_budget_and_oversized_row(self): + schema = pa.schema([('id', pa.int32())]) + table = self._create_table(pa_schema=schema, options={ + **self.table_options, 'file.block-size': '400 b'}) + writer = _RowIdUpdateFileWriter(table, (), ['id']) + data = pa.Table.from_pydict({'id': list(range(250))}, schema=schema) + try: + for batch_size in (1, 73, 250): + groups = list(writer._row_groups( + data.to_batches(max_chunksize=batch_size))) + self.assertEqual([group.num_rows for group in groups], [100, 100, 50]) + self.assertTrue(pa.concat_tables(groups).equals(data)) + self.assertTrue(all(group.nbytes <= 400 for group in groups)) + large = pa.Table.from_pydict({'text': ['a', 'x' * 500, 'b']}) + groups = list(writer._row_groups(large.to_batches())) + self.assertEqual([group.num_rows for group in groups], [1, 1, 1]) + self.assertTrue(pa.concat_tables(groups).equals(large)) + # Empty batches, nulls and sliced variable-width columns must not + # lose rows or make the output depend on input batch boundaries. + mixed = pa.Table.from_pydict({ + 'text': [ + 'skip', None, '', 'a' * 300, + 'b' * 300, None, 'last', 'skip', + ], + }).slice(1, 6) + layouts = [] + for batch_size in (1, 2, 6): + batches = mixed.to_batches(max_chunksize=batch_size) + batches.insert(0, batches[0].slice(0, 0)) + groups = list(writer._row_groups(batches)) + self.assertTrue(pa.concat_tables(groups).equals(mixed)) + self.assertTrue(all(group.nbytes <= 400 for group in groups)) + layouts.append([group.num_rows for group in groups]) + self.assertTrue(all(layout == layouts[0] for layout in layouts)) + self.assertEqual(list(writer._row_groups([])), []) + with mock.patch.object(_RowIdUpdateFileWriter, '_ROW_GROUP_MAX_ROWS', 17): + groups = list(writer._row_groups( + data.to_batches(max_chunksize=73))) + self.assertEqual([group.num_rows for group in groups], [17] * 14 + [12]) + self.assertTrue(pa.concat_tables(groups).equals(data)) + finally: + writer.close() + + def test_invalid_row_group_size_fails_before_opening_output(self): + schema = pa.schema([('id', pa.int32())]) + table = self._create_table(pa_schema=schema, options={ + **self.table_options, 'file.block-size': '0 b'}) + with mock.patch.object( + table.file_io, 'new_output_stream') as output_stream: + with self.assertRaisesRegex( + ValueError, 'file.block-size must be positive'): + _RowIdUpdateFileWriter(table, (), ['id']) + output_stream.assert_not_called() + + @mock.patch.object(_RowIdUpdateFileWriter, '_ROW_GROUP_MAX_ROWS', 2) + def test_partial_upsert_streams_original_file_group(self): + schema = pa.schema([ + ('id', pa.int32()), + ('payload', pa.list_(pa.struct([('text', pa.string())]))), + ('score', pa.int32()), + ]) + table = self._create_table(pa_schema=schema, options={ + **self.table_options, 'metadata.stats-mode': 'full'}) + expected = [{'id': i, 'payload': [{'text': str(i)}], 'score': i} + for i in range(12)] + + def as_table(rows): + return pa.Table.from_pydict( + {name: [row[name] for row in rows] for name in schema.names}, + schema=schema) + + self._write_arrow(table, as_table(expected)) + read_batches = TableRead._to_managed_arrow_batch_reader + write = pq.ParquetWriter.write_table + progress = {'read': 0, 'written': 0} + closed = [] + + @contextmanager + def bounded_reader(reader, splits, **kwargs): + source = read_batches(reader, splits, **kwargs) + + def batches(): + try: + for batch in source: + for start in range(0, batch.num_rows, 2): + self.assertEqual(progress['read'], progress['written']) + piece = batch.slice(start, 2) + progress['read'] += piece.num_rows + yield piece + finally: + source.close() + closed.append(True) + iterator = batches() + try: + yield iterator + finally: + iterator.close() + + def record_write(writer, batch, **kwargs): + result = write(writer, batch, **kwargs) + progress['written'] += batch.num_rows + return result + + for replacements in ([9, 1, 5], [0, 11]): + updates = [{'id': i, 'payload': None if i == 5 else + [{'text': 'updated-' + str(i)}], + 'score': None if i == 5 else -i} for i in replacements] + progress.update(read=0, written=0) + with mock.patch.object(TableRead, 'to_arrow', side_effect=AssertionError( + 'upsert must not materialize the original file group')): + with mock.patch.object(TableRead, '_to_managed_arrow_batch_reader', bounded_reader): + with mock.patch.object(pq.ParquetWriter, 'write_table', record_write): + messages = self._upsert( + table, as_table(updates), + ['id'], ['payload', 'score']) + self.assertEqual(progress, {'read': 12, 'written': 12}) + for row in updates: + expected[row['id']] = row + self.assertEqual(self._read_all(table).to_pydict(), as_table(expected).to_pydict()) + files = [f for msg in messages for f in msg.new_files] + self.assertEqual(len(files), 1) + self.assertEqual((files[0].first_row_id, files[0].row_count), (0, 12)) + scores = [r['score'] for r in expected if r['score'] is not None] + self.assertEqual(files[0].value_stats.min_values.values[1], min(scores)) + self.assertEqual(files[0].value_stats.max_values.values[1], max(scores)) + self.assertEqual(files[0].value_stats.null_counts, [1, 1]) + + # A failed streamed write must leave the committed table intact and + # remove the partially written overlay. + self.assertEqual(len(closed), 2) + progress.update(read=0, written=0) + + def fail_second_write(writer, batch, **kwargs): + if progress['written']: + raise OSError('injected write failure') + return record_write(writer, batch, **kwargs) + + with mock.patch.object(table.file_io, 'delete_quietly', + wraps=table.file_io.delete_quietly) as delete: + with mock.patch.object(pq.ParquetWriter, 'write_table', + fail_second_write): + with mock.patch.object(TableRead, '_to_managed_arrow_batch_reader', bounded_reader): + with self.assertRaisesRegex(OSError, 'injected write failure'): + self._upsert(table, as_table(updates), + ['id'], ['payload', 'score']) + self.assertEqual(len(closed), 3) + self.assertTrue(delete.called) + for call in delete.call_args_list: + self.assertFalse(table.file_io.exists(call[0][0])) + self.assertEqual(self._read_all(table).to_pydict(), as_table(expected).to_pydict()) + + # Closing the format writer is part of the file transaction too. + close = pq.ParquetWriter.close + + def fail_close(writer): + was_open = writer.is_open + close(writer) + if was_open: + raise OSError('injected close failure') + + with mock.patch.object(table.file_io, 'delete_quietly', + wraps=table.file_io.delete_quietly) as delete: + with mock.patch.object(pq.ParquetWriter, 'close', fail_close): + with self.assertRaisesRegex(OSError, 'injected close failure'): + self._upsert(table, as_table(updates), ['id'], ['payload', 'score']) + self.assertTrue(delete.called) + for call in delete.call_args_list: + self.assertFalse(table.file_io.exists(call[0][0])) + self.assertEqual(self._read_all(table).to_pydict(), as_table(expected).to_pydict()) + # ------------------------------------------------------------------ # Helpers built on the primitives # ------------------------------------------------------------------ @@ -953,6 +1148,61 @@ def test_update_cols_partial_update(self): self.assertEqual((1, 'Alice', 99, 'NYC'), rows[0]) self.assertEqual((2, 'Bob', 88, 'LA'), rows[1]) + def test_duplicate_update_cols_are_deduplicated(self): + table = self._create_table() + self._write_arrow(table, pa.Table.from_pydict({ + 'id': [1], + 'name': ['Alice'], + 'age': [25], + 'city': ['NYC'], + }, schema=self.pa_schema)) + + messages = self._upsert( + table, + pa.Table.from_pydict({ + 'id': [1], + 'name': ['ignored'], + 'age': [99], + 'city': ['ignored'], + }, schema=self.pa_schema), + upsert_keys=['id'], + # Matching the schema width must not mean "update all columns". + update_cols=['age'] * len(table.field_names), + ) + + self.assertEqual( + self._read_all(table).to_pydict(), + {'id': [1], 'name': ['Alice'], 'age': [99], 'city': ['NYC']}, + ) + files = [file for message in messages for file in message.new_files] + self.assertEqual([file.write_cols for file in files], [['age']]) + + def test_not_null_update_across_read_batches(self): + schema = pa.schema([ + pa.field('id', pa.int32(), nullable=False), + pa.field('score', pa.int32(), nullable=False), + ]) + table = self._create_table(pa_schema=schema, options={ + **self.table_options, 'read.batch-size': '2'}) + original = pa.Table.from_pydict({ + 'id': list(range(4)), + 'score': list(range(4)), + }, schema=schema) + self._write_arrow(table, original) + + updates = pa.Table.from_pydict({'id': [2], 'score': [99]}, schema=schema) + messages = self._upsert(table, updates, ['id'], ['score']) + + expected = original.set_column( + 1, + schema.field('score'), + pa.array([0, 1, 99, 3], type=pa.int32()), + ) + self.assertTrue(self._read_all(table).equals(expected)) + files = [file for message in messages for file in message.new_files] + self.assertEqual(len(files), 1) + self.assertFalse(pq.read_schema(files[0].file_path).field('score').nullable) + # ================================================================== # Duplicate-key dedup tests — parametrised # ================================================================== diff --git a/paimon-python/pypaimon/write/table_update.py b/paimon-python/pypaimon/write/table_update.py index c358644460ac..922b1bc8c33b 100644 --- a/paimon-python/pypaimon/write/table_update.py +++ b/paimon-python/pypaimon/write/table_update.py @@ -126,6 +126,7 @@ def __init__(self, table, commit_user): self.projection = None def with_update_type(self, update_cols: List[str]): + update_cols = list(dict.fromkeys(update_cols)) for col in update_cols: if col not in self.table.field_names: raise ValueError(f"Column {col} is not in table schema.") diff --git a/paimon-python/pypaimon/write/table_update_by_row_id.py b/paimon-python/pypaimon/write/table_update_by_row_id.py index 90c08d20d2aa..9226f8d8231f 100644 --- a/paimon-python/pypaimon/write/table_update_by_row_id.py +++ b/paimon-python/pypaimon/write/table_update_by_row_id.py @@ -16,6 +16,7 @@ # under the License. import bisect +import uuid from dataclasses import dataclass, field from typing import Any, Dict, List, Optional, Set, Tuple @@ -23,8 +24,10 @@ import pyarrow as pa import pyarrow.compute as pc +from pypaimon.common.options.core_options import ChangelogProducer, CoreOptions from pypaimon.manifest.schema.data_file_meta import DataFileMeta from pypaimon.manifest.schema.manifest_entry import ManifestEntry +from pypaimon.manifest.schema.simple_stats import SimpleStats from pypaimon.read.scanner.data_evolution_split_generator import ( DataEvolutionSplitGenerator, ) @@ -49,6 +52,13 @@ value_for_arrow, ) from pypaimon.write.writer.blob_writer import BlobWriter +from pypaimon.write.writer.append_only_data_writer import AppendOnlyDataWriter +from pypaimon.write.writer.single_file_writer import SingleFileWriter +from pypaimon.write.writer.write_buffer import WriteBuffer + +_ARROW_MAJOR = int(pa.__version__.split('.')[0]) +# Keep aligned with org.apache.parquet.hadoop.ParquetWriter.DEFAULT_BLOCK_SIZE. +_DEFAULT_PARQUET_BLOCK_SIZE = 128 * 1024 * 1024 @dataclass(frozen=True) @@ -66,6 +76,134 @@ class _FilesInfo: valid_row_id_ranges: List[Range] = field(default_factory=list) +class _RowIdUpdateFileWriter: + """Write one plain-Parquet update file for a row-id file group.""" + + _ROW_GROUP_MAX_ROWS = 1024 * 1024 + + @staticmethod + def supports_table(table): + options = table.options + return (not table.is_primary_key_table + and options.file_format(CoreOptions.FILE_FORMAT_PARQUET) + == CoreOptions.FILE_FORMAT_PARQUET + and not (options.variant_shredding_enabled() + and options.variant_shredding_schema()) + and not options.data_evolution_row_sidecar_enabled(False) + and not options.with_vector_format() + and options.changelog_producer() == ChangelogProducer.NONE + and not any(options.map_storage_layout(f.name) == 'shared-shredding' + for f in table.fields) + and not any(is_blob_file_field(f) for f in table.fields)) + + def __init__(self, table, partition, column_names): + if not self.supports_table(table): + raise ValueError('Row-id update file writer requires plain Parquet without sidecars') + configured = table.options.file_block_size() + self._target_bytes = (configured.get_bytes() if configured is not None + else _DEFAULT_PARQUET_BLOCK_SIZE) + if self._target_bytes <= 0: + raise ValueError('file.block-size must be positive') + self._data_writer = AppendOnlyDataWriter( + table, partition, 0, 0, table.options, write_cols=column_names) + + @staticmethod + def _row_group_slice(batch, offset, count): + piece = batch.slice(offset, count) + # Arrow 6 nbytes counts full backing buffers even for slices. Compact + # the slice there so both accounting and retained buffers stay bounded. + if _ARROW_MAJOR < 7: + piece = pa.RecordBatch.from_arrays( + [pa.concat_arrays([column]) for column in piece.columns], schema=piece.schema) + return piece + + def _row_groups(self, batches): + """Bound row-group buffering using Arrow-side bytes, not encoded size.""" + buffer = WriteBuffer(self._data_writer._merge_data) + try: + for batch in batches: + offset = 0 + while offset < batch.num_rows: + count = min(batch.num_rows - offset, + self._ROW_GROUP_MAX_ROWS - buffer.num_rows) + piece = self._row_group_slice(batch, offset, count) + available = self._target_bytes - buffer.nbytes + if piece.nbytes > available: + low, high = 0, count + while low < high: + middle = (low + high + 1) // 2 + if self._row_group_slice(piece, 0, middle).nbytes <= available: + low = middle + else: + high = middle - 1 + count = low + if count == 0 and buffer.num_rows: + yield buffer.take() + continue + count = max(1, count) + piece = self._row_group_slice(piece, 0, count) + buffer.append(pa.Table.from_batches([piece])) + offset += count + del piece + if buffer.nbytes >= self._target_bytes or buffer.num_rows >= self._ROW_GROUP_MAX_ROWS: + yield buffer.take() + del batch + if buffer.num_rows: + yield buffer.take() + finally: + buffer.reset() + + def write_batches(self, batches): + """Write one file while keeping input and output row groups bounded.""" + writer = self._data_writer + file_name = '{}{}-0.parquet'.format( + writer.options.data_file_prefix(), uuid.uuid4()) + file_path = writer._generate_file_path(file_name) + fields = [] + groups = self._row_groups(batches) + file_writer = None + try: + for batch in groups: + if not batch.num_rows: + continue + if file_writer is None: + if writer.options.metadata_stats_enabled(): + fields = PyarrowFieldParser.to_paimon_schema(batch.schema) + file_writer = SingleFileWriter( + writer.file_io, file_path, batch.schema, writer.file_format, + writer.compression, writer.zstd_level, + fields, writer._get_column_stats) + file_writer.write(batch, row_group_size=batch.num_rows) + del batch + if file_writer is None: + return [] + file_writer.close() + meta = writer._create_data_file_meta( + file_name=file_name, + file_path=file_path, + row_count=file_writer.row_count, + min_key=GenericRow([], []), max_key=GenericRow([], []), + key_stats=SimpleStats.empty_stats(), + value_stats=writer._collect_value_stats( + None, fields, file_writer.column_stats), + min_sequence_number=0, max_sequence_number=0, + ) + writer._finish_data_file(meta) + return [meta] + except Exception: + if file_writer is not None: + file_writer.abort() + raise + finally: + groups.close() + + def abort(self): + self._data_writer.abort() + + def close(self): + self._data_writer.close() + + class TableUpdateByRowId: """ Table update for partial column updates (data evolution). @@ -211,6 +349,7 @@ def update_columns(self, data: pa.Table, column_names: List[str]) -> List[Commit if not column_names: raise ValueError("column_names cannot be empty") + column_names = list(dict.fromkeys(column_names)) if SpecialFields.ROW_ID.name not in data.column_names: raise ValueError(f"Input data must contain {SpecialFields.ROW_ID.name} column") @@ -257,6 +396,7 @@ def update_rows_columns( ) -> List[CommitMessage]: if not column_names: raise ValueError("column_names cannot be empty") + column_names = list(dict.fromkeys(column_names)) if len(rows) != len(row_ids_by_row): raise ValueError( "rows and row_ids_by_row must have the same length: " @@ -390,7 +530,7 @@ def _write_by_first_row_id( group_blob_object_columns, ) - def _read_original_file_data(self, first_row_id: int, column_names: List[str]) -> Optional[pa.Table]: + def _read_original_file_data(self, first_row_id: int, column_names: List[str]) -> pa.Table: """Read original file data for the given first_row_id. Only reads columns that exist in the original file and need to be updated. @@ -402,16 +542,17 @@ def _read_original_file_data(self, first_row_id: int, column_names: List[str]) - column_names: The column names to update Returns: - PyArrow Table containing the original data for columns that exist in the file, - or None if no columns need to be read from the original file. + PyArrow Table containing the original values for the requested columns. """ + table_read, origin_split = self._original_file_read(first_row_id, column_names) + original = table_read.to_arrow([origin_split]) + return original.select(column_names) + + def _original_file_read(self, first_row_id, column_names): wanted = set(column_names) read_fields: List[DataField] = [ table_field for table_field in self.table.fields if table_field.name in wanted ] - if not read_fields: - return None - entry = self._first_row_id_index.get(first_row_id) if entry is None: raise ValueError(f"No file found for first_row_id {first_row_id}") @@ -432,8 +573,66 @@ def _read_original_file_data(self, first_row_id: int, column_names: List[str]) - predicate=None, read_type=read_fields + [SpecialFields.ROW_ID], ) - original = table_read.to_arrow([origin_split]) - return original.select([field.name for field in read_fields]) + return table_read, origin_split + + def _merged_batches(self, first_row_id, data, column_names): + """Merge ordinary columns a batch at a time in physical row order.""" + table_read, split = self._original_file_read(first_row_id, column_names) + updates = sorted(enumerate(data[SpecialFields.ROW_ID.name].to_pylist()), + key=lambda item: item[1]) + update_index = 0 + offset = first_row_id + with table_read._to_managed_arrow_batch_reader([split]) as reader: + for batch in reader: + if not batch.num_rows: + continue + row_ids = batch[SpecialFields.ROW_ID.name] + if row_ids.null_count: + raise ValueError( + 'Original file group contains null _ROW_ID values') + row_id_values = row_ids.to_numpy(zero_copy_only=True) + if (row_id_values[0] != offset + or (len(row_id_values) > 1 + and not np.all(np.diff(row_id_values) == 1))): + raise ValueError( + f'Original file group is not contiguous at row ID {offset}') + end = offset + batch.num_rows + selected = [] + while update_index < len(updates) and updates[update_index][1] < end: + if updates[update_index][1] < offset: + raise ValueError( + 'Update row IDs precede the original file group') + selected.append(updates[update_index][0]) + update_index += 1 + original = pa.Table.from_batches([batch]).select(column_names) + if selected: + merged, _ = self._merge_update_with_original( + original, data.take(selected), column_names, offset) + else: + merged = original + yield from merged.to_batches() + offset = end + del batch, original, merged + if update_index != len(updates): + raise ValueError('Update row IDs extend past the original file group') + + def _write_group_streaming(self, partition, first_row_id, data, column_names): + writer = _RowIdUpdateFileWriter( + self.table, tuple(partition.values), column_names) + batches = self._merged_batches(first_row_id, data, column_names) + try: + files = writer.write_batches(batches) + self._assign_update_file_metadata(files, first_row_id, column_names, {}) + if files: + self.commit_messages.append(CommitMessage( + partition=tuple(partition.values), bucket=0, new_files=files, + check_from_snapshot=self.snapshot_id)) + except Exception: + writer.abort() + raise + finally: + batches.close() + writer.close() def _merge_update_with_original( self, @@ -533,7 +732,13 @@ def _merge_update_with_original( merged_columns[col_name] = self._merge_chunked_column( original_col, update_col, sorted_updates) - merged_table = pa.table(merged_columns) if merged_columns else None + merged_table = None + if merged_columns: + merged_schema = pa.schema([ + original_data.schema.field(name) + for name in merged_columns + ]) + merged_table = pa.table(merged_columns, schema=merged_schema) return merged_table, blob_columns @@ -770,6 +975,10 @@ def _write_group( Reads the original file data, merges in the update values, and writes a single output file (rolling disabled) for the group. """ + # Specialized writers still own their sidecars and physical encoding. + if _RowIdUpdateFileWriter.supports_table(self.table): + self._write_group_streaming(partition, first_row_id, data, column_names) + return original_data = self._read_original_file_data(first_row_id, column_names) _, target_files = self._first_row_id_index[first_row_id] blob_columns_with_baseline = { diff --git a/paimon-python/pypaimon/write/writer/data_writer.py b/paimon-python/pypaimon/write/writer/data_writer.py index 46db0578ceb8..71f5004f7eaa 100644 --- a/paimon-python/pypaimon/write/writer/data_writer.py +++ b/paimon-python/pypaimon/write/writer/data_writer.py @@ -343,9 +343,9 @@ def _write_data_to_file(self, data: pa.Table): min_seq = self.sequence_generator.start max_seq = self.sequence_generator.current creation_time = Timestamp.now() - data_meta = DataFileMeta.create( + data_meta = self._create_data_file_meta( file_name=file_name, - file_size=self.file_io.get_file_size(file_path), + file_path=file_path, row_count=data.num_rows, min_key=GenericRow(min_key, self.trimmed_primary_keys_fields), max_key=GenericRow(max_key, self.trimmed_primary_keys_fields), @@ -353,17 +353,8 @@ def _write_data_to_file(self, data: pa.Table): value_stats=value_stats, min_sequence_number=min_seq, max_sequence_number=max_seq, - schema_id=self.table.table_schema.id, - level=0, extra_files=extra_files, creation_time=creation_time, - delete_row_count=0, - file_source=0, - value_stats_cols=None if value_stats_enabled else [], - external_path=external_path_str, - first_row_id=None, - write_cols=self.write_cols, - file_path=file_path, ) if self.changelog_producer == ChangelogProducer.INPUT: @@ -378,8 +369,14 @@ def _write_data_to_file(self, data: pa.Table): self.file_io.delete_quietly(row_sidecar_path) raise + self._finish_data_file(data_meta, changelog_meta, shared_shredding_stats) + + def _finish_data_file(self, data_meta, changelog_meta=None, + shared_shredding_stats=None): + """Record a fully written data file and its optional changelog.""" self.sequence_generator.start = self.sequence_generator.current - self._map_shared_shredding.file_completed(shared_shredding_stats) + if shared_shredding_stats is not None: + self._map_shared_shredding.file_completed(shared_shredding_stats) self.committed_files.append(data_meta) if changelog_meta is not None: self.committed_changelog_files.append(changelog_meta) @@ -391,6 +388,28 @@ def _write_parquet_data(self, path, data): self.file_io.write_parquet(path, data, compression=self.compression, zstd_level=self.zstd_level) return {} + def _create_data_file_meta(self, file_name, file_path, row_count, + min_key, max_key, key_stats, value_stats, + min_sequence_number, max_sequence_number, + extra_files=None, creation_time=None): + """Common metadata finalization for buffered and incremental files.""" + return DataFileMeta.create( + file_name=file_name, + file_size=self.file_io.get_file_size(file_path), + row_count=row_count, + min_key=min_key, max_key=max_key, + key_stats=key_stats, value_stats=value_stats, + min_sequence_number=min_sequence_number, + max_sequence_number=max_sequence_number, + schema_id=self.table.table_schema.id, level=0, + extra_files=extra_files if extra_files is not None else [], + creation_time=creation_time if creation_time is not None else Timestamp.now(), + delete_row_count=0, file_source=0, + value_stats_cols=None if self.options.metadata_stats_enabled() else [], + external_path=file_path if self.external_path_provider is not None else None, + first_row_id=None, write_cols=self.write_cols, file_path=file_path, + ) + def _apply_variant_shredding(self, data: pa.Table) -> pa.Table: """Transform VARIANT columns into shredded Parquet format. diff --git a/paimon-python/pypaimon/write/writer/single_file_writer.py b/paimon-python/pypaimon/write/writer/single_file_writer.py new file mode 100644 index 000000000000..184d013822ee --- /dev/null +++ b/paimon-python/pypaimon/write/writer/single_file_writer.py @@ -0,0 +1,106 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you 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 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +from contextlib import suppress + +import pyarrow.parquet as pq + +from pypaimon.common.options.core_options import CoreOptions + + +class SingleFileWriter: + """Write one file incrementally; currently only Parquet is supported.""" + + def __init__(self, file_io, path, schema, file_format, compression, zstd_level, + stats_fields, stats_collector): + if file_format != CoreOptions.FILE_FORMAT_PARQUET: + raise NotImplementedError( + 'SingleFileWriter only supports Parquet, got {}'.format(file_format)) + self._file_io = file_io + self._path = path + self._stats_fields = stats_fields + self._stats_collector = stats_collector + self._stream = None + self._writer = None + self._owns_file = False + self._closed = False + self.row_count = 0 + self.column_stats = {} + + kwargs = {'compression': compression} + if compression.lower() == 'zstd': + kwargs['compression_level'] = zstd_level + try: + self._stream = file_io.new_output_stream(path) + self._owns_file = True + self._writer = pq.ParquetWriter(self._stream, schema, **kwargs) + except Exception: + self.abort() + raise + + def write(self, data, row_group_size=None): + if self._closed: + raise RuntimeError('Writer is already closed') + try: + kwargs = {} + if row_group_size is not None: + kwargs['row_group_size'] = row_group_size + self._writer.write_table(data, **kwargs) + self.row_count += data.num_rows + for field in self._stats_fields: + current = self._stats_collector(data, field.name) + previous = self.column_stats.get(field.name) + if previous is not None: + current['null_counts'] += previous['null_counts'] + for key, choose in ( + ('min_values', min), ('max_values', max)): + values = [ + value for value in (previous[key], current[key]) + if value is not None + ] + current[key] = choose(values) if values else None + self.column_stats[field.name] = current + except Exception: + self.abort() + raise + + def close(self): + if self._closed: + return + try: + self._writer.close() + self._writer = None + self._stream.close() + self._stream = None + self._closed = True + except Exception: + self.abort() + raise + + def abort(self): + if self._writer is not None: + with suppress(Exception): + self._writer.close() + self._writer = None + if self._stream is not None: + with suppress(Exception): + self._stream.close() + self._stream = None + self._closed = True + if self._owns_file: + self._file_io.delete_quietly(self._path) + self._owns_file = False