This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new 1884aca55e [python] Stream batches when building generic global
indexes (#9751)
1884aca55e is described below
commit 1884aca55e7f4a6397a2b33bca574dd13b88cc67
Author: chaoyang <[email protected]>
AuthorDate: Mon Sep 14 08:42:04 2026 +0800
[python] Stream batches when building generic global indexes (#9751)
---
.../pypaimon/globalindex/create_global_index.py | 57 ++++---
.../pypaimon/tests/global_index_build_test.py | 189 ++++++++++++++++++++-
.../pypaimon/tests/vindex_batch_write_test.py | 7 +-
3 files changed, 227 insertions(+), 26 deletions(-)
diff --git a/paimon-python/pypaimon/globalindex/create_global_index.py
b/paimon-python/pypaimon/globalindex/create_global_index.py
index b2e125e421..d68221a0ae 100644
--- a/paimon-python/pypaimon/globalindex/create_global_index.py
+++ b/paimon-python/pypaimon/globalindex/create_global_index.py
@@ -288,6 +288,8 @@ class GlobalIndexBuilder:
def _build_generic_index(
self, splits, unindexed_ranges, index_field, table_read, index_path:
str
) -> List[CommitMessage]:
+ from pypaimon.read.table_read import _ClosableArrowBatchReader
+
rows_per_shard = self._core_options.global_index_row_count_per_shard()
if rows_per_shard <= 0:
raise ValueError(
@@ -298,26 +300,38 @@ class GlobalIndexBuilder:
for index_split, index_range in _split_by_global_index_shard(
splits, rows_per_shard, unindexed_ranges
):
- table = table_read.to_arrow([index_split])
- if table is None or table.num_rows == 0:
- continue
-
- writer = self._create_generic_index_writer(index_path, index_field)
+ writer = None
try:
- if self._index_type in VINDEX_IDENTIFIERS:
- if table.column(SpecialFields.ROW_ID.name).null_count:
- raise ValueError("Cannot build global index because
_ROW_ID is null.")
- for batch in
table.to_batches(max_chunksize=ADD_BATCH_SIZE):
- _write_vector_batch(
- writer, batch, self._index_columns[0], index_range)
- else:
- for value, row_id in _extract_index_rows(
- table,
- self._index_columns[0],
- SpecialFields.ROW_ID.name,
- index_range,
- ):
- writer.write(value, row_id - index_range.from_)
+ reader, batches =
table_read._new_arrow_batch_reader([index_split])
+ # Close the Python iterator explicitly on failure as well as
+ # the Arrow reader, which may retain a suspended generator.
+ with _ClosableArrowBatchReader(reader, batches) as
batch_reader:
+ for batch in batch_reader:
+ if batch.num_rows == 0:
+ continue
+ if writer is None:
+ writer = self._create_generic_index_writer(
+ index_path, index_field)
+ if self._index_type in VINDEX_IDENTIFIERS:
+ if
batch.column(SpecialFields.ROW_ID.name).null_count:
+ raise ValueError(
+ "Cannot build global index because _ROW_ID
is null.")
+ for offset in range(0, batch.num_rows,
ADD_BATCH_SIZE):
+ _write_vector_batch(
+ writer, batch.slice(offset,
ADD_BATCH_SIZE),
+ self._index_columns[0], index_range)
+ else:
+ for value, row_id in _extract_index_rows(
+ batch,
+ self._index_columns[0],
+ SpecialFields.ROW_ID.name,
+ index_range,
+ ):
+ writer.write(value, row_id - index_range.from_)
+ del batch
+
+ if writer is None:
+ continue
index_adds = _to_index_manifest_entries(
self._table,
@@ -328,7 +342,8 @@ class GlobalIndexBuilder:
writer.finish(),
)
finally:
- writer.close()
+ if writer is not None:
+ writer.close()
if index_adds:
messages.append(
CommitMessage(
@@ -469,7 +484,7 @@ def _write_vector_batch(writer, batch, index_column,
row_range):
def _extract_index_rows(
- table: pa.Table,
+ table: Union[pa.Table, pa.RecordBatch],
index_column: str,
row_id_column: str,
row_range: Optional[Range] = None,
diff --git a/paimon-python/pypaimon/tests/global_index_build_test.py
b/paimon-python/pypaimon/tests/global_index_build_test.py
index 186270fcb7..1633b5270f 100644
--- a/paimon-python/pypaimon/tests/global_index_build_test.py
+++ b/paimon-python/pypaimon/tests/global_index_build_test.py
@@ -22,6 +22,7 @@ import os
import struct
import sys
import types
+from unittest.mock import Mock, patch
import pyarrow as pa
@@ -556,7 +557,8 @@ class GlobalIndexBuildTest(
('id', pa.int32()),
('embedding', pa.list_(pa.float32())),
])
- table = self._create_table(pa_schema=schema,
options=self.table_options)
+ table = self._create_table(pa_schema=schema, options=dict(
+ self.table_options, **{'read.batch-size': '1'}))
vectors = pa.array(
[[1.0, 0.0], [0.0, 1.0], None],
type=pa.list_(pa.float32()),
@@ -610,12 +612,56 @@ class GlobalIndexBuildTest(
table.path_factory().global_index_path_factory().to_path(
entry.index_file.file_name)))
+ def test_create_vindex_streaming_failure_cleans_resources(self):
+ from pypaimon.read.table_read import TableRead
+
+ schema = pa.schema([('embedding', pa.list_(pa.float32()))])
+ table = self._create_table(pa_schema=schema, options=dict(
+ self.table_options, **{'read.batch-size': '1'}))
+ self._write_arrow(table, pa.table(
+ {'embedding': [[1.0, 0.0], [0.0, 1.0], [0.5, 0.5]]},
schema=schema))
+ snapshot_id = table.snapshot_manager().get_latest_snapshot().id
+ original_write = VindexVectorIndexWriter.write_batch
+ original_batches = TableRead._arrow_batch_generator
+ temp_paths = []
+ closed = []
+ generators = []
+
+ def failing_write(writer, vector, row_id):
+ original_write(writer, vector, row_id)
+ if 1 in row_id.to_pylist():
+ temp_paths.extend([writer._row_id_temp_path,
writer._vector_temp_path])
+ raise RuntimeError('injected write failure')
+
+ def tracked_batches(reader, *args):
+ def generate():
+ try:
+ yield from original_batches(reader, *args)
+ finally:
+ closed.append(True)
+ generator = generate()
+ generators.append(generator)
+ return generator
+
+ with patch.object(VindexVectorIndexWriter, 'write_batch',
failing_write), \
+ patch.object(TableRead, '_arrow_batch_generator',
tracked_batches):
+ with self.assertRaisesRegex(RuntimeError, 'injected write
failure'):
+ table.create_global_index('embedding', index_type='ivf-flat',
options={
+ 'ivf-flat.dimension': '2',
+ })
+
+ self.assertEqual([True], closed)
+ self.assertEqual(2, len(temp_paths))
+ self.assertTrue(all(not os.path.exists(path) for path in temp_paths))
+ self.assertEqual(snapshot_id,
table.snapshot_manager().get_latest_snapshot().id)
+
def test_create_vindex_global_index_respects_row_count_per_shard(self):
schema = pa.schema([
('id', pa.int32()),
('embedding', pa.list_(pa.float32())),
])
- table = self._create_table(pa_schema=schema,
options=self.table_options)
+ table = self._create_table(pa_schema=schema, options=dict(
+ self.table_options, **{'read.batch-size': '1'}))
vectors = pa.array(
[[1.0, 0.0], [0.0, 1.0], [0.5, 0.5], [0.2, 0.8], [0.9, 0.1]],
type=pa.list_(pa.float32()),
@@ -774,7 +820,8 @@ class GlobalIndexBuildTest(
('id', pa.int32()),
('content', pa.string()),
])
- table = self._create_table(pa_schema=schema,
options=self.table_options)
+ table = self._create_table(pa_schema=schema, options=dict(
+ self.table_options, **{'read.batch-size': '1'}))
self._write_arrow(table, pa.table(
{
'id': [1, 2, 3],
@@ -1181,5 +1228,141 @@ class GlobalIndexBuildTest(
self.assertEqual(value, actual)
+class GenericIndexStreamingTest(unittest.TestCase):
+
+ schema = pa.schema([
+ ('embedding', pa.list_(pa.float32())),
+ ('_ROW_ID', pa.int64()),
+ ])
+
+ def setUp(self):
+ self.builder = object.__new__(GlobalIndexBuilder)
+ self.builder._table = Mock()
+ self.builder._core_options = Mock()
+
self.builder._core_options.global_index_row_count_per_shard.return_value = 10
+ self.builder._index_columns = ['embedding']
+ self.builder._index_type = 'ivf-flat'
+ self.writer = Mock()
+ self.writer.finish.return_value = []
+ self.builder._create_generic_index_writer =
Mock(return_value=self.writer)
+ self.read = Mock()
+ self.read.to_arrow.side_effect = AssertionError('Must not materialize
a shard')
+ self.events = []
+
+ def _batch(self, values, row_ids):
+ return pa.RecordBatch.from_arrays([
+ pa.array(values, type=self.schema.field(0).type),
+ pa.array(row_ids, type=pa.int64()),
+ ], schema=self.schema)
+
+ def _build(self, batches):
+ def generate():
+ try:
+ for batch in batches:
+ self.events.append('read')
+ yield batch
+ finally:
+ self.events.append('reader closed')
+
+ # Retain the generator: cleanup must be explicit, not depend on GC.
+ self.generator = generate()
+ reader = pa.RecordBatchReader.from_batches(self.schema, self.generator)
+ self.read._new_arrow_batch_reader.return_value = reader, self.generator
+ module = 'pypaimon.globalindex.create_global_index'
+ with patch(module + '._split_by_global_index_shard', return_value=[
+ (_FakeSplit([]), Range(10, 19)),
+ ]), patch(module + '._to_index_manifest_entries', return_value=[]):
+ return self.builder._build_generic_index(
+ [], [], Mock(), self.read, '/unused')
+
+ def test_batches_are_written_before_reading_the_next_batch(self):
+ self.builder._index_type = 'lucene'
+ written = []
+
+ def write(value, row_id):
+ self.events.append('write')
+ written.append((value, row_id))
+
+ def finish():
+ self.assertEqual('reader closed', self.events[-1])
+ return []
+
+ self.writer.write.side_effect = write
+ self.writer.finish.side_effect = finish
+ self._build([
+ self._batch([], []),
+ self._batch([[9.0], [10.0], None], [9, 10, 11]),
+ self._batch([[19.0], [20.0]], [19, 20]),
+ ])
+ self.assertEqual([([10.0], 0), (None, 1), ([19.0], 9)], written)
+ self.assertEqual([
+ 'read', 'read', 'write', 'write', 'read', 'write', 'reader closed',
+ ], self.events)
+ self.writer.finish.assert_called_once()
+ self.writer.close.assert_called_once()
+ self.read.to_arrow.assert_not_called()
+
+ def test_vector_batches_are_bounded_and_written_before_next_read(self):
+ written = []
+
+ def write_batch(vectors, row_ids):
+ self.assertLessEqual(len(row_ids), 2)
+ self.events.append('write batch')
+ written.extend(zip(vectors.to_pylist(), row_ids.to_pylist()))
+
+ self.writer.write_batch.side_effect = write_batch
+ with patch('pypaimon.globalindex.create_global_index.ADD_BATCH_SIZE',
2):
+ self._build([
+ self._batch([], []),
+ self._batch([[9.0], [10.0], None, [12.0]], [9, 10, 11, 12]),
+ self._batch([[19.0], [20.0]], [19, 20]),
+ ])
+ self.assertEqual([([10.0], 0), (None, 1), ([12.0], 2), ([19.0], 9)],
written)
+ self.assertEqual([
+ 'read', 'read', 'write batch', 'write batch', 'read',
+ 'write batch', 'reader closed',
+ ], self.events)
+ self.writer.write.assert_not_called()
+ self.writer.finish.assert_called_once()
+ self.writer.close.assert_called_once()
+ self.read.to_arrow.assert_not_called()
+
+ def test_empty_input_does_not_create_a_writer(self):
+ for batches in ([], [self._batch([], [])]):
+ with self.subTest(batches=len(batches)):
+ self.assertEqual([], self._build(batches))
+ self.builder._create_generic_index_writer.assert_not_called()
+ self.assertEqual('reader closed', self.events[-1])
+
+ def test_failures_close_reader_and_writer(self):
+ for failure in ('create', 'read', 'write', 'finish', 'null_row_id'):
+ with self.subTest(failure=failure):
+ self.setUp()
+ error = RuntimeError('injected failure')
+ if failure == 'create':
+ self.builder._create_generic_index_writer.side_effect =
error
+ elif failure in ('write', 'finish'):
+ getattr(self.writer, 'write_batch' if failure == 'write'
else failure).side_effect = error
+
+ def batches():
+ yield self._batch([[10.0]], [10])
+ if failure == 'read':
+ raise error
+ yield self._batch([[11.0]], [
+ None if failure == 'null_row_id' else 11])
+
+ exception = ValueError if failure == 'null_row_id' else
RuntimeError
+ message = '_ROW_ID is null' if failure == 'null_row_id' else
'injected failure'
+ with self.assertRaisesRegex(exception, message):
+ self._build(batches())
+ self.assertEqual('reader closed', self.events[-1])
+ if failure == 'create':
+ self.writer.close.assert_not_called()
+ else:
+ self.writer.close.assert_called_once()
+ if failure != 'finish':
+ self.writer.finish.assert_not_called()
+
+
if __name__ == "__main__":
unittest.main()
diff --git a/paimon-python/pypaimon/tests/vindex_batch_write_test.py
b/paimon-python/pypaimon/tests/vindex_batch_write_test.py
index 72da3bb745..def8858c04 100644
--- a/paimon-python/pypaimon/tests/vindex_batch_write_test.py
+++ b/paimon-python/pypaimon/tests/vindex_batch_write_test.py
@@ -167,7 +167,7 @@ class VindexBatchWriteTest(unittest.TestCase):
with self.assertRaisesRegex(ValueError, '_ROW_ID is null'):
_write_vector_batch(self._writer(), batch, 'embedding', Range(10,
19))
- def test_null_row_ids_are_rejected_before_writing_any_batch(self):
+ def test_null_row_ids_are_rejected_before_writing_source_batch(self):
builder = object.__new__(GlobalIndexBuilder)
builder._core_options = Mock()
builder._core_options.global_index_row_count_per_shard.return_value =
10
@@ -176,10 +176,13 @@ class VindexBatchWriteTest(unittest.TestCase):
writer = Mock()
builder._create_generic_index_writer = Mock(return_value=writer)
read = Mock()
- read.to_arrow.return_value = pa.table({
+ table = pa.table({
'embedding': pa.array([[1], [3, 4]], type=pa.list_(pa.float32())),
'_ROW_ID': pa.array([0, None], type=pa.int64()),
})
+ batches = iter(table.to_batches())
+ read._new_arrow_batch_reader.return_value = (
+ pa.RecordBatchReader.from_batches(table.schema, batches), batches)
module = 'pypaimon.globalindex.create_global_index'
with patch(module + '._split_by_global_index_shard', return_value=[
(Mock(), Range(0, 9)),