Skip to content

Commit af9c2e2

Browse files
committed
Fix upsert after schema evolution (#3105)
* Cast only join_cols schema instead of full table schema in get_rows_to_update * Safely handle missing non-key columns in target_table when comparing rows * Add regression tests for upsert after add_column and union_by_name schema evolution * Verify underlying Parquet file replacement and snapshot operations
1 parent 68898e5 commit af9c2e2

2 files changed

Lines changed: 245 additions & 8 deletions

File tree

‎pyiceberg/table/upsert_util.py‎

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -62,8 +62,8 @@ def get_rows_to_update(source_table: pa.Table, target_table: pa.Table, join_cols
6262
"""
6363
all_columns = set(source_table.column_names)
6464
join_cols_set = set(join_cols)
65-
6665
non_key_cols = list(all_columns - join_cols_set)
66+
target_columns = set(target_table.column_names)
6767

6868
if has_duplicate_rows(target_table, join_cols):
6969
raise ValueError("Target table has duplicate rows, aborting upsert")
@@ -86,19 +86,20 @@ def get_rows_to_update(source_table: pa.Table, target_table: pa.Table, join_cols
8686
) from None
8787

8888
# Step 1: Prepare source index with join keys and a marker index
89-
# Cast to target table schema, so we can do the join
89+
# Cast join columns to target table schema, so we can do the join
9090
# See: https://github.com/apache/arrow/issues/37542
91+
join_schema = pa.schema([target_table.schema.field(col) for col in join_cols])
9192
source_index = (
92-
source_table.cast(target_table.schema)
93-
.select(join_cols_set)
93+
source_table.select(join_cols)
94+
.cast(join_schema)
9495
.append_column(SOURCE_INDEX_COLUMN_NAME, pa.array(range(len(source_table))))
9596
)
9697

9798
# Step 2: Prepare target index with join keys and a marker
98-
target_index = target_table.select(join_cols_set).append_column(TARGET_INDEX_COLUMN_NAME, pa.array(range(len(target_table))))
99+
target_index = target_table.select(join_cols).append_column(TARGET_INDEX_COLUMN_NAME, pa.array(range(len(target_table))))
99100

100101
# Step 3: Perform an inner join to find which rows from source exist in target
101-
matching_indices = source_index.join(target_index, keys=list(join_cols_set), join_type="inner")
102+
matching_indices = source_index.join(target_index, keys=join_cols, join_type="inner")
102103

103104
# Step 4: Compare all rows using Python
104105
to_update_indices = []
@@ -112,7 +113,7 @@ def get_rows_to_update(source_table: pa.Table, target_table: pa.Table, join_cols
112113

113114
for key in non_key_cols:
114115
source_val = source_row.column(key)[0].as_py()
115-
target_val = target_row.column(key)[0].as_py()
116+
target_val = target_row.column(key)[0].as_py() if key in target_columns else None
116117
if source_val != target_val:
117118
to_update_indices.append(source_idx)
118119
break

‎tests/table/test_upsert.py‎

Lines changed: 237 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@
3333
from pyiceberg.table.snapshots import Operation
3434
from pyiceberg.table.upsert_util import create_match_filter
3535
from pyiceberg.transforms import DayTransform
36-
from pyiceberg.types import IntegerType, NestedField, StringType, StructType, TimestampType
36+
from pyiceberg.types import IntegerType, LongType, NestedField, StringType, StructType, TimestampType
3737
from tests.catalog.test_base import InMemoryCatalog
3838

3939

@@ -927,3 +927,239 @@ def test_upsert_snapshot_properties(catalog: Catalog) -> None:
927927
for snapshot in snapshots[initial_snapshot_count:]:
928928
assert snapshot.summary is not None
929929
assert snapshot.summary.additional_properties.get("test_prop") == "test_value"
930+
931+
932+
def test_upsert_after_schema_evolution(catalog: Catalog) -> None:
933+
identifier = "default.test_upsert_after_schema_evolution"
934+
_drop_table(catalog, identifier)
935+
936+
schema = Schema(
937+
NestedField(1, "id", LongType(), required=True),
938+
NestedField(2, "name", StringType(), required=False),
939+
identifier_field_ids=[1],
940+
)
941+
tbl = catalog.create_table(identifier, schema=schema)
942+
943+
# Initial write with 2 columns
944+
arrow_schema_v0 = pa.schema(
945+
[
946+
pa.field("id", pa.int64(), nullable=False),
947+
pa.field("name", pa.string(), nullable=True),
948+
]
949+
)
950+
df_v0 = pa.Table.from_pylist([{"id": 1, "name": "Alice"}, {"id": 2, "name": "Bob"}], schema=arrow_schema_v0)
951+
tbl.append(df_v0)
952+
953+
# Record initial data file before evolution & upsert
954+
initial_files = [f.file.file_path for f in tbl.scan().plan_files()]
955+
assert len(initial_files) == 1
956+
957+
# Schema evolution: add column 'city'
958+
with tbl.update_schema() as update:
959+
update.add_column("city", StringType())
960+
961+
# Upsert with 3 columns: update id=1 with new city, insert id=3
962+
arrow_schema_v1 = pa.schema(
963+
[
964+
pa.field("id", pa.int64(), nullable=False),
965+
pa.field("name", pa.string(), nullable=True),
966+
pa.field("city", pa.string(), nullable=True),
967+
]
968+
)
969+
df_v1 = pa.Table.from_pylist(
970+
[
971+
{"id": 1, "name": "Alice", "city": "NYC"},
972+
{"id": 3, "name": "Charlie", "city": "LA"},
973+
],
974+
schema=arrow_schema_v1,
975+
)
976+
result = tbl.upsert(df_v1)
977+
assert result.rows_updated == 1
978+
assert result.rows_inserted == 1
979+
980+
# Verify that the old V0 data file was replaced (Copy-on-Write overwrite)
981+
current_files = [f.file.file_path for f in tbl.scan().plan_files()]
982+
assert initial_files[0] not in current_files
983+
984+
# Verify snapshot operations include both OVERWRITE and APPEND
985+
operations = [s.summary.operation for s in tbl.snapshots() if s.summary is not None]
986+
assert Operation.OVERWRITE in operations
987+
assert Operation.APPEND in operations
988+
989+
# Verify scanned table rows contain updated evolved values
990+
scanned_rows = tbl.scan().to_arrow().to_pylist()
991+
assert sorted(scanned_rows, key=lambda x: x["id"]) == [
992+
{"id": 1, "name": "Alice", "city": "NYC"},
993+
{"id": 2, "name": "Bob", "city": None},
994+
{"id": 3, "name": "Charlie", "city": "LA"},
995+
]
996+
997+
998+
def test_upsert_after_schema_evolution_union_by_name(catalog: Catalog) -> None:
999+
identifier = "default.test_upsert_after_schema_evolution_union_by_name"
1000+
_drop_table(catalog, identifier)
1001+
1002+
arrow_schema_v0 = pa.schema(
1003+
[
1004+
pa.field("id", pa.int64(), nullable=False),
1005+
pa.field("name", pa.string(), nullable=True),
1006+
pa.field("age", pa.int32(), nullable=True),
1007+
]
1008+
)
1009+
schema = Schema(
1010+
NestedField(1, "id", LongType(), required=True),
1011+
NestedField(2, "name", StringType(), required=False),
1012+
NestedField(3, "age", IntegerType(), required=False),
1013+
identifier_field_ids=[1],
1014+
)
1015+
tbl = catalog.create_table(identifier, schema=schema)
1016+
1017+
df_v0 = pa.Table.from_pylist(
1018+
[
1019+
{"id": 1, "name": "Alice", "age": 30},
1020+
{"id": 2, "name": "Bob", "age": 25},
1021+
],
1022+
schema=arrow_schema_v0,
1023+
)
1024+
tbl.append(df_v0)
1025+
1026+
# Schema evolution via union_by_name with a new column 'city'
1027+
arrow_schema_v1 = pa.schema(
1028+
[
1029+
pa.field("id", pa.int64(), nullable=False),
1030+
pa.field("name", pa.string(), nullable=True),
1031+
pa.field("age", pa.int32(), nullable=True),
1032+
pa.field("city", pa.string(), nullable=True),
1033+
]
1034+
)
1035+
with tbl.update_schema() as update:
1036+
update.union_by_name(arrow_schema_v1)
1037+
1038+
# Upsert with 4 columns: update id=1 (change age & city), unchanged id=2, insert id=3
1039+
df_v1 = pa.Table.from_pylist(
1040+
[
1041+
{"id": 1, "name": "Alice", "age": 31, "city": "Taipei"},
1042+
{"id": 2, "name": "Bob", "age": 25, "city": None},
1043+
{"id": 3, "name": "Charlie", "age": 40, "city": "Tokyo"},
1044+
],
1045+
schema=arrow_schema_v1,
1046+
)
1047+
result = tbl.upsert(df_v1)
1048+
assert result.rows_updated == 1
1049+
assert result.rows_inserted == 1
1050+
1051+
scanned_rows = tbl.scan().to_arrow().to_pylist()
1052+
assert sorted(scanned_rows, key=lambda x: x["id"]) == [
1053+
{"id": 1, "name": "Alice", "age": 31, "city": "Taipei"},
1054+
{"id": 2, "name": "Bob", "age": 25, "city": None},
1055+
{"id": 3, "name": "Charlie", "age": 40, "city": "Tokyo"},
1056+
]
1057+
1058+
1059+
def test_upsert_after_multiple_schema_evolutions_with_composite_keys(catalog: Catalog) -> None:
1060+
identifier = "default.test_upsert_after_multiple_schema_evolutions_with_composite_keys"
1061+
_drop_table(catalog, identifier)
1062+
1063+
# Step 1: Create table with V0 schema (id, dept, name)
1064+
schema = Schema(
1065+
NestedField(1, "id", LongType(), required=True),
1066+
NestedField(2, "dept", StringType(), required=True),
1067+
NestedField(3, "name", StringType(), required=False),
1068+
identifier_field_ids=[1, 2],
1069+
)
1070+
tbl = catalog.create_table(identifier, schema=schema)
1071+
1072+
arrow_schema_v0 = pa.schema(
1073+
[
1074+
pa.field("id", pa.int64(), nullable=False),
1075+
pa.field("dept", pa.string(), nullable=False),
1076+
pa.field("name", pa.string(), nullable=True),
1077+
]
1078+
)
1079+
tbl.append(pa.Table.from_pylist([{"id": 1, "dept": "ENG", "name": "Alice"}], schema=arrow_schema_v0))
1080+
1081+
# Step 2: Evolve to V1 by adding 'salary'
1082+
with tbl.update_schema() as update:
1083+
update.add_column("salary", LongType())
1084+
1085+
arrow_schema_v1 = pa.schema(
1086+
[
1087+
pa.field("id", pa.int64(), nullable=False),
1088+
pa.field("dept", pa.string(), nullable=False),
1089+
pa.field("name", pa.string(), nullable=True),
1090+
pa.field("salary", pa.int64(), nullable=True),
1091+
]
1092+
)
1093+
tbl.append(pa.Table.from_pylist([{"id": 2, "dept": "HR", "name": "Bob", "salary": 50000}], schema=arrow_schema_v1))
1094+
1095+
# Step 3: Evolve to V2 by adding 'city'
1096+
with tbl.update_schema() as update:
1097+
update.add_column("city", StringType())
1098+
1099+
arrow_schema_v2 = pa.schema(
1100+
[
1101+
pa.field("id", pa.int64(), nullable=False),
1102+
pa.field("dept", pa.string(), nullable=False),
1103+
pa.field("name", pa.string(), nullable=True),
1104+
pa.field("salary", pa.int64(), nullable=True),
1105+
pa.field("city", pa.string(), nullable=True),
1106+
]
1107+
)
1108+
1109+
# Step 4: Upsert spanning V0, V1, and V2 rows
1110+
df_v2 = pa.Table.from_pylist(
1111+
[
1112+
{"id": 1, "dept": "ENG", "name": "Alice", "salary": 80000, "city": "Taipei"}, # Update V0 row (salary + city added)
1113+
{"id": 2, "dept": "HR", "name": "Bob", "salary": 50000, "city": "London"}, # Update V1 row (city added)
1114+
{"id": 3, "dept": "MKT", "name": "Charlie", "salary": 60000, "city": "Tokyo"}, # Insert new V2 row
1115+
],
1116+
schema=arrow_schema_v2,
1117+
)
1118+
result = tbl.upsert(df_v2)
1119+
assert result.rows_updated == 2
1120+
assert result.rows_inserted == 1
1121+
1122+
scanned_rows = tbl.scan().to_arrow().to_pylist()
1123+
assert sorted(scanned_rows, key=lambda x: (x["id"], x["dept"])) == [
1124+
{"id": 1, "dept": "ENG", "name": "Alice", "salary": 80000, "city": "Taipei"},
1125+
{"id": 2, "dept": "HR", "name": "Bob", "salary": 50000, "city": "London"},
1126+
{"id": 3, "dept": "MKT", "name": "Charlie", "salary": 60000, "city": "Tokyo"},
1127+
]
1128+
1129+
1130+
def test_upsert_after_schema_evolution_noop_and_nulls(catalog: Catalog) -> None:
1131+
identifier = "default.test_upsert_after_schema_evolution_noop_and_nulls"
1132+
_drop_table(catalog, identifier)
1133+
1134+
schema = Schema(
1135+
NestedField(1, "id", LongType(), required=True),
1136+
NestedField(2, "name", StringType(), required=False),
1137+
identifier_field_ids=[1],
1138+
)
1139+
tbl = catalog.create_table(identifier, schema=schema)
1140+
1141+
arrow_schema_v0 = pa.schema(
1142+
[
1143+
pa.field("id", pa.int64(), nullable=False),
1144+
pa.field("name", pa.string(), nullable=True),
1145+
]
1146+
)
1147+
tbl.append(pa.Table.from_pylist([{"id": 1, "name": "Alice"}], schema=arrow_schema_v0))
1148+
1149+
# Evolve schema
1150+
with tbl.update_schema() as update:
1151+
update.add_column("extra", StringType())
1152+
1153+
arrow_schema_v1 = pa.schema(
1154+
[
1155+
pa.field("id", pa.int64(), nullable=False),
1156+
pa.field("name", pa.string(), nullable=True),
1157+
pa.field("extra", pa.string(), nullable=True),
1158+
]
1159+
)
1160+
1161+
# Upsert with identical row where new column is None -> should be no-op (0 updated, 0 inserted)
1162+
df_noop = pa.Table.from_pylist([{"id": 1, "name": "Alice", "extra": None}], schema=arrow_schema_v1)
1163+
result = tbl.upsert(df_noop)
1164+
assert result.rows_updated == 0
1165+
assert result.rows_inserted == 0

0 commit comments

Comments
 (0)