Skip to content

Commit 07461b7

Browse files
committed
fix: use the current schema when reading matched rows in upsert
1 parent 7de3484 commit 07461b7

2 files changed

Lines changed: 163 additions & 1 deletion

File tree

‎pyiceberg/table/__init__.py‎

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -938,7 +938,16 @@ def upsert(
938938
if branch in self.table_metadata.refs:
939939
matched_iceberg_record_batches_scan = matched_iceberg_record_batches_scan.use_ref(branch)
940940

941-
matched_iceberg_record_batches = matched_iceberg_record_batches_scan.to_arrow_batch_reader()
941+
# Project the current schema instead of DataScan.projection(). Pinning a ref sets the
942+
# snapshot id, which makes projection() fall back to that snapshot's historical schema.
943+
# A schema-only update does not create a data snapshot, so the branch tip can still carry
944+
# an older schema, and the matched rows would then be missing the newly added columns that
945+
# the input dataframe has.
946+
matched_iceberg_record_batches = _to_arrow_batch_reader_via_file_scan_tasks(
947+
matched_iceberg_record_batches_scan,
948+
self.table_metadata.schema(),
949+
matched_iceberg_record_batches_scan.plan_files(),
950+
)
942951

943952
batches_to_overwrite = []
944953
overwrite_predicates = []

‎tests/table/test_upsert.py‎

Lines changed: 153 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -927,3 +927,156 @@ 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_adding_column(catalog: Catalog) -> None:
933+
"""Rows written before an added column must still be comparable against the current schema."""
934+
identifier = "default.test_upsert_after_adding_column"
935+
_drop_table(catalog, identifier)
936+
937+
schema = Schema(
938+
NestedField(1, "city", StringType(), required=True),
939+
NestedField(2, "population", IntegerType(), required=True),
940+
identifier_field_ids=[1],
941+
)
942+
tbl = catalog.create_table(identifier, schema=schema)
943+
944+
tbl.append(
945+
pa.Table.from_pylist(
946+
[{"city": "Amsterdam", "population": 921402}],
947+
schema=schema_to_pyarrow(tbl.schema()),
948+
)
949+
)
950+
951+
# A schema-only update does not create a data snapshot, so the branch tip keeps the old schema
952+
with tbl.update_schema() as update:
953+
update.add_column("country", StringType())
954+
955+
df = pa.Table.from_pylist(
956+
[
957+
{"city": "Amsterdam", "population": 950000, "country": "NL"},
958+
{"city": "Berlin", "population": 3432000, "country": "DE"},
959+
],
960+
schema=schema_to_pyarrow(tbl.schema()),
961+
)
962+
result = tbl.upsert(df)
963+
964+
assert result.rows_updated == 1
965+
assert result.rows_inserted == 1
966+
assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda row: row["city"]) == [
967+
{"city": "Amsterdam", "population": 950000, "country": "NL"},
968+
{"city": "Berlin", "population": 3432000, "country": "DE"},
969+
]
970+
971+
972+
def test_upsert_after_renaming_column(catalog: Catalog) -> None:
973+
"""Renaming a non-key column must not make identical rows look changed."""
974+
identifier = "default.test_upsert_after_renaming_column"
975+
_drop_table(catalog, identifier)
976+
977+
schema = Schema(
978+
NestedField(1, "city", StringType(), required=True),
979+
NestedField(2, "population", IntegerType(), required=True),
980+
identifier_field_ids=[1],
981+
)
982+
tbl = catalog.create_table(identifier, schema=schema)
983+
984+
tbl.append(
985+
pa.Table.from_pylist(
986+
[{"city": "Amsterdam", "population": 921402}],
987+
schema=schema_to_pyarrow(tbl.schema()),
988+
)
989+
)
990+
991+
# A rename creates no data snapshot either, so the branch tip keeps the old field name
992+
with tbl.update_schema() as update:
993+
update.rename_column("population", "inhabitants")
994+
995+
# The very same row, only under the new column name, so there is nothing to update
996+
df = pa.Table.from_pylist(
997+
[{"city": "Amsterdam", "inhabitants": 921402}],
998+
schema=schema_to_pyarrow(tbl.schema()),
999+
)
1000+
result = tbl.upsert(df)
1001+
1002+
assert result.rows_updated == 0
1003+
assert result.rows_inserted == 0
1004+
1005+
1006+
def test_upsert_after_renaming_join_column(catalog: Catalog) -> None:
1007+
"""The join columns are looked up on the matched rows, so those must carry the current names."""
1008+
identifier = "default.test_upsert_after_renaming_join_column"
1009+
_drop_table(catalog, identifier)
1010+
1011+
schema = Schema(
1012+
NestedField(1, "city", StringType(), required=True),
1013+
NestedField(2, "population", IntegerType(), required=True),
1014+
identifier_field_ids=[1],
1015+
)
1016+
tbl = catalog.create_table(identifier, schema=schema)
1017+
1018+
tbl.append(
1019+
pa.Table.from_pylist(
1020+
[{"city": "Amsterdam", "population": 921402}],
1021+
schema=schema_to_pyarrow(tbl.schema()),
1022+
)
1023+
)
1024+
1025+
with tbl.update_schema() as update:
1026+
update.rename_column("city", "city_name")
1027+
1028+
df = pa.Table.from_pylist(
1029+
[
1030+
{"city_name": "Amsterdam", "population": 950000},
1031+
{"city_name": "Berlin", "population": 3432000},
1032+
],
1033+
schema=schema_to_pyarrow(tbl.schema()),
1034+
)
1035+
result = tbl.upsert(df)
1036+
1037+
assert result.rows_updated == 1
1038+
assert result.rows_inserted == 1
1039+
assert sorted(tbl.scan().to_arrow().to_pylist(), key=lambda row: row["city_name"]) == [
1040+
{"city_name": "Amsterdam", "population": 950000},
1041+
{"city_name": "Berlin", "population": 3432000},
1042+
]
1043+
1044+
1045+
def test_upsert_after_adding_column_in_transaction(catalog: Catalog) -> None:
1046+
"""An upsert must see a column added earlier in the same, still uncommitted, transaction."""
1047+
identifier = "default.test_upsert_after_adding_column_in_transaction"
1048+
_drop_table(catalog, identifier)
1049+
1050+
schema = Schema(
1051+
NestedField(1, "city", StringType(), required=True),
1052+
NestedField(2, "population", IntegerType(), required=True),
1053+
identifier_field_ids=[1],
1054+
)
1055+
tbl = catalog.create_table(identifier, schema=schema)
1056+
1057+
tbl.append(
1058+
pa.Table.from_pylist(
1059+
[{"city": "Amsterdam", "population": 921402}],
1060+
schema=schema_to_pyarrow(tbl.schema()),
1061+
)
1062+
)
1063+
1064+
evolved_schema = Schema(
1065+
NestedField(1, "city", StringType(), required=True),
1066+
NestedField(2, "population", IntegerType(), required=True),
1067+
NestedField(3, "country", StringType(), required=False),
1068+
identifier_field_ids=[1],
1069+
)
1070+
df = pa.Table.from_pylist(
1071+
[{"city": "Amsterdam", "population": 950000, "country": "NL"}],
1072+
schema=schema_to_pyarrow(evolved_schema),
1073+
)
1074+
1075+
with tbl.transaction() as txn:
1076+
with txn.update_schema() as update:
1077+
update.add_column("country", StringType())
1078+
result = txn.upsert(df)
1079+
1080+
assert result.rows_updated == 1
1081+
assert result.rows_inserted == 0
1082+
assert tbl.scan().to_arrow().to_pylist() == [{"city": "Amsterdam", "population": 950000, "country": "NL"}]

0 commit comments

Comments
 (0)