Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 36 additions & 21 deletions paimon-python/pypaimon/globalindex/create_global_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,8 @@ def _create_sorted_index_writer(self, index_path: str, key_serializer):
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(
Expand All @@ -298,26 +300,38 @@ def _build_generic_index(
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,
Expand All @@ -328,7 +342,8 @@ def _build_generic_index(
writer.finish(),
)
finally:
writer.close()
if writer is not None:
writer.close()
if index_adds:
messages.append(
CommitMessage(
Expand Down Expand Up @@ -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,
Expand Down
189 changes: 186 additions & 3 deletions paimon-python/pypaimon/tests/global_index_build_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import struct
import sys
import types
from unittest.mock import Mock, patch

import pyarrow as pa

Expand Down Expand Up @@ -556,7 +557,8 @@ def test_create_vindex_global_index_from_python(self):
('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()),
Expand Down Expand Up @@ -610,12 +612,56 @@ def test_create_vindex_global_index_from_python(self):
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()),
Expand Down Expand Up @@ -774,7 +820,8 @@ def test_create_native_fulltext_global_index_from_python(self):
('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],
Expand Down Expand Up @@ -1181,5 +1228,141 @@ def test_java_scalar_key_serializers_round_trip(self):
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()
7 changes: 5 additions & 2 deletions paimon-python/pypaimon/tests/vindex_batch_write_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,7 +167,7 @@ def test_builder_filters_ranges_before_vector_validation(self):
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
Expand All @@ -176,10 +176,13 @@ def test_null_row_ids_are_rejected_before_writing_any_batch(self):
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)),
Expand Down
Loading