diff --git a/CHANGELOG.md b/CHANGELOG.md index 20add07..e8ce014 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,8 +7,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- Share AEAD and deterministic AEAD primitives across fields using the same cached keyset, avoiding repeated wrapper construction while retaining bounded caching and weak manager tracking. + ### Fixed +- Synchronize primitive construction with cache invalidation so an in-flight load cannot republish a stale primitive after `clear_keyset_cache()`. + - Restore the database timezone when decrypting naive datetime representations under `USE_TZ=True`, preserving instants across reads, re-saves, and deterministic lookups without rewriting stored ciphertext. - Reject inherited JSON/date transforms and late-registered plaintext lookups on encrypted columns; keep deterministic exact and SQL null lookups explicit. diff --git a/README.md b/README.md index 4702815..a400784 100644 --- a/README.md +++ b/README.md @@ -167,6 +167,8 @@ from tink_fields import clear_keyset_cache clear_keyset_cache() ``` +Cache invalidation is synchronized with keyset loading and primitive construction. Operations that already obtained a primitive may finish with the old key; subsequent field operations load the replacement. The cache is local to each process, so reload or restart every worker. + Changing `keyset=` does not re-encrypt existing rows; it only changes how future reads and writes are processed. Likewise, changing an existing plaintext Django field to an encrypted field requires an explicit staged data migration. Back up data and test recovery before any key or ciphertext migration. ## Security limitations diff --git a/benchmarks/keyset_cache.py b/benchmarks/keyset_cache.py new file mode 100644 index 0000000..3dcdbdc --- /dev/null +++ b/benchmarks/keyset_cache.py @@ -0,0 +1,37 @@ +"""Run with `python -m benchmarks.keyset_cache` from the repository root.""" + +from pathlib import Path +from statistics import median +from timeit import repeat + +from django.conf import settings + +from tink_fields.fields import KeysetManager + + +def main() -> None: + settings.configure( + TINK_FIELDS_CONFIG={ + "default": { + "path": Path(__file__).resolve().parents[1] / "tink_fields/test/test_plaintext_keyset.json", + "cleartext": True, + } + } + ) + + def initialize_fields() -> None: + KeysetManager.clear_cache() + managers = [KeysetManager("default") for _ in range(100)] + for manager in managers: + _ = manager.aead_primitive + + cold = median(repeat(initialize_fields, number=20, repeat=7)) / 20 + manager = KeysetManager("default") + _ = manager.aead_primitive + warm = median(repeat(lambda: manager.aead_primitive.encrypt(b"secret", b""), number=100_000, repeat=7)) / 100_000 + print(f"Initialize 100 fields sharing a keyset: {cold * 1_000:.3f} ms") + print(f"Warm primitive access and encryption: {warm * 1_000_000:.3f} us") + + +if __name__ == "__main__": + main() diff --git a/tink_fields/fields.py b/tink_fields/fields.py index d7c9ee0..66fb3ac 100644 --- a/tink_fields/fields.py +++ b/tink_fields/fields.py @@ -10,12 +10,12 @@ import json from collections import OrderedDict from collections.abc import Callable, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import datetime from os import PathLike from pathlib import Path from threading import RLock -from typing import Any, ClassVar, cast +from typing import Any, ClassVar, TypeVar, cast from weakref import WeakSet from django.conf import settings @@ -119,6 +119,15 @@ def validate(self) -> None: raise ImproperlyConfigured("Encrypted keysets must specify `master_key_aead`.") +Primitive = TypeVar("Primitive", aead.Aead, daead.DeterministicAead) + + +@dataclass +class _KeysetEntry: + handle: Any + primitives: dict[type[Any], Any] = field(default_factory=dict) + + class KeysetManager: """Manages Tink keyset handles and primitives. @@ -128,7 +137,7 @@ class KeysetManager: _cache_size: ClassVar[int] = 32 _cache_lock: ClassVar[RLock] = RLock() - _handle_cache: ClassVar[OrderedDict[tuple[Any, ...], Any]] = OrderedDict() + _handle_cache: ClassVar[OrderedDict[tuple[Any, ...], _KeysetEntry]] = OrderedDict() _managers: ClassVar[WeakSet[KeysetManager]] = WeakSet() def __init__(self, keyset_name: str, aad_callback: AADCallback = _default_aad_callback) -> None: @@ -140,7 +149,7 @@ def __init__(self, keyset_name: str, aad_callback: AADCallback = _default_aad_ca """ self.keyset_name = keyset_name self.aad_callback = aad_callback - self._keyset_handle = None + self._entry: _KeysetEntry | None = None with self._cache_lock: self._managers.add(self) @@ -183,84 +192,80 @@ def clear_cache(cls) -> None: with cls._cache_lock: cls._handle_cache.clear() for manager in list(cls._managers): - manager._keyset_handle = None - manager.__dict__.pop("aead_primitive", None) - manager.__dict__.pop("daead_primitive", None) + manager._entry = None - def _get_tink_keyset_handle(self) -> Any: - """Read the configuration for the requested keyset and return a keyset handle. + def _get_keyset_entry(self) -> _KeysetEntry: + """Load or reuse a keyset entry while holding the cache lock.""" + if self._entry is not None: + return self._entry - Returns: - KeysetHandle: The configured Tink keyset handle - - Raises: - ImproperlyConfigured: If keyset configuration is invalid or missing - """ - if self._keyset_handle is None: - keyset_config = self._get_keyset_config() - keyset_path = Path(keyset_config.path).expanduser().resolve() - try: - stat = keyset_path.stat() - except OSError as error: - raise ImproperlyConfigured(f"Could not load keyset `{self.keyset_name}`.") from error - cache_key = ( - str(keyset_path), - stat.st_mtime_ns, - stat.st_size, - keyset_config.cleartext, - keyset_config.master_key_aead, - ) + keyset_config = self._get_keyset_config() + keyset_path = Path(keyset_config.path).expanduser().resolve() + try: + stat = keyset_path.stat() + except OSError as error: + raise ImproperlyConfigured(f"Could not load keyset `{self.keyset_name}`.") from error + cache_key = ( + str(keyset_path), + stat.st_mtime_ns, + stat.st_size, + keyset_config.cleartext, + keyset_config.master_key_aead, + ) + try: + hash(cache_key) + except TypeError: + cache_key = () + + cached_entry = self._handle_cache.get(cache_key) if cache_key else None + if cached_entry is not None: + self._handle_cache.move_to_end(cache_key) + self._entry = cached_entry + else: try: - hash(cache_key) - except TypeError: - cache_key = () - - with self._cache_lock: - cached_handle = self._handle_cache.get(cache_key) if cache_key else None - if cached_handle is not None: - self._handle_cache.move_to_end(cache_key) - self._keyset_handle = cached_handle + reader = JsonKeysetReader(keyset_path.read_text(encoding="utf-8")) + if keyset_config.cleartext: + handle = cleartext_keyset_handle.read(reader) else: - try: - reader = JsonKeysetReader(keyset_path.read_text(encoding="utf-8")) - if keyset_config.cleartext: - self._keyset_handle = cleartext_keyset_handle.read(reader) - else: - master_key_aead = keyset_config.master_key_aead - assert master_key_aead is not None - self._keyset_handle = read_keyset_handle(reader, master_key_aead) - except (OSError, TinkError) as error: - raise ImproperlyConfigured(f"Could not load keyset `{self.keyset_name}`.") from error - - if cache_key: - self._handle_cache[cache_key] = self._keyset_handle - self._handle_cache.move_to_end(cache_key) - while len(self._handle_cache) > self._cache_size: - self._handle_cache.popitem(last=False) - - return self._keyset_handle + master_key_aead = keyset_config.master_key_aead + assert master_key_aead is not None + handle = read_keyset_handle(reader, master_key_aead) + except (OSError, TinkError) as error: + raise ImproperlyConfigured(f"Could not load keyset `{self.keyset_name}`.") from error - @cached_property - def aead_primitive(self) -> aead.Aead: - """Get the AEAD primitive for encryption/decryption operations. + self._entry = _KeysetEntry(handle) + if cache_key: + self._handle_cache[cache_key] = self._entry + self._handle_cache.move_to_end(cache_key) + while len(self._handle_cache) > self._cache_size: + self._handle_cache.popitem(last=False) - Returns: - aead.Aead: The AEAD primitive instance - """ - return self._get_tink_keyset_handle().primitive(aead.Aead) + return self._entry - @cached_property - def daead_primitive(self) -> daead.DeterministicAead: - """Get the Deterministic AEAD primitive for encryption/decryption operations. + def _get_tink_keyset_handle(self) -> Any: + """Return this manager's configured handle.""" + with self._cache_lock: + return self._get_keyset_entry().handle - Returns: - daead.DeterministicAead: The Deterministic AEAD primitive instance + def _get_primitive(self, primitive_class: type[Primitive]) -> Primitive: + # Keep construction and publication atomic with respect to clear_cache(). + # cached_property publishes after its getter returns, outside this lock. + with self._cache_lock: + entry = self._get_keyset_entry() + if primitive_class not in entry.primitives: + entry.primitives[primitive_class] = entry.handle.primitive(primitive_class) + return entry.primitives[primitive_class] - Raises: - ImproperlyConfigured: If deterministic AEAD is not available or keyset doesn't support it - """ + @property + def aead_primitive(self) -> aead.Aead: + """Get the AEAD primitive shared by managers using this keyset.""" + return self._get_primitive(aead.Aead) + + @property + def daead_primitive(self) -> daead.DeterministicAead: + """Get the deterministic AEAD primitive shared by this keyset.""" try: - return self._get_tink_keyset_handle().primitive(daead.DeterministicAead) + return self._get_primitive(daead.DeterministicAead) except TinkError as error: raise ImproperlyConfigured( "Current keyset does not support deterministic AEAD. " diff --git a/tink_fields/test/test_keyset_cache.py b/tink_fields/test/test_keyset_cache.py new file mode 100644 index 0000000..ba8c995 --- /dev/null +++ b/tink_fields/test/test_keyset_cache.py @@ -0,0 +1,118 @@ +"""Shared primitive construction, bounded storage, and concurrent invalidation.""" + +import gc +import json +import weakref +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from threading import Event +from unittest.mock import patch + +import pytest +from django.test import override_settings +from tink import KeysetHandle, aead, daead, json_proto_keyset_format, new_keyset_handle, secret_key_access + +from tink_fields import clear_keyset_cache +from tink_fields.fields import KeysetManager + + +@pytest.fixture(autouse=True) +def empty_keyset_cache(): + clear_keyset_cache() + yield + clear_keyset_cache() + + +@pytest.mark.parametrize("keyset, attribute", [("default", "aead_primitive"), ("deterministic", "daead_primitive")]) +def test_managers_share_primitive_construction(keyset, attribute): + managers = [KeysetManager(keyset) for _ in range(100)] + original = KeysetHandle.primitive + with patch.object(KeysetHandle, "primitive", autospec=True, side_effect=original) as construct: + primitives = [getattr(manager, attribute) for manager in managers] + assert construct.call_count == 1 + assert all(primitive is primitives[0] for primitive in primitives) + + +@pytest.mark.parametrize("keyset, attribute", [("default", "aead_primitive"), ("deterministic", "daead_primitive")]) +def test_clear_cache_waits_for_inflight_primitive_construction(keyset, attribute): + manager = KeysetManager(keyset) + building = Event() + release = Event() + clearing = Event() + cleared = Event() + original = KeysetHandle.primitive + + def slow_primitive(handle, primitive_class): + building.set() + assert release.wait(5) + return original(handle, primitive_class) + + def clear(): + clearing.set() + clear_keyset_cache() + cleared.set() + + with ThreadPoolExecutor(max_workers=2) as pool: + with patch.object(KeysetHandle, "primitive", slow_primitive): + constructing = pool.submit(getattr, manager, attribute) + try: + assert building.wait(5) + invalidating = pool.submit(clear) + assert clearing.wait(5) + assert not cleared.wait(0.1), "cache cleared before the old primitive could finish publishing" + finally: + release.set() + old = constructing.result(timeout=5) + invalidating.result(timeout=5) + assert getattr(manager, attribute) is not old + + +@pytest.mark.parametrize( + "template, attribute, encrypt, decrypt", + [ + (aead.aead_key_templates.AES128_GCM, "aead_primitive", "encrypt", "decrypt"), + ( + daead.deterministic_aead_key_templates.AES256_SIV, + "daead_primitive", + "encrypt_deterministically", + "decrypt_deterministically", + ), + ], +) +def test_rotation_reloads_all_active_managers(tmp_path, template, attribute, encrypt, decrypt): + old_keyset = json.loads(json_proto_keyset_format.serialize(new_keyset_handle(template), secret_key_access.TOKEN)) + new_keyset = json.loads(json_proto_keyset_format.serialize(new_keyset_handle(template), secret_key_access.TOKEN)) + new_keyset["key"].extend(old_keyset["key"]) + path = tmp_path / "keys.json" + path.write_text(json.dumps(old_keyset), encoding="utf-8") + with override_settings(TINK_FIELDS_CONFIG={"default": {"path": path, "cleartext": True}}): + managers = [KeysetManager("default") for _ in range(3)] + before = [getattr(manager, attribute) for manager in managers] + ciphertext = getattr(before[0], encrypt)(b"secret", b"aad") + replacement = tmp_path / "replacement.json" + replacement.write_text(json.dumps(new_keyset), encoding="utf-8") + replacement.replace(path) + clear_keyset_cache() + for manager, old_primitive in zip(managers, before, strict=True): + primitive = getattr(manager, attribute) + assert primitive is not old_primitive + assert getattr(primitive, decrypt)(ciphertext, b"aad") == b"secret" + assert getattr(primitive, encrypt)(b"secret", b"aad")[1:5] == new_keyset["primaryKeyId"].to_bytes(4, "big") + + +def test_shared_cache_is_bounded_and_does_not_retain_managers(tmp_path): + config = {} + source = Path(__file__).with_name("test_plaintext_keyset.json").read_text(encoding="utf-8") + for index in range(3): + path = tmp_path / f"keys-{index}.json" + path.write_text(source, encoding="utf-8") + config[str(index)] = {"path": path, "cleartext": True} + with override_settings(TINK_FIELDS_CONFIG=config), patch.object(KeysetManager, "_cache_size", 2): + for name in config: + manager = KeysetManager(name) + _ = manager.aead_primitive + assert len(KeysetManager._handle_cache) == 2 + reference = weakref.ref(manager) + del manager + gc.collect() + assert reference() is None