From fbcc6fdcadb664129ca2280f7102664bc7cd3c5a Mon Sep 17 00:00:00 2001 From: Azwan Date: Tue, 29 Sep 2026 16:26:29 +0800 Subject: [PATCH 1/3] Keep the upsert match filter flat for composite keys An upsert with two or more join columns built one Or disjunct per key. PyArrow flattens that chain and overflows its stack from about 1,000 keys. The scan, the insert filter and the overwrite filter all crashed. Build one In per join column instead. Do the exact key match in Arrow with an anti join. Append back the rows that the overwrite filter removes but does not update. Closes #3508 --- pyiceberg/table/__init__.py | 46 +++++++++------------ pyiceberg/table/upsert_util.py | 35 ++++++++-------- tests/table/test_upsert.py | 75 +++++++++++++++++++++++++++++++++- 3 files changed, 112 insertions(+), 44 deletions(-) diff --git a/pyiceberg/table/__init__.py b/pyiceberg/table/__init__.py index 2c5c26800c..853a780a8c 100644 --- a/pyiceberg/table/__init__.py +++ b/pyiceberg/table/__init__.py @@ -895,7 +895,6 @@ def upsert( except ModuleNotFoundError as e: raise ModuleNotFoundError("For writes PyArrow needs to be installed") from e - from pyiceberg.io.pyarrow import expression_to_pyarrow from pyiceberg.table import upsert_util if join_cols is None: @@ -926,25 +925,16 @@ def upsert( format_version=self.table_metadata.format_version, ) - # get list of rows that exist so we don't have to load the entire target table - matched_predicate = upsert_util.create_match_filter(df, join_cols) - # We must use Transaction.table_metadata for the scan. This includes all uncommitted - but relevant - changes. + def scan_matching(row_filter: BooleanExpression) -> DataScan: + matching_scan = self._scan(row_filter=row_filter, case_sensitive=case_sensitive) + return matching_scan.use_ref(branch) if branch in self.table_metadata.refs else matching_scan - matched_iceberg_record_batches_scan = DataScan( - table_metadata=self.table_metadata, - io=self._table.io, - row_filter=matched_predicate, - case_sensitive=case_sensitive, - ) - - if branch in self.table_metadata.refs: - matched_iceberg_record_batches_scan = matched_iceberg_record_batches_scan.use_ref(branch) - - matched_iceberg_record_batches = matched_iceberg_record_batches_scan.to_arrow_batch_reader() + # get list of rows that exist so we don't have to load the entire target table + # The match filter can select more rows than the keys in df, get_rows_to_update and exclude_keys are exact + matched_iceberg_record_batches = scan_matching(upsert_util.create_match_filter(df, join_cols)).to_arrow_batch_reader() batches_to_overwrite = [] - overwrite_predicates = [] rows_to_insert = df for batch in matched_iceberg_record_batches: @@ -958,19 +948,10 @@ def upsert( rows_to_update = upsert_util.get_rows_to_update(df, rows, join_cols) if len(rows_to_update) > 0: - # build the match predicate filter - overwrite_mask_predicate = upsert_util.create_match_filter(rows_to_update, join_cols) - batches_to_overwrite.append(rows_to_update) - overwrite_predicates.append(overwrite_mask_predicate) if when_not_matched_insert_all: - expr_match = upsert_util.create_match_filter(rows, join_cols) - expr_match_bound = bind(self.table_metadata.schema(), expr_match, case_sensitive=case_sensitive) - expr_match_arrow = expression_to_pyarrow(expr_match_bound) - - # Filter rows per batch. - rows_to_insert = rows_to_insert.filter(~expr_match_arrow) + rows_to_insert = upsert_util.exclude_keys(rows_to_insert, rows, join_cols) update_row_cnt = 0 insert_row_cnt = 0 @@ -978,13 +959,24 @@ def upsert( if batches_to_overwrite: rows_to_update = pa.concat_tables(batches_to_overwrite) update_row_cnt = len(rows_to_update) + overwrite_filter = upsert_util.create_match_filter(rows_to_update, join_cols) + + # For a composite key the overwrite filter can remove rows outside the updated keys, write those back unchanged + unchanged_rows = None + if len(join_cols) > 1: + removed_rows = scan_matching(overwrite_filter).to_arrow() + unchanged_rows = upsert_util.exclude_keys(removed_rows, rows_to_update, join_cols) + self.overwrite( rows_to_update, - overwrite_filter=Or(*overwrite_predicates) if len(overwrite_predicates) > 1 else overwrite_predicates[0], + overwrite_filter=overwrite_filter, branch=branch, snapshot_properties=snapshot_properties, ) + if unchanged_rows is not None and len(unchanged_rows) > 0: + self.append(unchanged_rows, branch=branch, snapshot_properties=snapshot_properties) + if when_not_matched_insert_all: insert_row_cnt = len(rows_to_insert) if rows_to_insert: diff --git a/pyiceberg/table/upsert_util.py b/pyiceberg/table/upsert_util.py index 6f32826eb0..778358ff2e 100644 --- a/pyiceberg/table/upsert_util.py +++ b/pyiceberg/table/upsert_util.py @@ -22,30 +22,33 @@ from pyarrow import compute as pc from pyiceberg.expressions import ( - AlwaysFalse, BooleanExpression, - EqualTo, In, - Or, ) def create_match_filter(df: pyarrow_table, join_cols: list[str]) -> BooleanExpression: + """ + Return a filter that matches every row whose key is present in the given table. + + The filter is one In per join column, so it stays flat for any number of keys. For a composite key + it can also match rows whose key is not present, use exclude_keys to get the exact set afterwards. + See: https://github.com/apache/iceberg-python/issues/3508 + """ unique_keys = df.select(join_cols).group_by(join_cols).aggregate([]) - if len(join_cols) == 1: - return In(join_cols[0], unique_keys[0].to_pylist()) - else: - filters = [ - functools.reduce(operator.and_, [EqualTo(col, row[col]) for col in join_cols]) for row in unique_keys.to_pylist() - ] - - if len(filters) == 0: - return AlwaysFalse() - elif len(filters) == 1: - return filters[0] - else: - return Or(*filters) + return functools.reduce(operator.and_, [In(col, unique_keys[col].to_pylist()) for col in join_cols]) + + +def exclude_keys(table: pa.Table, keys: pa.Table, join_cols: list[str]) -> pa.Table: + """Return the rows of the table whose key is not present in the given keys, in the original order.""" + INDEX_COLUMN_NAME = "__index" + + table_keys = table.select(join_cols) + index = table_keys.append_column(INDEX_COLUMN_NAME, pa.array(range(len(table)))) + remaining = index.join(keys.select(join_cols).cast(table_keys.schema), keys=join_cols, join_type="left anti") + + return table.take(remaining.sort_by(INDEX_COLUMN_NAME)[INDEX_COLUMN_NAME]) def has_duplicate_rows(df: pyarrow_table, join_cols: list[str]) -> bool: diff --git a/tests/table/test_upsert.py b/tests/table/test_upsert.py index 78ddbc7c5c..5f60f9f49a 100644 --- a/tests/table/test_upsert.py +++ b/tests/table/test_upsert.py @@ -24,7 +24,7 @@ from pyiceberg.catalog import Catalog from pyiceberg.exceptions import NoSuchTableError -from pyiceberg.expressions import AlwaysTrue, And, EqualTo, Reference +from pyiceberg.expressions import AlwaysTrue, And, EqualTo, In, Reference from pyiceberg.expressions.literals import LongLiteral from pyiceberg.io.pyarrow import schema_to_pyarrow from pyiceberg.partitioning import PartitionField, PartitionSpec @@ -927,3 +927,76 @@ def test_upsert_snapshot_properties(catalog: Catalog) -> None: for snapshot in snapshots[initial_snapshot_count:]: assert snapshot.summary is not None assert snapshot.summary.additional_properties.get("test_prop") == "test_value" + + +def test_create_match_filter_composite_key_is_flat() -> None: + """ + Test create_match_filter with a composite key and several unique keys. + Expected: One In per key column, not one Or disjunct per key. + """ + schema = pa.schema([pa.field("order_id", pa.int32()), pa.field("order_line_id", pa.int32())]) + table = pa.Table.from_pylist([{"order_id": 101, "order_line_id": 1}, {"order_id": 102, "order_line_id": 2}], schema=schema) + expr = create_match_filter(table, ["order_id", "order_line_id"]) + assert expr == And(In("order_id", [101, 102]), In("order_line_id", [1, 2])) + + +def _composite_key_table(catalog: Catalog, identifier: str) -> tuple[Table, pa.Schema]: + _drop_table(catalog, identifier) + schema = Schema( + NestedField(1, "k1", IntegerType(), required=True), + NestedField(2, "k2", StringType(), required=True), + NestedField(3, "v", IntegerType(), required=True), + identifier_field_ids=[1, 2], + ) + arrow_schema = pa.schema( + [ + pa.field("k1", pa.int32(), nullable=False), + pa.field("k2", pa.string(), nullable=False), + pa.field("v", pa.int32(), nullable=False), + ] + ) + return catalog.create_table(identifier, schema=schema), arrow_schema + + +def test_upsert_composite_key_keeps_rows_outside_source_keys(catalog: Catalog) -> None: + # Every k1 and every k2 of the source exists in the target, but only two of the key tuples do + tbl, arrow_schema = _composite_key_table(catalog, "default.test_upsert_composite_key_keeps_rows_outside_source_keys") + tbl.append( + pa.Table.from_pylist( + [ + {"k1": 1, "k2": "a", "v": 1}, + {"k1": 1, "k2": "b", "v": 2}, + {"k1": 2, "k2": "a", "v": 3}, + {"k1": 2, "k2": "b", "v": 4}, + ], + schema=arrow_schema, + ) + ) + source = pa.Table.from_pylist( + [{"k1": 1, "k2": "b", "v": 20}, {"k1": 2, "k2": "a", "v": 30}, {"k1": 2, "k2": "c", "v": 5}], schema=arrow_schema + ) + + res = tbl.upsert(source) + + assert (res.rows_updated, res.rows_inserted) == (2, 1) + assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda r: (r["k1"], r["k2"])) == [ + {"k1": 1, "k2": "a", "v": 1}, + {"k1": 1, "k2": "b", "v": 20}, + {"k1": 2, "k2": "a", "v": 30}, + {"k1": 2, "k2": "b", "v": 4}, + {"k1": 2, "k2": "c", "v": 5}, + ] + + +def test_upsert_composite_key_large_batch(catalog: Catalog) -> None: + # One Or disjunct per key tuple overflows the Arrow expression stack from about 1,000 tuples (#3508) + tbl, arrow_schema = _composite_key_table(catalog, "default.test_upsert_composite_key_large_batch") + tbl.append(pa.Table.from_pylist([{"k1": i, "k2": f"k{i}", "v": i} for i in range(5000)], schema=arrow_schema)) + source = pa.Table.from_pylist([{"k1": i, "k2": f"k{i}", "v": i + 10} for i in range(2500, 7500)], schema=arrow_schema) + + res = tbl.upsert(source) + + assert (res.rows_updated, res.rows_inserted) == (2500, 2500) + assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda r: r["k1"]) == [ + {"k1": i, "k2": f"k{i}", "v": i if i < 2500 else i + 10} for i in range(7500) + ] From e51fde6c9630c911bdd441edb9082cd755643004 Mon Sep 17 00:00:00 2001 From: Azwan Date: Tue, 29 Sep 2026 16:35:11 +0800 Subject: [PATCH 2/3] Type the index column in exclude_keys so an empty table joins An empty table gave the index column the null type and the anti join rejected it. The insert path reaches this when one scan batch removes every source key and a later batch still matches the filter. --- pyiceberg/table/upsert_util.py | 2 +- tests/table/test_upsert.py | 17 +++++++++++++++++ 2 files changed, 18 insertions(+), 1 deletion(-) diff --git a/pyiceberg/table/upsert_util.py b/pyiceberg/table/upsert_util.py index 778358ff2e..8826f4d823 100644 --- a/pyiceberg/table/upsert_util.py +++ b/pyiceberg/table/upsert_util.py @@ -45,7 +45,7 @@ def exclude_keys(table: pa.Table, keys: pa.Table, join_cols: list[str]) -> pa.Ta INDEX_COLUMN_NAME = "__index" table_keys = table.select(join_cols) - index = table_keys.append_column(INDEX_COLUMN_NAME, pa.array(range(len(table)))) + index = table_keys.append_column(INDEX_COLUMN_NAME, pa.array(range(len(table)), pa.int64())) remaining = index.join(keys.select(join_cols).cast(table_keys.schema), keys=join_cols, join_type="left anti") return table.take(remaining.sort_by(INDEX_COLUMN_NAME)[INDEX_COLUMN_NAME]) diff --git a/tests/table/test_upsert.py b/tests/table/test_upsert.py index 5f60f9f49a..0043bc7610 100644 --- a/tests/table/test_upsert.py +++ b/tests/table/test_upsert.py @@ -1000,3 +1000,20 @@ def test_upsert_composite_key_large_batch(catalog: Catalog) -> None: assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda r: r["k1"]) == [ {"k1": i, "k2": f"k{i}", "v": i if i < 2500 else i + 10} for i in range(7500) ] + + +def test_upsert_composite_key_all_source_keys_matched_across_files(catalog: Catalog) -> None: + # The first file removes every source key from the insert set, the second file still matches the superset filter + tbl, arrow_schema = _composite_key_table(catalog, "default.test_upsert_composite_key_all_source_keys_matched_across_files") + tbl.append(pa.Table.from_pylist([{"k1": 1, "k2": "b", "v": 1}], schema=arrow_schema)) + tbl.append(pa.Table.from_pylist([{"k1": 1, "k2": "a", "v": 2}, {"k1": 2, "k2": "b", "v": 3}], schema=arrow_schema)) + source = pa.Table.from_pylist([{"k1": 1, "k2": "a", "v": 20}, {"k1": 2, "k2": "b", "v": 3}], schema=arrow_schema) + + res = tbl.upsert(source) + + assert (res.rows_updated, res.rows_inserted) == (1, 0) + assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda r: (r["k1"], r["k2"])) == [ + {"k1": 1, "k2": "a", "v": 20}, + {"k1": 1, "k2": "b", "v": 1}, + {"k1": 2, "k2": "b", "v": 3}, + ] From 9c0e0fb566243b697c67d3e938ca5265d1e018fc Mon Sep 17 00:00:00 2001 From: Azwan Date: Tue, 29 Sep 2026 16:55:12 +0800 Subject: [PATCH 3/3] Reject the reserved index column name in exclude_keys Match the check that get_rows_to_update already does, so a join column named __index gives a clear error instead of an Arrow join failure. --- pyiceberg/table/upsert_util.py | 3 +++ tests/table/test_upsert.py | 13 ++++++++----- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/pyiceberg/table/upsert_util.py b/pyiceberg/table/upsert_util.py index 8826f4d823..2ecf0b138e 100644 --- a/pyiceberg/table/upsert_util.py +++ b/pyiceberg/table/upsert_util.py @@ -44,6 +44,9 @@ def exclude_keys(table: pa.Table, keys: pa.Table, join_cols: list[str]) -> pa.Ta """Return the rows of the table whose key is not present in the given keys, in the original order.""" INDEX_COLUMN_NAME = "__index" + if INDEX_COLUMN_NAME in join_cols: + raise ValueError(f"{INDEX_COLUMN_NAME} is reserved for joining DataFrames, and cannot be used as a column name") + table_keys = table.select(join_cols) index = table_keys.append_column(INDEX_COLUMN_NAME, pa.array(range(len(table)), pa.int64())) remaining = index.join(keys.select(join_cols).cast(table_keys.schema), keys=join_cols, join_type="left anti") diff --git a/tests/table/test_upsert.py b/tests/table/test_upsert.py index 0043bc7610..58c0877a3d 100644 --- a/tests/table/test_upsert.py +++ b/tests/table/test_upsert.py @@ -31,7 +31,7 @@ from pyiceberg.schema import Schema from pyiceberg.table import Table, UpsertResult from pyiceberg.table.snapshots import Operation -from pyiceberg.table.upsert_util import create_match_filter +from pyiceberg.table.upsert_util import create_match_filter, exclude_keys from pyiceberg.transforms import DayTransform from pyiceberg.types import IntegerType, NestedField, StringType, StructType, TimestampType from tests.catalog.test_base import InMemoryCatalog @@ -930,10 +930,7 @@ def test_upsert_snapshot_properties(catalog: Catalog) -> None: def test_create_match_filter_composite_key_is_flat() -> None: - """ - Test create_match_filter with a composite key and several unique keys. - Expected: One In per key column, not one Or disjunct per key. - """ + # One In per key column, not one Or disjunct per key schema = pa.schema([pa.field("order_id", pa.int32()), pa.field("order_line_id", pa.int32())]) table = pa.Table.from_pylist([{"order_id": 101, "order_line_id": 1}, {"order_id": 102, "order_line_id": 2}], schema=schema) expr = create_match_filter(table, ["order_id", "order_line_id"]) @@ -1017,3 +1014,9 @@ def test_upsert_composite_key_all_source_keys_matched_across_files(catalog: Cata {"k1": 1, "k2": "b", "v": 1}, {"k1": 2, "k2": "b", "v": 3}, ] + + +def test_exclude_keys_rejects_reserved_column_name() -> None: + table = pa.table({"__index": [1], "k": ["a"]}) + with pytest.raises(ValueError, match="__index is reserved"): + exclude_keys(table, table, ["__index", "k"])