diff --git a/pyiceberg/table/__init__.py b/pyiceberg/table/__init__.py index fca718f5ec..303b3db135 100644 --- a/pyiceberg/table/__init__.py +++ b/pyiceberg/table/__init__.py @@ -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 diff --git a/pyiceberg/table/metadata.py b/pyiceberg/table/metadata.py index 84be7b07bf..8236f12229 100644 --- a/pyiceberg/table/metadata.py +++ b/pyiceberg/table/metadata.py @@ -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 @@ -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" @@ -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") + """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. @@ -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) + """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") diff --git a/pyiceberg/table/snapshots.py b/pyiceberg/table/snapshots.py index 5e9e519a01..f862529e83 100644 --- a/pyiceberg/table/snapshots.py +++ b/pyiceberg/table/snapshots.py @@ -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.""" @@ -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)})" diff --git a/tests/table/test_metadata.py b/tests/table/test_metadata.py index c163c90626..d696a16205 100644 --- a/tests/table/test_metadata.py +++ b/tests/table/test_metadata.py @@ -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, @@ -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": {}} + ] diff --git a/tests/table/test_snapshots.py b/tests/table/test_snapshots.py index 0f72b08087..439375cc69 100644 --- a/tests/table/test_snapshots.py +++ b/tests/table/test_snapshots.py @@ -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(