diff --git a/pyiceberg/encryption/kms.py b/pyiceberg/encryption/kms.py new file mode 100644 index 0000000000..a53a3a45b0 --- /dev/null +++ b/pyiceberg/encryption/kms.py @@ -0,0 +1,114 @@ +# 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. +"""Key management client interface for table encryption, and an in-memory implementation.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field + +from pyiceberg.encryption.ciphers import AesGcmCipher, AesKeySize, SecureKey + + +@dataclass(frozen=True) +class GeneratedKey: + """A newly generated key, both in the clear and wrapped by the key management service.""" + + key: bytes = field(repr=False) + wrapped_key: bytes + + +class KeyManagementClient(ABC): + """A base class for key management service implementations. + + Wraps and unwraps table encryption keys using master keys that the service holds. + """ + + @abstractmethod + def wrap_key(self, key: bytes, wrapping_key_id: str) -> bytes: + """Wrap a key using the master key identified by `wrapping_key_id`. + + Args: + key (bytes): The key to wrap. + wrapping_key_id (str): Identifies the master key held by the service. + """ + + @abstractmethod + def unwrap_key(self, wrapped_key: bytes, wrapping_key_id: str) -> bytes: + """Unwrap a key using the master key identified by `wrapping_key_id`. + + Args: + wrapped_key (bytes): The wrapped key, as returned by `wrap_key`. + wrapping_key_id (str): Identifies the master key held by the service. + """ + + def supports_key_generation(self) -> bool: + """Whether the service generates keys itself, rather than only wrapping them.""" + return False + + def generate_key(self, wrapping_key_id: str) -> GeneratedKey: + """Generate a new key, wrapped by the master key identified by `wrapping_key_id`. + + Args: + wrapping_key_id (str): Identifies the master key held by the service. + """ + raise NotImplementedError(f"{type(self).__name__} does not support key generation") + + +class MemoryKeyManagementClient(KeyManagementClient): + """A key management service that holds its master keys in memory, for testing and demonstration. + + Master keys live only in this process, with no durability or access control, so this is + not for production use. Mirrors Java's `MemoryMockKMS` and iceberg-rust's + `MemoryKeyManagementClient`. + """ + + def __init__(self, master_key_size: AesKeySize = AesKeySize.BITS_128) -> None: + self._master_key_size = master_key_size + self._master_keys: dict[str, SecureKey] = {} + + def __repr__(self) -> str: + """Return a representation that counts the master keys without exposing them.""" + return f"MemoryKeyManagementClient(master_key_size={self._master_key_size!r}, key_count={len(self._master_keys)})" + + def add_master_key(self, wrapping_key_id: str, key: SecureKey | None = None) -> SecureKey: + """Register a master key under `wrapping_key_id`, generating one when `key` is omitted. + + Args: + wrapping_key_id (str): The id to register the master key under. + key (SecureKey | None): Known key material, for tests that share it with another client. + """ + if wrapping_key_id in self._master_keys: + raise ValueError(f"Master key already exists: {wrapping_key_id}") + + master_key = SecureKey.generate(self._master_key_size) if key is None else key + self._master_keys[wrapping_key_id] = master_key + return master_key + + def _cipher(self, wrapping_key_id: str) -> AesGcmCipher: + if (master_key := self._master_keys.get(wrapping_key_id)) is None: + raise ValueError(f"Master key not found: {wrapping_key_id}") + + return AesGcmCipher(master_key) + + def wrap_key(self, key: bytes, wrapping_key_id: str) -> bytes: + """Wrap a key with the registered master key, without AAD, as Java and iceberg-rust do.""" + return self._cipher(wrapping_key_id).encrypt(key) + + def unwrap_key(self, wrapped_key: bytes, wrapping_key_id: str) -> bytes: + """Unwrap a key wrapped by `wrap_key`.""" + return self._cipher(wrapping_key_id).decrypt(wrapped_key) diff --git a/tests/encryption/test_kms.py b/tests/encryption/test_kms.py new file mode 100644 index 0000000000..779a9d5e91 --- /dev/null +++ b/tests/encryption/test_kms.py @@ -0,0 +1,141 @@ +# 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. + +import pytest + +from pyiceberg.encryption.ciphers import AesGcmCipher, AesKeySize, SecureKey +from pyiceberg.encryption.kms import GeneratedKey, KeyManagementClient, MemoryKeyManagementClient + +MASTER_KEY_ID = "master-key" +MASTER_KEY = SecureKey(b"0123456789012345") +DEK = b"6543210987654321" + + +@pytest.fixture +def kms() -> MemoryKeyManagementClient: + client = MemoryKeyManagementClient() + client.add_master_key(MASTER_KEY_ID) + return client + + +def test_key_management_client_cannot_be_instantiated() -> None: + with pytest.raises(TypeError, match="abstract"): + KeyManagementClient() # type: ignore[abstract] + + +def test_generated_key_repr_redacts_key() -> None: + generated = GeneratedKey(key=DEK, wrapped_key=b"wrapped") + + assert repr(generated) == "GeneratedKey(wrapped_key=b'wrapped')" + assert repr(DEK) not in repr(generated) + + +def test_key_generation_is_unsupported_by_default(kms: MemoryKeyManagementClient) -> None: + assert kms.supports_key_generation() is False + + with pytest.raises(NotImplementedError, match="MemoryKeyManagementClient does not support key generation"): + kms.generate_key(MASTER_KEY_ID) + + +def test_wrap_unwrap_round_trip(kms: MemoryKeyManagementClient) -> None: + wrapped = kms.wrap_key(DEK, MASTER_KEY_ID) + + assert wrapped != DEK + assert kms.unwrap_key(wrapped, MASTER_KEY_ID) == DEK + + +@pytest.mark.parametrize("key_size", list(AesKeySize)) +def test_wrap_unwrap_round_trip_for_each_master_key_size(key_size: AesKeySize) -> None: + kms = MemoryKeyManagementClient(key_size) + master_key = kms.add_master_key(MASTER_KEY_ID) + + assert master_key.key_size == key_size + assert kms.unwrap_key(kms.wrap_key(DEK, MASTER_KEY_ID), MASTER_KEY_ID) == DEK + + +def test_wrap_key_does_not_reuse_nonce(kms: MemoryKeyManagementClient) -> None: + first, second = kms.wrap_key(DEK, MASTER_KEY_ID), kms.wrap_key(DEK, MASTER_KEY_ID) + + assert first != second + assert kms.unwrap_key(first, MASTER_KEY_ID) == kms.unwrap_key(second, MASTER_KEY_ID) == DEK + + +def test_wrap_key_is_not_bound_to_the_wrapping_key_id(kms: MemoryKeyManagementClient) -> None: + """No AAD is used when wrapping, matching Java's `MemoryMockKMS` and iceberg-rust.""" + kms.add_master_key("other-key", MASTER_KEY) + kms.add_master_key("same-key-different-id", MASTER_KEY) + + wrapped = kms.wrap_key(DEK, "other-key") + + assert kms.unwrap_key(wrapped, "same-key-different-id") == DEK + + +def test_generated_master_keys_are_unique() -> None: + kms = MemoryKeyManagementClient() + + assert kms.add_master_key("first") != kms.add_master_key("second") + + +def test_add_master_key_with_known_key_material() -> None: + kms = MemoryKeyManagementClient() + + assert kms.add_master_key(MASTER_KEY_ID, MASTER_KEY) == MASTER_KEY + assert AesGcmCipher(MASTER_KEY).decrypt(kms.wrap_key(DEK, MASTER_KEY_ID)) == DEK + + +def test_add_master_key_rejects_a_duplicate_id(kms: MemoryKeyManagementClient) -> None: + with pytest.raises(ValueError, match=f"Master key already exists: {MASTER_KEY_ID}"): + kms.add_master_key(MASTER_KEY_ID) + + +@pytest.mark.parametrize("key_length", [0, 15, 33]) +def test_add_master_key_rejects_an_invalid_key_length(key_length: int) -> None: + with pytest.raises(ValueError, match="Unsupported key length"): + MemoryKeyManagementClient().add_master_key(MASTER_KEY_ID, SecureKey(bytes(key_length))) + + +def test_wrap_key_with_an_unknown_master_key_id(kms: MemoryKeyManagementClient) -> None: + with pytest.raises(ValueError, match="Master key not found: missing-key"): + kms.wrap_key(DEK, "missing-key") + + +def test_unwrap_key_with_an_unknown_master_key_id(kms: MemoryKeyManagementClient) -> None: + with pytest.raises(ValueError, match="Master key not found: missing-key"): + kms.unwrap_key(kms.wrap_key(DEK, MASTER_KEY_ID), "missing-key") + + +def test_unwrap_key_with_the_wrong_master_key(kms: MemoryKeyManagementClient) -> None: + wrapped = kms.wrap_key(DEK, MASTER_KEY_ID) + kms.add_master_key("other-key") + + with pytest.raises(ValueError, match="wrong decryption key; or corrupt/tampered data"): + kms.unwrap_key(wrapped, "other-key") + + +def test_unwrap_tampered_key(kms: MemoryKeyManagementClient) -> None: + wrapped = bytearray(kms.wrap_key(DEK, MASTER_KEY_ID)) + wrapped[-1] ^= 0xFF + + with pytest.raises(ValueError, match="wrong decryption key; or corrupt/tampered data"): + kms.unwrap_key(bytes(wrapped), MASTER_KEY_ID) + + +def test_repr_redacts_master_keys(kms: MemoryKeyManagementClient) -> None: + kms.add_master_key("other-key", MASTER_KEY) + + assert repr(kms) == "MemoryKeyManagementClient(master_key_size=, key_count=2)" + assert repr(MASTER_KEY.key) not in repr(kms)