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
46 changes: 19 additions & 27 deletions pyiceberg/table/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -958,33 +948,35 @@ 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

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:
Expand Down
38 changes: 22 additions & 16 deletions pyiceberg/table/upsert_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,30 +22,36 @@
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"

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")

return table.take(remaining.sort_by(INDEX_COLUMN_NAME)[INDEX_COLUMN_NAME])


def has_duplicate_rows(df: pyarrow_table, join_cols: list[str]) -> bool:
Expand Down
97 changes: 95 additions & 2 deletions tests/table/test_upsert.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,14 +24,14 @@

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
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
Expand Down Expand Up @@ -927,3 +927,96 @@ 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:
# 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)
]


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},
]


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"])
Loading