diff --git a/pyiceberg/table/__init__.py b/pyiceberg/table/__init__.py index 9624eac981..7e02c4f534 100644 --- a/pyiceberg/table/__init__.py +++ b/pyiceberg/table/__init__.py @@ -90,7 +90,7 @@ from pyiceberg.table.update.sorting import UpdateSortOrder from pyiceberg.table.update.spec import UpdateSpec from pyiceberg.table.update.statistics import UpdateStatistics -from pyiceberg.transforms import IdentityTransform +from pyiceberg.transforms import IdentityTransform, VoidTransform from pyiceberg.typedef import ( EMPTY_DICT, IcebergBaseModel, @@ -236,6 +236,11 @@ class TableProperties: WRITE_ISOLATION_LEVEL_DEFAULT = "serializable" +def _active_partition_source_ids(spec: PartitionSpec) -> set[int]: + """Return the set of active source field IDs in a partition spec.""" + return {field.source_id for field in spec.fields if not isinstance(field.transform, VoidTransform)} + + class Transaction: _table: Table _autocommit: bool @@ -390,12 +395,18 @@ def _set_ref_snapshot( return updates, requirements - def _build_partition_predicate(self, partition_records: set[Record], partition_fields: list[str]) -> BooleanExpression: + def _build_partition_predicate( + self, + partition_records: set[Record], + partition_fields: list[str], + evolved_fields: set[str] | None = None, + ) -> BooleanExpression: """Build a filter predicate matching any of the input partition records. Args: partition_records: A set of partition records to match partition_fields: The field names to reference for each position in a partition record + evolved_fields: Optional set of field names added during partition spec evolution Returns: A predicate matching any of the input partition records. @@ -403,18 +414,42 @@ def _build_partition_predicate(self, partition_records: set[Record], partition_f if not partition_records or not partition_fields: return AlwaysFalse() + evolved = evolved_fields or set() per_record_exprs: list[BooleanExpression] = [] for partition_record in partition_records: - predicates: list[BooleanExpression] = [ - EqualTo(Reference(partition_field), partition_record[pos]) - if partition_record[pos] is not None - else IsNull(Reference(partition_field)) - for pos, partition_field in enumerate(partition_fields) - ] + predicates: list[BooleanExpression] = [] + for pos, field in enumerate(partition_fields): + ref = Reference(field) + val = partition_record[pos] + if val is None: + predicates.append(IsNull(ref)) + elif field in evolved: + predicates.append(Or(EqualTo(ref, val), IsNull(ref))) + else: + predicates.append(EqualTo(ref, val)) + per_record_exprs.append(And(*predicates) if len(predicates) > 1 else predicates[0]) return Or(*per_record_exprs) if len(per_record_exprs) > 1 else per_record_exprs[0] + def _get_evolved_partition_fields(self, current_spec: PartitionSpec) -> set[str]: + """Find partition fields in the current spec that were absent in at least one historical partitioned spec.""" + historical_specs = [ + spec + for spec in self.table_metadata.specs().values() + if spec.spec_id != current_spec.spec_id and not spec.is_unpartitioned() + ] + if not historical_specs: + return set() + + common_historical_source_ids = set.intersection(*(_active_partition_source_ids(spec) for spec in historical_specs)) + evolved_source_ids = _active_partition_source_ids(current_spec) - common_historical_source_ids + if not evolved_source_ids: + return set() + + schema = self.table_metadata.schema() + return {field.name for source_id in evolved_source_ids if (field := schema.find_field(source_id)) is not None} + def _append_snapshot_producer( self, snapshot_properties: dict[str, str], branch: str | None = MAIN_BRANCH ) -> _FastAppendFiles: @@ -619,11 +654,13 @@ def dynamic_partition_overwrite( ) partitions_to_overwrite = {data_file.partition for data_file in data_files} - partitions_fields = [ - self.table_metadata.schema().find_field(field.source_id).name for field in self.table_metadata.spec().fields - ] + current_spec = self.table_metadata.spec() + partitions_fields = [self.table_metadata.schema().find_field(field.source_id).name for field in current_spec.fields] + evolved_fields = self._get_evolved_partition_fields(current_spec) delete_filter = self._build_partition_predicate( - partition_records=partitions_to_overwrite, partition_fields=partitions_fields + partition_records=partitions_to_overwrite, + partition_fields=partitions_fields, + evolved_fields=evolved_fields, ) self.delete( delete_filter=delete_filter, diff --git a/tests/table/test_init.py b/tests/table/test_init.py index 739039debb..10a92e9948 100644 --- a/tests/table/test_init.py +++ b/tests/table/test_init.py @@ -32,6 +32,9 @@ And, EqualTo, In, + IsNull, + Or, + Reference, ) from pyiceberg.expressions.visitors import bind from pyiceberg.io import PY_IO_IMPL, FileIO, load_file_io @@ -2036,3 +2039,68 @@ def _spy(*args: Any, **kwargs: Any) -> FileIO: assert seen_locations, "expected at least one load_file_io call" assert all(loc is not None for loc in seen_locations), f"load_file_io called without a location: {seen_locations}" + + +def test_build_partition_predicate_with_evolved_fields(table_v2: Table) -> None: + tx = table_v2.transaction() + records = {Record("A", "us")} + fields = ["category", "region"] + + # Without evolved fields + pred = tx._build_partition_predicate(records, fields) + assert pred == And(EqualTo(Reference("category"), "A"), EqualTo(Reference("region"), "us")) + + # With evolved fields + pred_evolved = tx._build_partition_predicate(records, fields, evolved_fields={"region"}) + assert pred_evolved == And( + EqualTo(Reference("category"), "A"), + Or(EqualTo(Reference("region"), "us"), IsNull(Reference("region"))), + ) + + +def test_dynamic_partition_overwrite_with_partition_spec_evolution(catalog: Catalog) -> None: + import pyarrow as pa + + catalog.create_namespace("default") + schema = Schema( + NestedField(1, "category", StringType(), required=False), + NestedField(2, "region", StringType(), required=False), + NestedField(3, "value", LongType(), required=False), + ) + spec_v0 = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="category")) + table = catalog.create_table("default.test_evolve", schema=schema, partition_spec=spec_v0) + + # Write under spec 0 (category only) + table.append( + pa.table( + { + "category": ["A", "A", "B"], + "region": pa.array([None, None, None], type=pa.string()), + "value": [1, 2, 10], + } + ) + ) + + # Evolve spec to add region + with table.update_spec() as u: + u.add_field("region", IdentityTransform(), "region") + table = catalog.load_table("default.test_evolve") + + # Write under spec 1 (category + region) + table.append(pa.table({"category": ["A", "B"], "region": ["us", "us"], "value": [100, 200]})) + table.append(pa.table({"category": ["A"], "region": ["eu"], "value": [555]})) + + # Overwrite category=A, region=us — should delete A under spec-0 and (A, us) under spec-1, + # while preserving (A, eu) and all B rows + table.dynamic_partition_overwrite(pa.table({"category": ["A"], "region": ["us"], "value": [999]})) + + result = table.scan().to_arrow().to_pydict() + rows = list(zip(result["category"], result["region"], result["value"], strict=True)) + + # Verify category A rows + a_rows = [r for r in rows if r[0] == "A"] + assert sorted(a_rows, key=lambda x: (x[0], x[1] or "", x[2])) == [("A", "eu", 555), ("A", "us", 999)] + + # Verify category B rows are untouched + b_rows = [r for r in rows if r[0] == "B"] + assert sorted(b_rows, key=lambda x: (x[0], x[1] or "", x[2])) == [("B", None, 10), ("B", "us", 200)]