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
237 changes: 237 additions & 0 deletions paimon-python/pypaimon/tests/table_upsert_by_key_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.writer.append_only_data_writer import AppendOnlyDataWriter
from pypaimon.tests.data_evolution_test_helpers import (
BatchModeMixin,
DataEvolutionTestBase,
Expand Down Expand Up @@ -61,6 +65,184 @@ 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 = AppendOnlyDataWriter(table, (), 0, 0, table.options)
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(AppendOnlyDataWriter, '_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()

@mock.patch.object(AppendOnlyDataWriter, '_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
# ------------------------------------------------------------------
Expand Down Expand Up @@ -953,6 +1135,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
# ==================================================================
Expand Down
1 change: 1 addition & 0 deletions paimon-python/pypaimon/write/table_update.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
Loading
Loading