Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
146 changes: 146 additions & 0 deletions pyiceberg/table/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -113,6 +113,7 @@
SnapshotLogEntry,
SnapshotSummaryCollector,
Summary,
ancestors_of,
update_snapshot_summaries,
)
from pyiceberg.table.sorting import UNSORTED_SORT_ORDER, SortOrder
Expand DownExpand Up@@ -140,6 +141,7 @@
from pyiceberg.utils.datetime import datetime_to_millis
from pyiceberg.utils.deprecated import deprecated
from pyiceberg.utils.singleton import _convert_to_hashable_type
from pyiceberg.utils.wap import get_wap_id, validate_wap_publish

if TYPE_CHECKING:
import daft
Expand DownExpand Up@@ -1956,6 +1958,10 @@ def _commit(self) -> UpdatesAndRequirements:
"""Apply the pending changes and commit."""
return self._updates, self._requirements

def _commit_if_ref_updates_exist(self) -> None:
self.commit()
self._updates, self._requirements = (), ()

def create_tag(self, snapshot_id: int, tag_name: str, max_ref_age_ms: Optional[int] = None) -> ManageSnapshots:
"""
Create a new tag pointing to the given snapshot id.
Expand DownExpand Up@@ -2010,6 +2016,53 @@ def create_branch(
self._requirements += requirement
return self

def cherrypick_snapshot(self, snapshot_id: int) -> ManageSnapshots:
"""Create a new snapshot from an existing snapshot without altering or removing the original.

Only append & overwrite snapshots can be cherry-picked.

Args:
snapshot_id: snapshot id of the snapshot to cherry-pick
"""
self._commit_if_ref_updates_exist()
if snapshot := self._transaction._table.metadata.snapshot_by_id(snapshot_id):
if snapshot.summary:
CherryPickSnapshot(
snapshot_id=snapshot_id,
operation=snapshot.summary.operation,
transaction=self._transaction,
io=self._transaction._table.io,
).cherrypick().commit()
return self

def publish_changes(self, staged_wap_id: str) -> ManageSnapshots:
"""Apply changes in a snapshot created within a Write-Audit-Publish workflow with a wap_id. Then create a new snapshot which will be set as the current snapshot in a table.

Only append and dynamic overwrite snapshots can be successfully published.

Args:
staged_wap_id: staged wap id of the snapshot to cherry-pick
"""
self._commit_if_ref_updates_exist()
if wap_snapshot := next(
(
snapshot
for snapshot in self._transaction._table.metadata.snapshots
if staged_wap_id == get_wap_id(snapshot, STAGED_WAP_ID_PROP)
),
None,
):
if wap_snapshot.summary:
CherryPickSnapshot(
snapshot_id=wap_snapshot.snapshot_id,
operation=wap_snapshot.summary.operation,
transaction=self._transaction,
io=self._transaction._table.io,
).cherrypick().commit()
else:
raise ValidationError(f"Cannot apply unknown WAP ID {staged_wap_id}")
return self


class UpdateSchema(UpdateTableMetadata["UpdateSchema"]):
_schema: Schema
Expand DownExpand Up@@ -2948,6 +3001,7 @@ class _MergingSnapshotProducer(UpdateTableMetadata["_MergingSnapshotProducer"]):
_snapshot_id: int
_parent_snapshot_id: Optional[int]
_added_data_files: List[DataFile]
snapshot_properties: Dict[str, str]

def __init__(
self,
Expand DownExpand Up@@ -3104,6 +3158,98 @@ def _commit(self) -> UpdatesAndRequirements:
)


STAGED_WAP_ID_PROP = "wap.id"
PUBLISHED_WAP_ID_PROP = "published-wap-id"
SOURCE_SNAPSHOT_ID_PROP = "source-snapshot-id"
REPLACE_PARTITIONS_PROP = "replace-partitions"


class CherryPickSnapshot(_MergingSnapshotProducer):
_cherry_pick_snapshot_id: int
_cherry_pick_snapshot: Optional[Snapshot]
_table_metadata: TableMetadata
_replaced_partitions: set
_require_fast_forward: bool
_properties: Dict[str, str]

def __init__(self, snapshot_id: int, operation: Operation, transaction: Transaction, io: FileIO):
super().__init__(operation=operation, transaction=transaction, io=io)
self._cherry_pick_snapshot_id = snapshot_id
self._cherry_pick_snapshot = transaction.table_metadata.snapshot_by_id(snapshot_id)
self._table_metadata = transaction.table_metadata
self._replaced_partitions = set()
self._properties = dict()

def is_fast_forward(self):
if self._table_metadata.current_snapshot():
# can fast-forward if the cherry-picked snapshot's parent is the current snapshot
return (
self._cherry_pick_snapshot.parent_snapshot_id is not None
and self._table_metadata.current_snapshot().snapshot_id == self._cherry_pick_snapshot.snapshot_id
)
else:
# ... or if the parent and current snapshot are both null
return self._cherry_pick_snapshot is None

def cherrypick(self) -> CherryPickSnapshot:
if not self._cherry_pick_snapshot:
raise ValidationError(f"Cannot cherry-pick unknown snapshot ID: {self._cherry_pick_snapshot_id}")

summary = self._properties
if self._operation == Operation.APPEND:
if (wap_id := validate_wap_publish(self._table_metadata, self._cherry_pick_snapshot)) not in (None, ""):
summary[STAGED_WAP_ID_PROP] = wap_id
summary[SOURCE_SNAPSHOT_ID_PROP] = str(self._cherry_pick_snapshot_id)

for manifest in self._cherry_pick_snapshot.manifests(self._io):
if manifest.content == ManifestContent.DATA:
for entry in manifest.fetch_manifest_entry(self._io):
if entry.status == ManifestEntryStatus.ADDED:
self.append_data_file(entry.data_file)

elif self._operation == Operation.OVERWRITE and (
self._cherry_pick_snapshot.summary
and self._cherry_pick_snapshot.summary.get(REPLACE_PARTITIONS_PROP, "").lower() == "true"
):
if self._cherry_pick_snapshot.parent_snapshot_id is not None and (
self._cherry_pick_snapshot.parent_snapshot_id
not in {
ancestor.snapshot_id
for ancestor in ancestors_of(self._table_metadata.current_snapshot(), self._table_metadata)
}
):
raise ValidationError(
f"Cannot cherry-pick overwrite not based on an ancestor of the current state: {self._cherry_pick_snapshot_id}"
)

if (wap_id := validate_wap_publish(self._table_metadata, self._cherry_pick_snapshot)) not in (None, ""):
summary[STAGED_WAP_ID_PROP] = wap_id
summary[SOURCE_SNAPSHOT_ID_PROP] = str(self._cherry_pick_snapshot_id)

# TODO: failMissingDeletePaths() from Java
for manifest in self._cherry_pick_snapshot.manifests(self._io):
if manifest.content == ManifestContent.DATA:
for entry in manifest.fetch_manifest_entry(self._io):
if entry.status == ManifestEntryStatus.ADDED:
self.append_data_file(entry.data_file)
self._replaced_partitions.add((entry.data_file.spec_id, entry.data_file.partition))
elif entry.status == ManifestEntryStatus.DELETED:
self.delete_data_file(entry.data_file)

elif not self.is_fast_forward():
raise ValidationError(
f"Cannot cherry-pick snapshot {self._cherry_pick_snapshot.snapshot_id}: not append, dynamic overwrite, or fast-forward"
)
self.snapshot_properties = self._properties
return self

def _deleted_entries(self) -> List[ManifestEntry]:
return []

def _existing_manifests(self) -> List[ManifestFile]:
return []


class FastAppendFiles(_MergingSnapshotProducer):
def _existing_manifests(self) -> List[ManifestFile]:
"""To determine if there are any existing manifest files.
Expand Down
40 changes: 40 additions & 0 deletions pyiceberg/utils/wap.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,40 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from __future__ import annotations

from typing import Optional

import pyiceberg.table as tbl


def get_wap_id(snapshot: tbl.Snapshot, state: str) -> Optional[str]:
return snapshot.summary.get(state, None) if snapshot.summary else None


def validate_wap_publish(table_metadata: tbl.TableMetadata, snapshot: tbl.Snapshot) -> str:
wap_id = get_wap_id(snapshot=snapshot, state=tbl.STAGED_WAP_ID_PROP)
if wap_id not in [None, ""]:
if is_wap_id_published(table_metadata, wap_id):
raise ValueError(f"Duplicate request to cherry pick wap id that was published already: {wap_id}")
return wap_id if wap_id else ""


def is_wap_id_published(table_metadata: tbl.TableMetadata, wap_id: str) -> bool:
for ancestor in ancestors_of(table_metadata.current_snapshot(), table_metadata):
if wap_id in (get_wap_id(ancestor, tbl.STAGED_WAP_ID_PROP), get_wap_id(ancestor, tbl.PUBLISHED_WAP_ID_PROP)):
return True
return False
15 changes: 15 additions & 0 deletions tests/integration/test_snapshot_operations.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -17,6 +17,7 @@
import pytest

from pyiceberg.catalog import Catalog
from pyiceberg.table import SOURCE_SNAPSHOT_ID_PROP
from pyiceberg.table.refs import SnapshotRef


Expand All@@ -40,3 +41,17 @@ def test_create_branch(catalog: Catalog) -> None:
branch_snapshot_id = tbl.history()[-2].snapshot_id
tbl.manage_snapshots().create_branch(snapshot_id=branch_snapshot_id, branch_name="branch123").commit()
assert tbl.metadata.refs["branch123"] == SnapshotRef(snapshot_id=branch_snapshot_id, snapshot_ref_type="branch")


@pytest.mark.integration
@pytest.mark.parametrize("catalog", [pytest.lazy_fixture("session_catalog_hive"), pytest.lazy_fixture("session_catalog")])
def test_cherrypick_snapshot(catalog: Catalog):
identifier = "default.test_table_snapshot_operations"
tbl = catalog.load_table(identifier)
assert len(tbl.history()) > 3
current_snapshot_id = tbl.current_snapshot().snapshot_id
cherrypick_snapshot_id = tbl.history()[-2].snapshot_id
tbl.manage_snapshots().cherrypick_snapshot(snapshot_id=cherrypick_snapshot_id).commit()
assert tbl.current_snapshot().snapshot_id is not current_snapshot_id
assert tbl.current_snapshot().snapshot_id is not cherrypick_snapshot_id
assert tbl.current_snapshot().summary[SOURCE_SNAPSHOT_ID_PROP] == str(cherrypick_snapshot_id)