Skip to content

Commit c3754e6

Browse files
committed
IO: Dissolve upsert_util.py into io/pyarrow.py with deprecation shim
1 parent 73392ef commit c3754e6

4 files changed

Lines changed: 72 additions & 46 deletions

File tree

‎pyiceberg/io/pyarrow.py‎

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,18 @@
6868

6969
from pyiceberg.conversions import to_bytes
7070
from pyiceberg.exceptions import ResolveError
71-
from pyiceberg.expressions import AlwaysTrue, BooleanExpression, BoundIsNaN, BoundIsNull, BoundTerm, Not, Or
71+
from pyiceberg.expressions import (
72+
AlwaysFalse,
73+
AlwaysTrue,
74+
BooleanExpression,
75+
BoundIsNaN,
76+
BoundIsNull,
77+
BoundTerm,
78+
EqualTo,
79+
In,
80+
Not,
81+
Or,
82+
)
7283
from pyiceberg.expressions.literals import Literal
7384
from pyiceberg.expressions.visitors import (
7485
BoundBooleanExpressionVisitor,
@@ -3147,6 +3158,25 @@ def upsert_unique_keys(df: pa.Table, join_cols: list[str]) -> pa.Table:
31473158
return df.select(join_cols).group_by(join_cols).aggregate([])
31483159

31493160

3161+
def upsert_create_match_filter(df: pa.Table, join_cols: list[str]) -> BooleanExpression:
3162+
"""Build an Iceberg filter expression matching the unique keys in df."""
3163+
unique_keys = upsert_unique_keys(df, join_cols)
3164+
3165+
if len(join_cols) == 1:
3166+
return In(join_cols[0], unique_keys[0].to_pylist())
3167+
else:
3168+
filters = [
3169+
functools.reduce(operator.and_, [EqualTo(col, row[col]) for col in join_cols]) for row in unique_keys.to_pylist()
3170+
]
3171+
3172+
if len(filters) == 0:
3173+
return AlwaysFalse()
3174+
elif len(filters) == 1:
3175+
return filters[0]
3176+
else:
3177+
return Or(*filters)
3178+
3179+
31503180
def upsert_has_duplicate_rows(df: pa.Table, join_cols: list[str]) -> bool:
31513181
"""Check for duplicate rows in a PyArrow table based on the join columns."""
31523182
return len(df.select(join_cols).group_by(join_cols).aggregate([([], "count_all")]).filter(pc.field("count_all") > 1)) > 0

‎pyiceberg/table/__init__.py‎

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -892,8 +892,12 @@ def upsert(
892892
except ModuleNotFoundError as e:
893893
raise ModuleNotFoundError("For writes PyArrow needs to be installed") from e
894894

895-
from pyiceberg.io.pyarrow import expression_to_pyarrow
896-
from pyiceberg.table import upsert_util
895+
from pyiceberg.io.pyarrow import (
896+
expression_to_pyarrow,
897+
upsert_create_match_filter,
898+
upsert_get_rows_to_update,
899+
upsert_has_duplicate_rows,
900+
)
897901

898902
if join_cols is None:
899903
join_cols = []
@@ -910,7 +914,7 @@ def upsert(
910914
if not when_matched_update_all and not when_not_matched_insert_all:
911915
raise ValueError("no upsert options selected...exiting")
912916

913-
if upsert_util.has_duplicate_rows(df, join_cols):
917+
if upsert_has_duplicate_rows(df, join_cols):
914918
raise ValueError("Duplicate rows found in source dataset based on the key columns. No upsert executed")
915919

916920
from pyiceberg.io.pyarrow import _check_pyarrow_schema_compatible
@@ -924,7 +928,7 @@ def upsert(
924928
)
925929

926930
# get list of rows that exist so we don't have to load the entire target table
927-
matched_predicate = upsert_util.create_match_filter(df, join_cols)
931+
matched_predicate = upsert_create_match_filter(df, join_cols)
928932

929933
# We must use Transaction.table_metadata for the scan. This includes all uncommitted - but relevant - changes.
930934

@@ -952,17 +956,17 @@ def upsert(
952956
# values have actually changed. We don't want to do just a blanket overwrite for matched
953957
# rows if the actual non-key column data hasn't changed.
954958
# this extra step avoids unnecessary IO and writes
955-
rows_to_update = upsert_util.get_rows_to_update(df, rows, join_cols)
959+
rows_to_update = upsert_get_rows_to_update(df, rows, join_cols)
956960

957961
if len(rows_to_update) > 0:
958962
# build the match predicate filter
959-
overwrite_mask_predicate = upsert_util.create_match_filter(rows_to_update, join_cols)
963+
overwrite_mask_predicate = upsert_create_match_filter(rows_to_update, join_cols)
960964

961965
batches_to_overwrite.append(rows_to_update)
962966
overwrite_predicates.append(overwrite_mask_predicate)
963967

964968
if when_not_matched_insert_all:
965-
expr_match = upsert_util.create_match_filter(rows, join_cols)
969+
expr_match = upsert_create_match_filter(rows, join_cols)
966970
expr_match_bound = bind(self.table_metadata.schema(), expr_match, case_sensitive=case_sensitive)
967971
expr_match_arrow = expression_to_pyarrow(expr_match_bound)
968972

@@ -2663,8 +2667,9 @@ def plan_files(self) -> Iterable[FileScanTask]:
26632667
options=self.options,
26642668
).plan_files(
26652669
manifests=manifests,
2666-
manifest_entry_filter=lambda manifest_entry: manifest_entry.snapshot_id in append_snapshot_ids
2667-
and manifest_entry.status == ManifestEntryStatus.ADDED,
2670+
manifest_entry_filter=lambda manifest_entry: (
2671+
manifest_entry.snapshot_id in append_snapshot_ids and manifest_entry.status == ManifestEntryStatus.ADDED
2672+
),
26682673
)
26692674

26702675
def to_arrow(self) -> pa.Table:

‎pyiceberg/table/upsert_util.py‎

Lines changed: 25 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -14,55 +14,47 @@
1414
# KIND, either express or implied. See the License for the
1515
# specific language governing permissions and limitations
1616
# under the License.
17-
import functools
18-
import operator
17+
18+
"""Deprecated: upsert helpers have moved to pyiceberg.io.pyarrow.
19+
20+
All functions in this module are re-exported from ``pyiceberg.io.pyarrow``
21+
and will emit a ``DeprecationWarning`` when called. Import directly from
22+
``pyiceberg.io.pyarrow`` instead.
23+
"""
24+
25+
from __future__ import annotations
26+
1927
from typing import TYPE_CHECKING
2028

21-
from pyiceberg.expressions import (
22-
AlwaysFalse,
23-
BooleanExpression,
24-
EqualTo,
25-
In,
26-
Or,
27-
)
29+
from pyiceberg.expressions import BooleanExpression
2830
from pyiceberg.io.pyarrow import (
31+
upsert_create_match_filter,
2932
upsert_get_rows_to_update,
3033
upsert_has_duplicate_rows,
31-
upsert_unique_keys,
3234
)
35+
from pyiceberg.utils.deprecated import deprecated
3336

3437
if TYPE_CHECKING:
3538
import pyarrow as pa
3639

40+
_DEPRECATION_IN = "0.13.0"
41+
_REMOVAL_IN = "0.14.0"
42+
_HELP = "Use the equivalent function from pyiceberg.io.pyarrow instead"
3743

38-
def create_match_filter(df: "pa.Table", join_cols: list[str]) -> BooleanExpression:
39-
"""Build an Iceberg filter expression matching the unique keys in df."""
40-
unique_keys = upsert_unique_keys(df, join_cols)
41-
42-
if len(join_cols) == 1:
43-
return In(join_cols[0], unique_keys[0].to_pylist())
44-
else:
45-
filters = [
46-
functools.reduce(operator.and_, [EqualTo(col, row[col]) for col in join_cols]) for row in unique_keys.to_pylist()
47-
]
4844

49-
if len(filters) == 0:
50-
return AlwaysFalse()
51-
elif len(filters) == 1:
52-
return filters[0]
53-
else:
54-
return Or(*filters)
45+
@deprecated(deprecated_in=_DEPRECATION_IN, removed_in=_REMOVAL_IN, help_message=_HELP)
46+
def create_match_filter(df: pa.Table, join_cols: list[str]) -> BooleanExpression:
47+
"""Build an Iceberg filter expression matching the unique keys in df."""
48+
return upsert_create_match_filter(df, join_cols)
5549

5650

57-
def has_duplicate_rows(df: "pa.Table", join_cols: list[str]) -> bool:
51+
@deprecated(deprecated_in=_DEPRECATION_IN, removed_in=_REMOVAL_IN, help_message=_HELP)
52+
def has_duplicate_rows(df: pa.Table, join_cols: list[str]) -> bool:
5853
"""Check for duplicate rows in a table based on the join columns."""
5954
return upsert_has_duplicate_rows(df, join_cols)
6055

6156

62-
def get_rows_to_update(source_table: "pa.Table", target_table: "pa.Table", join_cols: list[str]) -> "pa.Table":
63-
"""Return rows from source that need to be updated in the target table based on the join columns.
64-
65-
The table is joined on the identifier columns, and then checked if there are any updated rows.
66-
Those are selected and everything is renamed correctly.
67-
"""
57+
@deprecated(deprecated_in=_DEPRECATION_IN, removed_in=_REMOVAL_IN, help_message=_HELP)
58+
def get_rows_to_update(source_table: pa.Table, target_table: pa.Table, join_cols: list[str]) -> pa.Table:
59+
"""Return rows from source that need to be updated in the target table."""
6860
return upsert_get_rows_to_update(source_table, target_table, join_cols)

‎tests/table/test_upsert.py‎

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -26,12 +26,11 @@
2626
from pyiceberg.exceptions import NoSuchTableError
2727
from pyiceberg.expressions import AlwaysTrue, And, EqualTo, Reference
2828
from pyiceberg.expressions.literals import LongLiteral
29-
from pyiceberg.io.pyarrow import schema_to_pyarrow
29+
from pyiceberg.io.pyarrow import schema_to_pyarrow, upsert_create_match_filter
3030
from pyiceberg.partitioning import PartitionField, PartitionSpec
3131
from pyiceberg.schema import Schema
3232
from pyiceberg.table import Table, UpsertResult
3333
from pyiceberg.table.snapshots import Operation
34-
from pyiceberg.table.upsert_util import create_match_filter
3534
from pyiceberg.transforms import DayTransform
3635
from pyiceberg.types import IntegerType, NestedField, StringType, StructType, TimestampType
3736
from tests.catalog.test_base import InMemoryCatalog
@@ -439,7 +438,7 @@ def test_create_match_filter_single_condition() -> None:
439438
]
440439
schema = pa.schema([pa.field("order_id", pa.int32()), pa.field("order_line_id", pa.int32()), pa.field("extra", pa.string())])
441440
table = pa.Table.from_pylist(data, schema=schema)
442-
expr = create_match_filter(table, ["order_id", "order_line_id"])
441+
expr = upsert_create_match_filter(table, ["order_id", "order_line_id"])
443442
assert expr == And(
444443
EqualTo(term=Reference(name="order_id"), literal=LongLiteral(101)),
445444
EqualTo(term=Reference(name="order_line_id"), literal=LongLiteral(1)),

0 commit comments

Comments
 (0)