Skip to content

Commit a47b47e

Browse files
committed
fix: use starting_snapshot_id for validation window in commit retry
The concurrency validation was using parent_snapshot (current head) for both the starting point and ending point of the validation window. When multiple concurrent commits occur during retry sleep, the validation would only inspect the latest head and miss conflicting commits below it. Introduce _starting_snapshot_id that is fixed at operation init time and does not change on retry. Also fix _validation_history to use exclusive semantics for from_snapshot, matching Java Iceberg's ancestorsBetween. Signed-off-by: Sotaro Hikita <bering1814@gmail.com>
1 parent bcf76cc commit a47b47e

3 files changed

Lines changed: 44 additions & 16 deletions

File tree

‎pyiceberg/table/update/snapshot.py‎

Lines changed: 20 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,7 @@ class _SnapshotProducer(UpdateTableMetadata[U], Generic[U]):
100100
_operation: Operation
101101
_snapshot_id: int
102102
_parent_snapshot_id: int | None
103+
_starting_snapshot_id: int | None
103104
_added_data_files: list[DataFile]
104105
_manifest_num_counter: itertools.count[int]
105106
_deleted_data_files: set[DataFile]
@@ -139,6 +140,7 @@ def __init__(
139140
self._parent_snapshot_id = (
140141
snapshot.snapshot_id if (snapshot := self._transaction.table_metadata.snapshot_by_name(self._target_branch)) else None
141142
)
143+
self._starting_snapshot_id = self._parent_snapshot_id
142144
self._predicate = AlwaysFalse()
143145
self._case_sensitive = True
144146
self._isolation_level_property: str = TableProperties.WRITE_DELETE_ISOLATION_LEVEL
@@ -579,22 +581,27 @@ def _validate_concurrency(self) -> None:
579581
if parent_snapshot is None:
580582
return
581583

584+
starting_snapshot_id = self._starting_snapshot_id if self._starting_snapshot_id is not None else self._parent_snapshot_id
585+
starting_snapshot = table.metadata.snapshot_by_id(starting_snapshot_id)
586+
if starting_snapshot is None:
587+
return
588+
582589
isolation_level_str = table.metadata.properties.get(
583590
self._isolation_level_property, TableProperties.WRITE_ISOLATION_LEVEL_DEFAULT
584591
)
585592
isolation_level = IsolationLevel(isolation_level_str)
586593
conflict_detection_filter = self._predicate if self._predicate != AlwaysFalse() else None
587594

588595
if isolation_level == IsolationLevel.SERIALIZABLE:
589-
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
596+
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, starting_snapshot)
590597

591598
if conflict_detection_filter is not None:
592-
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, parent_snapshot)
593-
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
599+
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, starting_snapshot)
600+
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, starting_snapshot)
594601

595602
if self._deleted_data_files:
596603
_validate_no_new_deletes_for_data_files(
597-
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, parent_snapshot
604+
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, starting_snapshot
598605
)
599606

600607

@@ -792,22 +799,27 @@ def _validate_concurrency(self) -> None:
792799
if parent_snapshot is None:
793800
return
794801

802+
starting_snapshot_id = self._starting_snapshot_id if self._starting_snapshot_id is not None else self._parent_snapshot_id
803+
starting_snapshot = table.metadata.snapshot_by_id(starting_snapshot_id)
804+
if starting_snapshot is None:
805+
return
806+
795807
isolation_level_str = table.metadata.properties.get(
796808
self._isolation_level_property, TableProperties.WRITE_ISOLATION_LEVEL_DEFAULT
797809
)
798810
isolation_level = IsolationLevel(isolation_level_str)
799811
conflict_detection_filter = self._predicate if self._predicate != AlwaysFalse() else None
800812

801813
if isolation_level == IsolationLevel.SERIALIZABLE:
802-
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
814+
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, starting_snapshot)
803815

804816
if conflict_detection_filter is not None:
805-
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, parent_snapshot)
806-
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
817+
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, starting_snapshot)
818+
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, starting_snapshot)
807819

808820
if self._deleted_data_files:
809821
_validate_no_new_deletes_for_data_files(
810-
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, parent_snapshot
822+
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, starting_snapshot
811823
)
812824

813825

‎pyiceberg/table/update/validate.py‎

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -47,10 +47,13 @@ def _validation_history(
4747
) -> tuple[list[ManifestFile], set[int]]:
4848
"""Return newly added manifests and snapshot IDs between the starting snapshot and parent snapshot.
4949
50+
Walks from to_snapshot backwards towards from_snapshot, collecting manifests from
51+
snapshots whose operations match. from_snapshot is excluded from results.
52+
5053
Args:
5154
table: Table to get the history from
52-
from_snapshot: Parent snapshot to get the history from
53-
to_snapshot: Starting snapshot
55+
from_snapshot: Snapshot where the walk stops (exclusive)
56+
to_snapshot: Snapshot where the walk starts
5457
matching_operations: Operations to match on
5558
manifest_content_filter: Manifest content type to filter
5659
@@ -60,11 +63,17 @@ def _validation_history(
6063
Returns:
6164
List of manifest files and set of snapshots ID's matching conditions
6265
"""
66+
if from_snapshot.snapshot_id == to_snapshot.snapshot_id:
67+
return [], set()
68+
6369
manifests_files: list[ManifestFile] = []
6470
snapshots: set[int] = set()
6571

6672
last_snapshot = None
6773
for snapshot in ancestors_between(from_snapshot, to_snapshot, table.metadata):
74+
if snapshot.snapshot_id == from_snapshot.snapshot_id:
75+
last_snapshot = snapshot
76+
break
6877
last_snapshot = snapshot
6978
summary = snapshot.summary
7079
if summary is None:
@@ -73,7 +82,6 @@ def _validation_history(
7382
continue
7483

7584
snapshots.add(snapshot.snapshot_id)
76-
# TODO: Maybe do the IO in a separate thread at some point, and collect at the bottom (we can easily merge the sets
7785
manifests_files.extend(
7886
[
7987
manifest
@@ -82,7 +90,7 @@ def _validation_history(
8290
]
8391
)
8492

85-
if last_snapshot is not None and last_snapshot.snapshot_id != from_snapshot.snapshot_id:
93+
if last_snapshot is None or last_snapshot.snapshot_id != from_snapshot.snapshot_id:
8694
raise ValidationException("No matching snapshot found.")
8795

8896
return manifests_files, snapshots

‎tests/table/test_validate.py‎

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -65,11 +65,18 @@ def test_validation_history(table_v2_with_extensive_snapshots_and_manifests: tup
6565
"""Test the validation history function."""
6666
table, mock_manifests = table_v2_with_extensive_snapshots_and_manifests
6767

68-
expected_manifest_data_counts = len([m for m in mock_manifests.values() if m[0].content == ManifestContent.DATA])
69-
7068
oldest_snapshot = table.snapshots()[0]
7169
newest_snapshot = cast(Snapshot, table.current_snapshot())
7270

71+
# from_snapshot (oldest) is excluded from the results since it's the base
72+
expected_manifest_data_counts = len(
73+
[
74+
m
75+
for snap_id, m in mock_manifests.items()
76+
if m[0].content == ManifestContent.DATA and snap_id != oldest_snapshot.snapshot_id
77+
]
78+
)
79+
7380
def mock_read_manifest_side_effect(self: Snapshot, io: FileIO) -> list[ManifestFile]:
7481
"""Mock the manifests method to use the snapshot_id for lookup."""
7582
snapshot_id = self.snapshot_id
@@ -246,7 +253,8 @@ def test_validate_added_data_files_conflicting_count(
246253
update={"snapshots": snapshots},
247254
)
248255

249-
oldest_snapshot = table.snapshots()[-snapshot_history]
256+
# Use one snapshot before the altered range as the boundary (exclusive stop point)
257+
boundary_snapshot = table.snapshots()[-(snapshot_history + 1)]
250258
newest_snapshot = cast(Snapshot, table.current_snapshot())
251259

252260
def mock_read_manifest_side_effect(self: Snapshot, io: FileIO) -> list[ManifestFile]:
@@ -273,7 +281,7 @@ def mock_fetch_manifest_entry(self: ManifestFile, io: FileIO, discard_deleted: b
273281
table=table,
274282
starting_snapshot=newest_snapshot,
275283
data_filter=None,
276-
parent_snapshot=oldest_snapshot,
284+
parent_snapshot=boundary_snapshot,
277285
partition_set=None,
278286
)
279287
)

0 commit comments

Comments
 (0)