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
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
37 changes: 37 additions & 0 deletions benchmarks/keyset_cache.py
Original file line number Diff line number Diff line change
@@ -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()
149 changes: 77 additions & 72 deletions tink_fields/fields.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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. "
Expand Down
118 changes: 118 additions & 0 deletions tink_fields/test/test_keyset_cache.py
Original file line number Diff line number Diff line change
@@ -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
Loading