Skip to content

Commit e015a56

Browse files
Gayathri Srividya RajavarapuGayathri Srividya Rajavarapu
authored andcommitted
fix(table): handle dynamic overwrite across partition spec evolution
1 parent 2c75523 commit e015a56

3 files changed

Lines changed: 210 additions & 8 deletions

File tree

pyiceberg/table/__init__.py

Lines changed: 120 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -593,10 +593,127 @@ def dynamic_partition_overwrite(
593593
)
594594

595595
partitions_to_overwrite = {data_file.partition for data_file in data_files}
596-
delete_filter = self._build_partition_predicate(
597-
partition_records=partitions_to_overwrite, spec=self.table_metadata.spec(), schema=self.table_metadata.schema()
596+
current_spec = self.table_metadata.spec()
597+
all_specs = self.table_metadata.specs()
598+
schema = self.table_metadata.schema()
599+
600+
# Keep the existing dynamic overwrite behavior for non-evolution tables.
601+
# We only need per-spec predicate handling when some historical specs are
602+
# missing current partition source IDs.
603+
current_source_ids = {field.source_id for field in current_spec.fields}
604+
has_missing_partition_fields_in_history = any(
605+
current_source_ids - {field.source_id for field in historical_spec.fields} for historical_spec in all_specs.values()
598606
)
599-
self.delete(delete_filter=delete_filter, snapshot_properties=snapshot_properties, branch=branch)
607+
608+
if not has_missing_partition_fields_in_history:
609+
delete_filter = self._build_partition_predicate(
610+
partition_records=partitions_to_overwrite,
611+
spec=current_spec,
612+
schema=schema,
613+
)
614+
self.delete(delete_filter=delete_filter, snapshot_properties=snapshot_properties, branch=branch)
615+
else:
616+
# Build per-spec delete predicates to handle partition spec evolution correctly.
617+
#
618+
# When a partition field was added via spec evolution, data files written under
619+
# older specs carry NULL for that field (because it was absent from the schema at
620+
# write time). A single "category=A AND region=us" predicate would never match
621+
# those files because the strict-metrics evaluator sees region=NULL != "us".
622+
#
623+
# To fix this, we compute a per-spec predicate for every historical spec:
624+
# - For specs that include all current partition fields -> use exact-match predicate.
625+
# - For specs that are missing some current partition fields -> also accept NULL
626+
# for the missing fields.
627+
#
628+
# These per-spec predicates are stored on the delete snapshot producer so that
629+
# _compute_deletes uses the right predicate when evaluating each manifest file.
630+
source_id_to_pos = {field.source_id: pos for pos, field in enumerate(current_spec.fields)}
631+
source_id_to_col = {field.source_id: schema.find_field(field.source_id).name for field in current_spec.fields}
632+
exact_delete_filter = self._build_partition_predicate(
633+
partition_records=partitions_to_overwrite,
634+
spec=current_spec,
635+
schema=schema,
636+
)
637+
638+
per_spec_predicates: dict[int, BooleanExpression] = {}
639+
for spec_id, hist_spec in all_specs.items():
640+
hist_source_ids = {field.source_id for field in hist_spec.fields}
641+
missing_source_ids = current_source_ids - hist_source_ids
642+
has_overlap_with_current = bool(hist_source_ids & current_source_ids)
643+
644+
per_record_exprs: list[BooleanExpression] = []
645+
for partition_record in partitions_to_overwrite:
646+
predicates: list[BooleanExpression] = []
647+
for source_id, col_name in source_id_to_col.items():
648+
value = partition_record[source_id_to_pos[source_id]]
649+
if value is not None:
650+
field_pred: BooleanExpression = EqualTo(Reference(col_name), value)
651+
if source_id in missing_source_ids and has_overlap_with_current:
652+
field_pred = Or(field_pred, IsNull(Reference(col_name)))
653+
else:
654+
field_pred = IsNull(Reference(col_name))
655+
predicates.append(field_pred)
656+
657+
per_record_exprs.append(And(*predicates) if len(predicates) > 1 else predicates[0])
658+
659+
per_spec_predicates[spec_id] = Or(*per_record_exprs) if len(per_record_exprs) > 1 else per_record_exprs[0]
660+
661+
# Open the delete snapshot and set per-spec predicates before committing.
662+
# This mirrors Transaction.delete() but injects per_spec_predicates so that
663+
# _compute_deletes uses the right predicate for each historical spec.
664+
from pyiceberg.io.pyarrow import ArrowScan, _dataframe_to_data_files, _expression_to_complementary_pyarrow
665+
666+
with self.update_snapshot(snapshot_properties=snapshot_properties, branch=branch).delete() as delete_snapshot:
667+
delete_snapshot._per_spec_predicates = per_spec_predicates
668+
delete_snapshot.delete_by_predicate(exact_delete_filter)
669+
670+
# Handle partial-match files that need to be rewritten (copy-on-write).
671+
if delete_snapshot.rewrites_needed is True:
672+
bound_delete_filter = bind(self.table_metadata.schema(), exact_delete_filter, case_sensitive=True)
673+
preserve_row_filter = _expression_to_complementary_pyarrow(bound_delete_filter, self.table_metadata.schema())
674+
675+
file_scan = self._scan(row_filter=exact_delete_filter)
676+
if branch is not None:
677+
file_scan = file_scan.use_ref(branch)
678+
679+
rewrite_uuid = uuid.uuid4()
680+
rewrite_counter = itertools.count(0)
681+
replaced_files: list[tuple[DataFile, list[DataFile]]] = []
682+
for original_file in file_scan.plan_files():
683+
df_orig = ArrowScan(
684+
table_metadata=self.table_metadata,
685+
io=self._table.io,
686+
projected_schema=self.table_metadata.schema(),
687+
row_filter=AlwaysTrue(),
688+
).to_table(tasks=[original_file])
689+
filtered_df = df_orig.filter(preserve_row_filter)
690+
if len(filtered_df) == 0:
691+
replaced_files.append((original_file.file, []))
692+
elif len(df_orig) != len(filtered_df):
693+
replaced_files.append(
694+
(
695+
original_file.file,
696+
list(
697+
_dataframe_to_data_files(
698+
io=self._table.io,
699+
df=filtered_df,
700+
table_metadata=self.table_metadata,
701+
write_uuid=rewrite_uuid,
702+
counter=rewrite_counter,
703+
)
704+
),
705+
)
706+
)
707+
708+
if replaced_files:
709+
with self.update_snapshot(
710+
snapshot_properties=snapshot_properties, branch=branch
711+
).overwrite() as overwrite_snapshot:
712+
overwrite_snapshot.commit_uuid = rewrite_uuid
713+
for original_data_file, replacement_data_files in replaced_files:
714+
overwrite_snapshot.delete_data_file(original_data_file)
715+
for replacement_data_file in replacement_data_files:
716+
overwrite_snapshot.append_data_file(replacement_data_file)
600717

601718
with self._append_snapshot_producer(snapshot_properties, branch=branch) as append_files:
602719
append_files.commit_uuid = append_snapshot_commit_uuid

pyiceberg/table/update/snapshot.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -104,6 +104,7 @@ class _SnapshotProducer(UpdateTableMetadata[U], Generic[U]):
104104
_target_branch: str | None
105105
_predicate: BooleanExpression
106106
_case_sensitive: bool
107+
_per_spec_predicates: dict[int, BooleanExpression]
107108

108109
def __init__(
109110
self,
@@ -134,6 +135,7 @@ def __init__(
134135
)
135136
self._predicate = AlwaysFalse()
136137
self._case_sensitive = True
138+
self._per_spec_predicates = {}
137139

138140
def _validate_target_branch(self, branch: str | None) -> str | None:
139141
# if branch is none, write will be written into a staging snapshot
@@ -360,7 +362,8 @@ def fetch_manifest_entry(self, manifest: ManifestFile, discard_deleted: bool = T
360362

361363
def _build_partition_projection(self, spec_id: int) -> BooleanExpression:
362364
project = inclusive_projection(self.schema(), self.spec(spec_id), self._case_sensitive)
363-
return project(self._predicate)
365+
predicate = self._per_spec_predicates.get(spec_id, self._predicate)
366+
return project(predicate)
364367

365368
@cached_property
366369
def partition_filters(self) -> KeyDefaultDict[int, BooleanExpression]:
@@ -431,10 +434,14 @@ def _copy_with_new_status(entry: ManifestEntry, status: ManifestEntryStatus) ->
431434
schema = table_metadata.schema()
432435

433436
manifest_evaluators: dict[int, Callable[[ManifestFile], bool]] = KeyDefaultDict(self._build_manifest_evaluator)
434-
strict_metrics_evaluator = _StrictMetricsEvaluator(schema, self._predicate, case_sensitive=self._case_sensitive).eval
435-
inclusive_metrics_evaluator = _InclusiveMetricsEvaluator(
436-
schema, self._predicate, case_sensitive=self._case_sensitive
437-
).eval
437+
438+
def _strict_metrics_for_spec(spec_id: int) -> Callable[[DataFile], bool]:
439+
predicate = self._per_spec_predicates.get(spec_id, self._predicate)
440+
return _StrictMetricsEvaluator(schema, predicate, case_sensitive=self._case_sensitive).eval
441+
442+
def _inclusive_metrics_for_spec(spec_id: int) -> Callable[[DataFile], bool]:
443+
predicate = self._per_spec_predicates.get(spec_id, self._predicate)
444+
return _InclusiveMetricsEvaluator(schema, predicate, case_sensitive=self._case_sensitive).eval
438445

439446
existing_manifests = []
440447
total_deleted_entries = []
@@ -454,6 +461,9 @@ def _copy_with_new_status(entry: ManifestEntry, status: ManifestEntryStatus) ->
454461
existing_manifests.append(manifest_file)
455462
else:
456463
# It is relevant, let's check out the content
464+
spec_id = manifest_file.partition_spec_id
465+
strict_metrics_evaluator = _strict_metrics_for_spec(spec_id)
466+
inclusive_metrics_evaluator = _inclusive_metrics_for_spec(spec_id)
457467
deleted_entries = []
458468
existing_entries = []
459469
for entry in manifest_file.fetch_manifest_entry(io=self._io, discard_deleted=True):

tests/table/test_init.py

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1989,3 +1989,78 @@ def test_build_large_partition_predicate(table_v2: Table) -> None:
19891989
)
19901990

19911991
bind(table_v2.metadata.schema(), expr, case_sensitive=True)
1992+
1993+
1994+
def test_dynamic_partition_overwrite_spec_evolution(tmp_path: Any) -> None:
1995+
"""Regression test for https://github.com/apache/iceberg-python/issues/3148.
1996+
1997+
After partition spec evolution, dynamic_partition_overwrite must delete data files
1998+
written under the old spec (where the new partition field was absent / NULL) when
1999+
overwriting the matching logical partition.
2000+
"""
2001+
import tempfile
2002+
2003+
import pyarrow as pa
2004+
2005+
from pyiceberg.catalog import load_catalog
2006+
from pyiceberg.transforms import IdentityTransform
2007+
from pyiceberg.types import LongType
2008+
2009+
with tempfile.TemporaryDirectory() as warehouse:
2010+
catalog = load_catalog("test", type="sql", uri=f"sqlite:///{warehouse}/catalog.db", warehouse=f"file://{warehouse}")
2011+
catalog.create_namespace("default")
2012+
2013+
schema = Schema(
2014+
NestedField(1, "category", StringType(), required=False),
2015+
NestedField(2, "region", StringType(), required=False),
2016+
NestedField(3, "value", LongType(), required=False),
2017+
)
2018+
spec_v0 = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="category"))
2019+
table = catalog.create_table("default.test_spec_evo", schema=schema, partition_spec=spec_v0)
2020+
2021+
# Write under spec-0 (region is NULL — field exists in schema but not in partition spec)
2022+
table.append(
2023+
pa.table(
2024+
{
2025+
"category": pa.array(["A", "A", "B"], type=pa.string()),
2026+
"region": pa.array([None, None, None], type=pa.string()),
2027+
"value": pa.array([1, 2, 10], type=pa.int64()),
2028+
}
2029+
)
2030+
)
2031+
2032+
# Evolve to spec-1: add region as a partition field
2033+
with table.update_spec() as u:
2034+
u.add_field("region", IdentityTransform(), "region")
2035+
table = catalog.load_table("default.test_spec_evo")
2036+
2037+
# Write under spec-1
2038+
table.append(
2039+
pa.table(
2040+
{
2041+
"category": pa.array(["A", "B"], type=pa.string()),
2042+
"region": pa.array(["us", "us"], type=pa.string()),
2043+
"value": pa.array([100, 200], type=pa.int64()),
2044+
}
2045+
)
2046+
)
2047+
2048+
# Overwrite partition {A, us} — must also delete stale spec-0 {A} files
2049+
table.dynamic_partition_overwrite(
2050+
pa.table(
2051+
{
2052+
"category": pa.array(["A"], type=pa.string()),
2053+
"region": pa.array(["us"], type=pa.string()),
2054+
"value": pa.array([999], type=pa.int64()),
2055+
}
2056+
)
2057+
)
2058+
2059+
result = table.scan().to_arrow().to_pydict()
2060+
a_values = sorted([v for c, v in zip(result["category"], result["value"], strict=True) if c == "A"])
2061+
b_values = sorted([v for c, v in zip(result["category"], result["value"], strict=True) if c == "B"])
2062+
2063+
# Spec-0 rows 1,2 (category=A, region=NULL) should be gone; only 999 remains
2064+
assert a_values == [999], f"Expected [999] but got {a_values}"
2065+
# B rows from both specs should be untouched
2066+
assert b_values == [10, 200], f"Expected [10, 200] but got {b_values}"

0 commit comments

Comments
 (0)