Skip to content
Merged
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
4 changes: 4 additions & 0 deletions pyiceberg/table/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,10 @@ class TableProperties:
FORMAT_VERSION = "format-version"
DEFAULT_FORMAT_VERSION: TableVersion = 2

ENCRYPTION_KEY_ID = "encryption.key-id"
ENCRYPTION_DATA_KEY_LENGTH = "encryption.data-key-length"
ENCRYPTION_DATA_KEY_LENGTH_DEFAULT = 16

MANIFEST_TARGET_SIZE_BYTES = "commit.manifest.target-size-bytes"
MANIFEST_TARGET_SIZE_BYTES_DEFAULT = 8 * 1024 * 1024 # 8 MB

Expand Down
49 changes: 47 additions & 2 deletions pyiceberg/table/metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,14 @@
# under the License.
from __future__ import annotations

import base64
import datetime
import uuid
from collections.abc import Iterable
from copy import copy
from typing import Annotated, Any, Literal
from typing import TYPE_CHECKING, Annotated, Any, Literal

from pydantic import Field, field_serializer, field_validator, model_validator
from pydantic import Field, field_serializer, field_validator, model_serializer, model_validator
from pydantic import ValidationError as PydanticValidationError

from pyiceberg.exceptions import ValidationError
Expand All @@ -48,6 +49,9 @@
from pyiceberg.utils.config import Config
from pyiceberg.utils.datetime import datetime_to_millis

if TYPE_CHECKING:
from pydantic.functional_serializers import ModelWrapSerializerWithoutInfo

CURRENT_SNAPSHOT_ID = "current-snapshot-id"
CURRENT_SCHEMA_ID = "current-schema-id"
SCHEMAS = "schemas"
Expand Down Expand Up @@ -125,6 +129,36 @@ def construct_refs(table_metadata: TableMetadata) -> TableMetadata:
return table_metadata


class EncryptedKey(IcebergBaseModel):
"""A key used for table encryption, tracked in v3 metadata under `encryption-keys`.

https://iceberg.apache.org/spec/#encryption-keys
"""

key_id: str = Field(alias="key-id")
"""ID of the encryption key."""

encrypted_key_metadata: bytes = Field(alias="encrypted-key-metadata")
Comment thread
kevinjqliu marked this conversation as resolved.
"""The encrypted key and metadata, base64 encoded in JSON."""

encrypted_by_id: str | None = Field(alias="encrypted-by-id", default=None)
"""ID of the key used to encrypt or wrap `encrypted-key-metadata`."""

properties: dict[str, str] = Field(default_factory=dict)
"""Additional metadata used by the table's encryption scheme."""

@field_validator("encrypted_key_metadata", mode="before")
def decode_encrypted_key_metadata(cls, encrypted_key_metadata: Any) -> Any:
# validate=True so that malformed base64 raises instead of silently discarding characters
if isinstance(encrypted_key_metadata, str):
return base64.b64decode(encrypted_key_metadata, validate=True)
return encrypted_key_metadata

@field_serializer("encrypted_key_metadata")
def serialize_encrypted_key_metadata(self, encrypted_key_metadata: bytes) -> str:
return base64.b64encode(encrypted_key_metadata).decode("utf-8")


class TableMetadataCommonFields(IcebergBaseModel):
"""Metadata for an Iceberg table as specified in the Apache Iceberg spec.

Expand Down Expand Up @@ -584,6 +618,17 @@ def construct_refs(self) -> TableMetadata:
next_row_id: int | None = Field(alias="next-row-id", default=None)
"""A long higher than all assigned row IDs; the next snapshot's `first-row-id`."""

encryption_keys: list[EncryptedKey] = Field(alias="encryption-keys", default_factory=list)
Comment thread
kevinjqliu marked this conversation as resolved.
"""An optional list of encryption keys used for table encryption."""

@model_serializer(mode="wrap")
def serialize_model(self, handler: ModelWrapSerializerWithoutInfo) -> dict[str, Any]:
"""Set custom serializer to leave out `encryption-keys` when it is empty."""
serialized: dict[str, Any] = handler(self)
if not self.encryption_keys:
serialized.pop("encryption-keys", None)
return serialized

def model_dump_json(self, exclude_none: bool = True, exclude: Any | None = None, by_alias: bool = True, **kwargs: Any) -> str:
raise NotImplementedError("Writing V3 is not yet supported, see: https://github.com/apache/iceberg-python/issues/1551")

Expand Down
4 changes: 4 additions & 0 deletions pyiceberg/table/snapshots.py
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,9 @@ class Snapshot(IcebergBaseModel):
added_rows: int | None = Field(
alias="added-rows", default=None, description="The upper bound of the number of rows with assigned row IDs"
)
key_id: str | None = Field(
alias="key-id", default=None, description="ID of the encryption key that encrypts the manifest list key metadata"
)

def __str__(self) -> str:
"""Return the string representation of the Snapshot class."""
Expand All @@ -280,6 +283,7 @@ def __repr__(self) -> str:
f"schema_id={self.schema_id}" if self.schema_id is not None else None,
f"first_row_id={self.first_row_id}" if self.first_row_id is not None else None,
f"added_rows={self.added_rows}" if self.added_rows is not None else None,
f"key_id='{self.key_id}'" if self.key_id is not None else None,
]
filtered_fields = [field for field in fields if field is not None]
return f"Snapshot({', '.join(filtered_fields)})"
Expand Down
116 changes: 116 additions & 0 deletions tests/table/test_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,12 +24,14 @@
from uuid import UUID

import pytest
from pydantic import ValidationError as PydanticValidationError

from pyiceberg.exceptions import ValidationError
from pyiceberg.partitioning import PartitionField, PartitionSpec
from pyiceberg.schema import Schema
from pyiceberg.serializers import FromByteStream
from pyiceberg.table.metadata import (
EncryptedKey,
TableMetadataUtil,
TableMetadataV1,
TableMetadataV2,
Expand Down Expand Up @@ -876,3 +878,117 @@ def test_new_table_metadata_format_v2_with_v3_schema_fails(field_type: Primitive
location="s3://some_v1_location/",
properties={"format-version": "2"},
)


def test_encrypted_key_minimal() -> None:
# Mirrors Java's TestEncryptedKeyParser, where the key metadata is base64 of b"key"
key = EncryptedKey.model_validate_json('{"key-id": "a", "encrypted-key-metadata": "a2V5"}')

assert key.key_id == "a"
assert key.encrypted_key_metadata == b"key"
assert key.encrypted_by_id is None
assert key.properties == {}


def test_encrypted_key_full() -> None:
key = EncryptedKey.model_validate_json(
'{"key-id": "a", "encrypted-key-metadata": "a2V5", "encrypted-by-id": "b", "properties": {"test": "value"}}'
)

assert key.key_id == "a"
assert key.encrypted_key_metadata == b"key"
assert key.encrypted_by_id == "b"
assert key.properties == {"test": "value"}


def test_encrypted_key_serialize() -> None:
key = EncryptedKey(key_id="a", encrypted_key_metadata=b"key", encrypted_by_id="b", properties={"test": "value"})

expected = '{"key-id":"a","encrypted-key-metadata":"a2V5","encrypted-by-id":"b","properties":{"test":"value"}}'
assert key.model_dump_json() == expected


def test_encrypted_key_serialize_minimal() -> None:
key = EncryptedKey(key_id="a", encrypted_key_metadata=b"key")

assert key.model_dump_json() == '{"key-id":"a","encrypted-key-metadata":"a2V5","properties":{}}'


@pytest.mark.parametrize(
"payload, missing",
[
('{"encrypted-key-metadata": "a2V5"}', "key-id"),
('{"key-id": "a"}', "encrypted-key-metadata"),
],
)
def test_encrypted_key_missing_required_field(payload: str, missing: str) -> None:
with pytest.raises(PydanticValidationError) as exc_info:
EncryptedKey.model_validate_json(payload)

assert missing in str(exc_info.value)


@pytest.mark.parametrize(
"encrypted_key_metadata",
[
"a2V*5", # character outside the base64 alphabet
"a2V5=extra",
"a2V", # incorrect padding
"not base64",
],
)
def test_encrypted_key_malformed_base64(encrypted_key_metadata: str) -> None:
with pytest.raises(PydanticValidationError, match="encrypted-key-metadata"):
EncryptedKey.model_validate({"key-id": "a", "encrypted-key-metadata": encrypted_key_metadata})


def test_v3_metadata_parsing_encryption_keys(example_table_metadata_v3: dict[str, Any]) -> None:
metadata = {
**example_table_metadata_v3,
"encryption-keys": [
{
"key-id": "table-key-1",
"encrypted-key-metadata": "a2V5",
"encrypted-by-id": "external-key-1",
"properties": {"a": "b"},
},
{"key-id": "table-key-2", "encrypted-key-metadata": "a2V5", "encrypted-by-id": "table-key-1"},
],
"snapshots": [
{**snapshot, "key-id": "table-key-2"} if snapshot["snapshot-id"] == 3055729675574597004 else snapshot
for snapshot in example_table_metadata_v3["snapshots"]
],
}

table_metadata = TableMetadataUtil.parse_obj(metadata)

assert isinstance(table_metadata, TableMetadataV3)
assert [key.key_id for key in table_metadata.encryption_keys] == ["table-key-1", "table-key-2"]
assert table_metadata.encryption_keys[0].encrypted_key_metadata == b"key"
assert table_metadata.encryption_keys[0].properties == {"a": "b"}
assert table_metadata.encryption_keys[1].encrypted_by_id == "table-key-1"

current_snapshot = table_metadata.snapshot_by_id(3055729675574597004)
assert current_snapshot is not None
assert current_snapshot.key_id == "table-key-2"


def test_v3_metadata_without_encryption_keys(example_table_metadata_v3: dict[str, Any]) -> None:
table_metadata = TableMetadataUtil.parse_obj(example_table_metadata_v3)

assert isinstance(table_metadata, TableMetadataV3)
assert table_metadata.encryption_keys == []
for snapshot in table_metadata.snapshots:
assert snapshot.key_id is None

assert "encryption-keys" not in table_metadata.model_dump(mode="json")


def test_v3_metadata_with_encryption_keys_serializes_them(example_table_metadata_v3: dict[str, Any]) -> None:
table_metadata = TableMetadataUtil.parse_obj(
{**example_table_metadata_v3, "encryption-keys": [{"key-id": "a", "encrypted-key-metadata": "a2V5"}]}
)

assert table_metadata.model_dump(mode="json")["encryption-keys"] == [
{"key-id": "a", "encrypted-key-metadata": "a2V5", "properties": {}}
]
42 changes: 42 additions & 0 deletions tests/table/test_snapshots.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,48 @@ def test_snapshot_with_properties_repr(snapshot_with_properties: Snapshot) -> No
assert snapshot_with_properties == eval(repr(snapshot_with_properties))


@pytest.fixture
def snapshot_with_key_id() -> Snapshot:
return Snapshot(
snapshot_id=25,
parent_snapshot_id=19,
sequence_number=200,
timestamp_ms=1602638573590,
manifest_list="s3:/a/b/c.avro",
summary=Summary(Operation.APPEND),
schema_id=3,
key_id="table-key-1",
)


def test_serialize_snapshot_with_key_id(snapshot_with_key_id: Snapshot) -> None:
assert snapshot_with_key_id.model_dump_json() == (
'{"snapshot-id":25,"parent-snapshot-id":19,"sequence-number":200,"timestamp-ms":1602638573590,'
'"manifest-list":"s3:/a/b/c.avro","summary":{"operation":"append"},"schema-id":3,"key-id":"table-key-1"}'
)


def test_deserialize_snapshot_with_key_id(snapshot_with_key_id: Snapshot) -> None:
payload = (
'{"snapshot-id": 25, "parent-snapshot-id": 19, "sequence-number": 200, "timestamp-ms": 1602638573590, '
'"manifest-list": "s3:/a/b/c.avro", "summary": {"operation": "append"}, "schema-id": 3, "key-id": "table-key-1"}'
)
assert Snapshot.model_validate_json(payload) == snapshot_with_key_id


def test_snapshot_without_key_id_omits_it(snapshot: Snapshot) -> None:
assert snapshot.key_id is None
assert "key-id" not in snapshot.model_dump_json()


def test_snapshot_with_key_id_repr(snapshot_with_key_id: Snapshot) -> None:
assert repr(snapshot_with_key_id) == (
"Snapshot(snapshot_id=25, parent_snapshot_id=19, sequence_number=200, timestamp_ms=1602638573590, "
"manifest_list='s3:/a/b/c.avro', summary=Summary(Operation.APPEND), schema_id=3, key_id='table-key-1')"
)
assert snapshot_with_key_id == eval(repr(snapshot_with_key_id))


@pytest.fixture
def manifest_file() -> ManifestFile:
return ManifestFile.from_args(
Expand Down
Loading