Skip to content

Commit 43337e6

Browse files
committed
Add commit retry with data conflict validation
Add automatic retry with exponential backoff when catalog commits fail due to concurrent transactions (CommitFailedException), and integrate the existing validation functions from validate.py into the write path to detect incompatible concurrent modifications (ValidationException). The retry loop is placed in Transaction.commit_transaction(). On each retry attempt, table metadata is refreshed, registered snapshot producers are re-executed to regenerate manifests, and data conflict validation is run. Uncommitted manifests from failed attempts are cleaned up after a successful commit. Validation is performed for _OverwriteFiles and _DeleteFiles based on the table's isolation level (serializable/snapshot). _FastAppendFiles and _MergeAppendFiles do not require validation since appends never conflict. Signed-off-by: Sotaro Hikita <bering1814@gmail.com>
1 parent 714c24f commit 43337e6

5 files changed

Lines changed: 698 additions & 29 deletions

File tree

‎pyiceberg/table/__init__.py‎

Lines changed: 83 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
from pydantic import Field
3232

3333
import pyiceberg.expressions.parser as parser
34+
from pyiceberg.exceptions import CommitFailedException
3435
from pyiceberg.expressions import AlwaysFalse, AlwaysTrue, And, BooleanExpression, EqualTo, IsNull, Or, Reference
3536
from pyiceberg.expressions.visitors import (
3637
ResidualEvaluator,
@@ -205,6 +206,23 @@ class TableProperties:
205206
MIN_SNAPSHOTS_TO_KEEP = "history.expire.min-snapshots-to-keep"
206207
MIN_SNAPSHOTS_TO_KEEP_DEFAULT = 1
207208

209+
COMMIT_NUM_RETRIES = "commit.retry.num-retries"
210+
COMMIT_NUM_RETRIES_DEFAULT = 4
211+
212+
COMMIT_MIN_RETRY_WAIT_MS = "commit.retry.min-wait-ms"
213+
COMMIT_MIN_RETRY_WAIT_MS_DEFAULT = 100
214+
215+
COMMIT_MAX_RETRY_WAIT_MS = "commit.retry.max-wait-ms"
216+
COMMIT_MAX_RETRY_WAIT_MS_DEFAULT = 60000
217+
218+
COMMIT_TOTAL_RETRY_TIME_MS = "commit.retry.total-timeout-ms"
219+
COMMIT_TOTAL_RETRY_TIME_MS_DEFAULT = 1800000 # 30 minutes
220+
221+
WRITE_DELETE_ISOLATION_LEVEL = "write.delete.isolation-level"
222+
WRITE_UPDATE_ISOLATION_LEVEL = "write.update.isolation-level"
223+
WRITE_MERGE_ISOLATION_LEVEL = "write.merge.isolation-level"
224+
WRITE_ISOLATION_LEVEL_DEFAULT = "serializable"
225+
208226

209227
class Transaction:
210228
_table: Table
@@ -223,6 +241,7 @@ def __init__(self, table: Table, autocommit: bool = False):
223241
self._autocommit = autocommit
224242
self._updates = ()
225243
self._requirements = ()
244+
self._snapshot_producers: list[Any] = []
226245

227246
@property
228247
def table_metadata(self) -> TableMetadata:
@@ -265,6 +284,10 @@ def _stage(
265284

266285
return self
267286

287+
def _register_snapshot_producer(self, producer: Any) -> None:
288+
"""Register a snapshot producer for retry support."""
289+
self._snapshot_producers.append(producer)
290+
268291
def _apply(
269292
self,
270293
updates: tuple[TableUpdate, ...],
@@ -703,6 +726,7 @@ def delete(
703726
snapshot_properties=snapshot_properties, branch=branch
704727
).overwrite() as overwrite_snapshot:
705728
overwrite_snapshot.commit_uuid = commit_uuid
729+
overwrite_snapshot.delete_by_predicate(delete_filter, case_sensitive)
706730
for original_data_file, replaced_data_files in replaced_files:
707731
overwrite_snapshot.delete_data_file(original_data_file)
708732
for replaced_data_file in replaced_data_files:
@@ -939,17 +963,67 @@ def commit_transaction(self) -> Table:
939963
The table with the updates applied.
940964
"""
941965
if len(self._updates) > 0:
942-
self._requirements += (AssertTableUUID(uuid=self.table_metadata.table_uuid),)
943-
self._table._do_commit( # pylint: disable=W0212
944-
updates=self._updates,
945-
requirements=self._requirements,
966+
from pyiceberg.utils.properties import property_as_int
967+
968+
properties = self._table.metadata.properties
969+
num_retries_val = property_as_int(
970+
properties, TableProperties.COMMIT_NUM_RETRIES, TableProperties.COMMIT_NUM_RETRIES_DEFAULT
971+
)
972+
num_retries = num_retries_val if num_retries_val is not None else TableProperties.COMMIT_NUM_RETRIES_DEFAULT
973+
min_wait_val = property_as_int(
974+
properties, TableProperties.COMMIT_MIN_RETRY_WAIT_MS, TableProperties.COMMIT_MIN_RETRY_WAIT_MS_DEFAULT
946975
)
976+
min_wait_ms = min_wait_val if min_wait_val is not None else TableProperties.COMMIT_MIN_RETRY_WAIT_MS_DEFAULT
977+
max_wait_val = property_as_int(
978+
properties, TableProperties.COMMIT_MAX_RETRY_WAIT_MS, TableProperties.COMMIT_MAX_RETRY_WAIT_MS_DEFAULT
979+
)
980+
max_wait_ms = max_wait_val if max_wait_val is not None else TableProperties.COMMIT_MAX_RETRY_WAIT_MS_DEFAULT
981+
982+
for attempt in range(num_retries + 1):
983+
try:
984+
self._requirements += (AssertTableUUID(uuid=self.table_metadata.table_uuid),)
985+
self._table._do_commit( # pylint: disable=W0212
986+
updates=self._updates,
987+
requirements=self._requirements,
988+
)
989+
self._cleanup_uncommitted_manifests()
990+
break
991+
except CommitFailedException:
992+
if attempt == num_retries or not self._snapshot_producers:
993+
raise
994+
import random
995+
import time
996+
997+
wait = min(min_wait_ms * (2**attempt), max_wait_ms)
998+
jitter = random.uniform(0, 0.25 * wait)
999+
time.sleep((wait + jitter) / 1000.0)
1000+
1001+
self._table.refresh()
1002+
self._rebuild_snapshot_updates()
9471003

9481004
self._updates = ()
9491005
self._requirements = ()
9501006

9511007
return self._table
9521008

1009+
def _cleanup_uncommitted_manifests(self) -> None:
1010+
"""Clean up manifests from failed retry attempts after a successful commit."""
1011+
for producer in self._snapshot_producers:
1012+
producer._cleanup_uncommitted()
1013+
1014+
def _rebuild_snapshot_updates(self) -> None:
1015+
"""Rebuild snapshot updates for retry by re-executing registered producers."""
1016+
from pyiceberg.table.update import AddSnapshotUpdate, AssertRefSnapshotId, SetSnapshotRefUpdate
1017+
1018+
self._updates = tuple(u for u in self._updates if not isinstance(u, (AddSnapshotUpdate, SetSnapshotRefUpdate)))
1019+
self._requirements = tuple(r for r in self._requirements if not isinstance(r, (AssertRefSnapshotId, AssertTableUUID)))
1020+
1021+
for producer in self._snapshot_producers:
1022+
producer._refresh_for_retry()
1023+
producer._validate_concurrency()
1024+
updates, requirements = producer._commit()
1025+
self._stage(updates, requirements)
1026+
9531027

9541028
class CreateTableTransaction(Transaction):
9551029
"""A transaction that involves the creation of a new table."""
@@ -1961,13 +2035,11 @@ def _build_residual_evaluator(self, spec_id: int) -> Callable[[DataFile], Residu
19612035
# The lambda created here is run in multiple threads.
19622036
# So we avoid creating _EvaluatorExpression methods bound to a single
19632037
# shared instance across multiple threads.
1964-
return lambda datafile: (
1965-
residual_evaluator_of(
1966-
spec=spec,
1967-
expr=self.row_filter,
1968-
case_sensitive=self.case_sensitive,
1969-
schema=self.table_metadata.schema(),
1970-
)
2038+
return lambda datafile: residual_evaluator_of(
2039+
spec=spec,
2040+
expr=self.row_filter,
2041+
case_sensitive=self.case_sensitive,
2042+
schema=self.table_metadata.schema(),
19712043
)
19722044

19732045
@staticmethod

‎pyiceberg/table/snapshots.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -86,6 +86,13 @@ def __repr__(self) -> str:
8686
return f"Operation.{self.name}"
8787

8888

89+
class IsolationLevel(str, Enum):
90+
"""Transaction isolation level for concurrent write validation."""
91+
92+
SERIALIZABLE = "serializable"
93+
SNAPSHOT = "snapshot"
94+
95+
8996
class UpdateMetrics:
9097
added_file_size: int
9198
removed_file_size: int

‎pyiceberg/table/update/snapshot.py‎

Lines changed: 113 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from __future__ import annotations
1818

1919
import itertools
20+
import logging
2021
import uuid
2122
from abc import abstractmethod
2223
from collections import defaultdict
@@ -80,6 +81,8 @@
8081
if TYPE_CHECKING:
8182
from pyiceberg.table import Transaction
8283

84+
logger = logging.getLogger(__name__)
85+
8386

8487
def _new_manifest_file_name(num: int, commit_uuid: uuid.UUID) -> str:
8588
return f"{commit_uuid}-m{num}.avro"
@@ -104,6 +107,8 @@ class _SnapshotProducer(UpdateTableMetadata[U], Generic[U]):
104107
_target_branch: str | None
105108
_predicate: BooleanExpression
106109
_case_sensitive: bool
110+
_written_manifests: list[str]
111+
_uncommitted_manifests: list[str]
107112

108113
def __init__(
109114
self,
@@ -123,6 +128,8 @@ def __init__(
123128
self._deleted_data_files = set()
124129
self.snapshot_properties = snapshot_properties
125130
self._manifest_num_counter = itertools.count(0)
131+
self._written_manifests = []
132+
self._uncommitted_manifests = []
126133
from pyiceberg.table import TableProperties
127134

128135
self._compression = self._transaction.table_metadata.properties.get( # type: ignore
@@ -351,11 +358,39 @@ def new_manifest_output(self) -> OutputFile:
351358
location_provider = self._transaction._table.location_provider()
352359
file_name = _new_manifest_file_name(num=next(self._manifest_num_counter), commit_uuid=self.commit_uuid)
353360
file_path = location_provider.new_metadata_location(file_name)
361+
self._written_manifests.append(file_path)
354362
return self._io.new_output(file_path)
355363

356364
def fetch_manifest_entry(self, manifest: ManifestFile, discard_deleted: bool = True) -> list[ManifestEntry]:
357365
return manifest.fetch_manifest_entry(io=self._io, discard_deleted=discard_deleted)
358366

367+
def commit(self) -> None:
368+
self._transaction._register_snapshot_producer(self)
369+
self._transaction._apply(*self._commit())
370+
371+
def _cleanup_uncommitted(self) -> None:
372+
"""Delete manifest files from failed retry attempts."""
373+
for path in self._uncommitted_manifests:
374+
try:
375+
self._io.delete(path)
376+
except Exception:
377+
logger.warning("Failed to delete uncommitted manifest: %s", path, exc_info=True)
378+
self._uncommitted_manifests.clear()
379+
380+
def _refresh_for_retry(self) -> None:
381+
"""Reset state for a retry attempt with refreshed metadata."""
382+
self._uncommitted_manifests.extend(self._written_manifests)
383+
self._written_manifests.clear()
384+
self._parent_snapshot_id = (
385+
snapshot.snapshot_id if (snapshot := self._transaction.table_metadata.snapshot_by_name(self._target_branch)) else None
386+
)
387+
self._snapshot_id = self._transaction.table_metadata.new_snapshot_id()
388+
self._manifest_num_counter = itertools.count(0)
389+
self.commit_uuid = uuid.uuid4()
390+
391+
def _validate_concurrency(self) -> None:
392+
"""Validate that concurrent changes do not conflict with this operation. No-op by default."""
393+
359394
def _build_partition_projection(self, spec_id: int) -> BooleanExpression:
360395
project = inclusive_projection(self.schema(), self.spec(spec_id), self._case_sensitive)
361396
return project(self._predicate)
@@ -495,6 +530,48 @@ def files_affected(self) -> bool:
495530
"""Indicate if any manifest-entries can be dropped."""
496531
return len(self._deleted_entries()) > 0
497532

533+
def _refresh_for_retry(self) -> None:
534+
"""Reset state for a retry attempt, clearing the cached delete computation."""
535+
super()._refresh_for_retry()
536+
if "_compute_deletes" in self.__dict__:
537+
del self.__dict__["_compute_deletes"]
538+
539+
def _validate_concurrency(self) -> None:
540+
"""Validate that concurrent changes do not conflict with this delete."""
541+
from pyiceberg.table import TableProperties
542+
from pyiceberg.table.snapshots import IsolationLevel
543+
from pyiceberg.table.update.validate import (
544+
_validate_added_data_files,
545+
_validate_deleted_data_files,
546+
_validate_no_new_delete_files,
547+
_validate_no_new_deletes_for_data_files,
548+
)
549+
550+
if self._parent_snapshot_id is None:
551+
return
552+
553+
table = self._transaction._table
554+
parent_snapshot = table.metadata.snapshot_by_id(self._parent_snapshot_id)
555+
if parent_snapshot is None:
556+
return
557+
558+
isolation_level_str = table.metadata.properties.get(
559+
TableProperties.WRITE_DELETE_ISOLATION_LEVEL, TableProperties.WRITE_ISOLATION_LEVEL_DEFAULT
560+
)
561+
isolation_level = IsolationLevel(isolation_level_str)
562+
conflict_detection_filter = self._predicate if self._predicate != AlwaysFalse() else None
563+
564+
if isolation_level == IsolationLevel.SERIALIZABLE:
565+
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
566+
567+
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, parent_snapshot)
568+
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
569+
570+
if self._deleted_data_files:
571+
_validate_no_new_deletes_for_data_files(
572+
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, parent_snapshot
573+
)
574+
498575

499576
class _FastAppendFiles(_SnapshotProducer["_FastAppendFiles"]):
500577
def _existing_manifests(self) -> list[ManifestFile]:
@@ -666,6 +743,42 @@ def _get_entries(manifest: ManifestFile) -> list[ManifestEntry]:
666743
else:
667744
return []
668745

746+
def _validate_concurrency(self) -> None:
747+
"""Validate that concurrent changes do not conflict with this overwrite."""
748+
from pyiceberg.table import TableProperties
749+
from pyiceberg.table.snapshots import IsolationLevel
750+
from pyiceberg.table.update.validate import (
751+
_validate_added_data_files,
752+
_validate_deleted_data_files,
753+
_validate_no_new_delete_files,
754+
_validate_no_new_deletes_for_data_files,
755+
)
756+
757+
if self._parent_snapshot_id is None:
758+
return
759+
760+
table = self._transaction._table
761+
parent_snapshot = table.metadata.snapshot_by_id(self._parent_snapshot_id)
762+
if parent_snapshot is None:
763+
return
764+
765+
isolation_level_str = table.metadata.properties.get(
766+
TableProperties.WRITE_DELETE_ISOLATION_LEVEL, TableProperties.WRITE_ISOLATION_LEVEL_DEFAULT
767+
)
768+
isolation_level = IsolationLevel(isolation_level_str)
769+
conflict_detection_filter = self._predicate if self._predicate != AlwaysFalse() else None
770+
771+
if isolation_level == IsolationLevel.SERIALIZABLE:
772+
_validate_added_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
773+
774+
_validate_no_new_delete_files(table, parent_snapshot, conflict_detection_filter, None, parent_snapshot)
775+
_validate_deleted_data_files(table, parent_snapshot, conflict_detection_filter, parent_snapshot)
776+
777+
if self._deleted_data_files:
778+
_validate_no_new_deletes_for_data_files(
779+
table, parent_snapshot, conflict_detection_filter, self._deleted_data_files, parent_snapshot
780+
)
781+
669782

670783
class UpdateSnapshot:
671784
_transaction: Transaction

0 commit comments

Comments
 (0)