Skip to content

Commit 1fd76c7

Browse files
committed
Fix dynamic_partition_overwrite with partition spec evolution (#3148)
* Identify evolved partition fields added in historical partitioned specs * Extend _build_partition_predicate to match IS NULL for evolved fields * Add unit and regression tests for dynamic partition overwrite with spec evolution
1 parent 68898e5 commit 1fd76c7

2 files changed

Lines changed: 121 additions & 12 deletions

File tree

‎pyiceberg/table/__init__.py‎

Lines changed: 49 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -90,7 +90,7 @@
9090
from pyiceberg.table.update.sorting import UpdateSortOrder
9191
from pyiceberg.table.update.spec import UpdateSpec
9292
from pyiceberg.table.update.statistics import UpdateStatistics
93-
from pyiceberg.transforms import IdentityTransform
93+
from pyiceberg.transforms import IdentityTransform, VoidTransform
9494
from pyiceberg.typedef import (
9595
EMPTY_DICT,
9696
IcebergBaseModel,
@@ -236,6 +236,11 @@ class TableProperties:
236236
WRITE_ISOLATION_LEVEL_DEFAULT = "serializable"
237237

238238

239+
def _active_partition_source_ids(spec: PartitionSpec) -> set[int]:
240+
"""Return the set of active source field IDs in a partition spec."""
241+
return {field.source_id for field in spec.fields if not isinstance(field.transform, VoidTransform)}
242+
243+
239244
class Transaction:
240245
_table: Table
241246
_autocommit: bool
@@ -390,31 +395,61 @@ def _set_ref_snapshot(
390395

391396
return updates, requirements
392397

393-
def _build_partition_predicate(self, partition_records: set[Record], partition_fields: list[str]) -> BooleanExpression:
398+
def _build_partition_predicate(
399+
self,
400+
partition_records: set[Record],
401+
partition_fields: list[str],
402+
evolved_fields: set[str] | None = None,
403+
) -> BooleanExpression:
394404
"""Build a filter predicate matching any of the input partition records.
395405
396406
Args:
397407
partition_records: A set of partition records to match
398408
partition_fields: The field names to reference for each position in a partition record
409+
evolved_fields: Optional set of field names added during partition spec evolution
399410
400411
Returns:
401412
A predicate matching any of the input partition records.
402413
"""
403414
if not partition_records or not partition_fields:
404415
return AlwaysFalse()
405416

417+
evolved = evolved_fields or set()
406418
per_record_exprs: list[BooleanExpression] = []
407419
for partition_record in partition_records:
408-
predicates: list[BooleanExpression] = [
409-
EqualTo(Reference(partition_field), partition_record[pos])
410-
if partition_record[pos] is not None
411-
else IsNull(Reference(partition_field))
412-
for pos, partition_field in enumerate(partition_fields)
413-
]
420+
predicates: list[BooleanExpression] = []
421+
for pos, field in enumerate(partition_fields):
422+
ref = Reference(field)
423+
val = partition_record[pos]
424+
if val is None:
425+
predicates.append(IsNull(ref))
426+
elif field in evolved:
427+
predicates.append(Or(EqualTo(ref, val), IsNull(ref)))
428+
else:
429+
predicates.append(EqualTo(ref, val))
430+
414431
per_record_exprs.append(And(*predicates) if len(predicates) > 1 else predicates[0])
415432

416433
return Or(*per_record_exprs) if len(per_record_exprs) > 1 else per_record_exprs[0]
417434

435+
def _get_evolved_partition_fields(self, current_spec: PartitionSpec) -> set[str]:
436+
"""Find partition fields in the current spec that were absent in any historical partitioned spec."""
437+
historical_specs = [
438+
spec
439+
for spec in self.table_metadata.specs().values()
440+
if spec.spec_id != current_spec.spec_id and not spec.is_unpartitioned()
441+
]
442+
if not historical_specs:
443+
return set()
444+
445+
common_historical_source_ids = set.intersection(*(_active_partition_source_ids(spec) for spec in historical_specs))
446+
evolved_source_ids = _active_partition_source_ids(current_spec) - common_historical_source_ids
447+
if not evolved_source_ids:
448+
return set()
449+
450+
schema = self.table_metadata.schema()
451+
return {field.name for source_id in evolved_source_ids if (field := schema.find_field(source_id)) is not None}
452+
418453
def _append_snapshot_producer(
419454
self, snapshot_properties: dict[str, str], branch: str | None = MAIN_BRANCH
420455
) -> _FastAppendFiles:
@@ -619,11 +654,13 @@ def dynamic_partition_overwrite(
619654
)
620655

621656
partitions_to_overwrite = {data_file.partition for data_file in data_files}
622-
partitions_fields = [
623-
self.table_metadata.schema().find_field(field.source_id).name for field in self.table_metadata.spec().fields
624-
]
657+
current_spec = self.table_metadata.spec()
658+
partitions_fields = [self.table_metadata.schema().find_field(field.source_id).name for field in current_spec.fields]
659+
evolved_fields = self._get_evolved_partition_fields(current_spec)
625660
delete_filter = self._build_partition_predicate(
626-
partition_records=partitions_to_overwrite, partition_fields=partitions_fields
661+
partition_records=partitions_to_overwrite,
662+
partition_fields=partitions_fields,
663+
evolved_fields=evolved_fields,
627664
)
628665
self.delete(
629666
delete_filter=delete_filter,

‎tests/table/test_init.py‎

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
import json
1919
import uuid
2020
from copy import copy
21+
from pathlib import Path
2122
from typing import Any
2223

2324
import pytest
@@ -32,6 +33,9 @@
3233
And,
3334
EqualTo,
3435
In,
36+
IsNull,
37+
Or,
38+
Reference,
3539
)
3640
from pyiceberg.expressions.visitors import bind
3741
from pyiceberg.io import PY_IO_IMPL, FileIO, load_file_io
@@ -2036,3 +2040,71 @@ def _spy(*args: Any, **kwargs: Any) -> FileIO:
20362040

20372041
assert seen_locations, "expected at least one load_file_io call"
20382042
assert all(loc is not None for loc in seen_locations), f"load_file_io called without a location: {seen_locations}"
2043+
2044+
2045+
def test_build_partition_predicate_with_evolved_fields(table_v2: Table) -> None:
2046+
tx = table_v2.transaction()
2047+
records = {Record("A", "us")}
2048+
fields = ["category", "region"]
2049+
2050+
# Without evolved fields
2051+
pred = tx._build_partition_predicate(records, fields)
2052+
assert pred == And(EqualTo(Reference("category"), "A"), EqualTo(Reference("region"), "us"))
2053+
2054+
# With evolved fields
2055+
pred_evolved = tx._build_partition_predicate(records, fields, evolved_fields={"region"})
2056+
assert pred_evolved == And(
2057+
EqualTo(Reference("category"), "A"),
2058+
Or(EqualTo(Reference("region"), "us"), IsNull(Reference("region"))),
2059+
)
2060+
2061+
2062+
def test_dynamic_partition_overwrite_with_partition_spec_evolution(warehouse: Path) -> None:
2063+
import pyarrow as pa
2064+
2065+
from pyiceberg.catalog.sql import SqlCatalog
2066+
2067+
catalog = SqlCatalog(name="test", uri=f"sqlite:///{warehouse.as_posix()}/test_dpo_evolve.db", warehouse=f"file://{warehouse}")
2068+
catalog.create_namespace("default")
2069+
schema = Schema(
2070+
NestedField(1, "category", StringType(), required=False),
2071+
NestedField(2, "region", StringType(), required=False),
2072+
NestedField(3, "value", LongType(), required=False),
2073+
)
2074+
spec_v0 = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="category"))
2075+
table = catalog.create_table("default.test_evolve", schema=schema, partition_spec=spec_v0)
2076+
2077+
# Write under spec 0 (category only)
2078+
table.append(
2079+
pa.table(
2080+
{
2081+
"category": ["A", "A", "B"],
2082+
"region": pa.array([None, None, None], type=pa.string()),
2083+
"value": [1, 2, 10],
2084+
}
2085+
)
2086+
)
2087+
2088+
# Evolve spec to add region
2089+
with table.update_spec() as u:
2090+
u.add_field("region", IdentityTransform(), "region")
2091+
table = catalog.load_table("default.test_evolve")
2092+
2093+
# Write under spec 1 (category + region)
2094+
table.append(pa.table({"category": ["A", "B"], "region": ["us", "us"], "value": [100, 200]}))
2095+
table.append(pa.table({"category": ["A"], "region": ["eu"], "value": [555]}))
2096+
2097+
# Overwrite category=A, region=us — should delete A under spec-0 and (A, us) under spec-1,
2098+
# while preserving (A, eu) and all B rows
2099+
table.dynamic_partition_overwrite(pa.table({"category": ["A"], "region": ["us"], "value": [999]}))
2100+
2101+
result = table.scan().to_arrow().to_pydict()
2102+
rows = list(zip(result["category"], result["region"], result["value"], strict=True))
2103+
2104+
# Verify category A rows
2105+
a_rows = [r for r in rows if r[0] == "A"]
2106+
assert sorted(a_rows, key=lambda x: (x[0], x[1] or "", x[2])) == [("A", "eu", 555), ("A", "us", 999)]
2107+
2108+
# Verify category B rows are untouched
2109+
b_rows = [r for r in rows if r[0] == "B"]
2110+
assert sorted(b_rows, key=lambda x: (x[0], x[1] or "", x[2])) == [("B", None, 10), ("B", "us", 200)]

0 commit comments

Comments
 (0)