From c97526499c23fed4120de29e11a95291d90d348a Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Fri, 8 May 2026 17:03:13 +0200 Subject: [PATCH 01/12] encryption: implement DEK caching --- encryption/cache.go | 179 +++++++++++++ encryption/cache_test.go | 384 ++++++++++++++++++++++++++++ encryption/pool.go | 129 +++++++--- encryption/pool_test.go | 536 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 1199 insertions(+), 29 deletions(-) create mode 100644 encryption/cache.go create mode 100644 encryption/cache_test.go diff --git a/encryption/cache.go b/encryption/cache.go new file mode 100644 index 0000000..54ed107 --- /dev/null +++ b/encryption/cache.go @@ -0,0 +1,179 @@ +package encryption + +import ( + "container/list" + "sync" + "time" +) + +// CacheConfig configures the optional DEK (data encryption key) cache. +// A zero-value config disables caching. +type CacheConfig struct { + // MaxSize is the maximum number of DEKs to cache. Must be > 0 to enable caching. + MaxSize int + // TTL is the time-to-live for each cache entry. Must be > 0 to enable caching. + TTL time.Duration +} + +// dekCacheEntry holds a cached DEK and its LRU/TTL metadata. +type dekCacheEntry struct { + dek []byte // 32-byte AES-256 key (owned copy) + keyRef string // for reverse lookup from LRU list element + expiresAt time.Time + element *list.Element // back-pointer into LRU list +} + +// dekCache is a thread-safe LRU cache for decrypted data encryption keys. +// It zeroes key material on every eviction path (TTL, LRU, delete, clear). +type dekCache struct { + mu sync.Mutex + entries map[string]*dekCacheEntry + order *list.List // front = most recently used + maxSize int + ttl time.Duration + + inflightMu sync.Mutex + inflight map[string]*inflightEntry +} + +// inflightEntry coordinates singleflight deduplication for concurrent +// cache misses on the same keyRef. +type inflightEntry struct { + done chan struct{} + dek []byte + err error +} + +func newDEKCache(maxSize int, ttl time.Duration) *dekCache { + return &dekCache{ + entries: make(map[string]*dekCacheEntry), + order: list.New(), + maxSize: maxSize, + ttl: ttl, + inflight: make(map[string]*inflightEntry), + } +} + +// get returns a copy of the cached DEK for keyRef, or ok=false on miss/expiry. +func (c *dekCache) get(keyRef string) ([]byte, bool) { + c.mu.Lock() + defer c.mu.Unlock() + + entry, ok := c.entries[keyRef] + if !ok { + return nil, false + } + + if time.Now().After(entry.expiresAt) { + c.evictLocked(entry) + return nil, false + } + + c.order.MoveToFront(entry.element) + return copyBytes(entry.dek), true +} + +// put stores a copy of dek in the cache, evicting the LRU entry if full. +func (c *dekCache) put(keyRef string, dek []byte) { + c.mu.Lock() + defer c.mu.Unlock() + + if entry, ok := c.entries[keyRef]; ok { + // Update existing entry. + clear(entry.dek) + entry.dek = copyBytes(dek) + entry.expiresAt = time.Now().Add(c.ttl) + c.order.MoveToFront(entry.element) + return + } + + entry := &dekCacheEntry{ + dek: copyBytes(dek), + keyRef: keyRef, + expiresAt: time.Now().Add(c.ttl), + } + entry.element = c.order.PushFront(entry) + c.entries[keyRef] = entry + + if len(c.entries) > c.maxSize { + back := c.order.Back() + if back != nil { + c.evictLocked(back.Value.(*dekCacheEntry)) + } + } +} + +// delete removes and zeroes a specific entry. Called by RotateKey. +func (c *dekCache) delete(keyRef string) { + c.mu.Lock() + defer c.mu.Unlock() + + if entry, ok := c.entries[keyRef]; ok { + c.evictLocked(entry) + } +} + +// clear removes and zeroes all entries. +func (c *dekCache) clear() { + c.mu.Lock() + defer c.mu.Unlock() + + for _, entry := range c.entries { + clear(entry.dek) + } + c.entries = make(map[string]*dekCacheEntry) + c.order.Init() +} + +// evictLocked removes an entry, zeroing its DEK. Caller must hold c.mu. +func (c *dekCache) evictLocked(entry *dekCacheEntry) { + clear(entry.dek) + c.order.Remove(entry.element) + delete(c.entries, entry.keyRef) +} + +// waitOrStart implements singleflight deduplication. If another goroutine is +// already fetching the DEK for keyRef, started=false and wait blocks until the +// result is available. Otherwise started=true and the caller must call finish. +func (c *dekCache) waitOrStart(keyRef string) (started bool, wait func() ([]byte, error)) { + c.inflightMu.Lock() + + if entry, ok := c.inflight[keyRef]; ok { + c.inflightMu.Unlock() + return false, func() ([]byte, error) { + <-entry.done + if entry.err != nil { + return nil, entry.err + } + return copyBytes(entry.dek), nil + } + } + + entry := &inflightEntry{done: make(chan struct{})} + c.inflight[keyRef] = entry + c.inflightMu.Unlock() + + return true, nil +} + +// finish signals all waiters for keyRef with the fetch result. +func (c *dekCache) finish(keyRef string, dek []byte, err error) { + c.inflightMu.Lock() + entry, ok := c.inflight[keyRef] + if ok { + entry.dek = copyBytes(dek) + entry.err = err + close(entry.done) + delete(c.inflight, keyRef) + } + c.inflightMu.Unlock() +} + +func copyBytes(b []byte) []byte { + if b == nil { + return nil + } + cp := make([]byte, len(b)) + copy(cp, b) + return cp +} diff --git a/encryption/cache_test.go b/encryption/cache_test.go new file mode 100644 index 0000000..4de3f05 --- /dev/null +++ b/encryption/cache_test.go @@ -0,0 +1,384 @@ +package encryption + +import ( + "sync" + "testing" + "time" +) + +func TestDEKCache_GetPut(t *testing.T) { + c := newDEKCache(10, time.Minute) + dek := []byte("0123456789abcdef0123456789abcdef") + + c.put("ref1", dek) + + got, ok := c.get("ref1") + if !ok { + t.Fatal("expected cache hit") + } + if string(got) != string(dek) { + t.Fatalf("got %x, want %x", got, dek) + } + + // Returned slice must be an independent copy. + got[0] = 0xFF + got2, ok := c.get("ref1") + if !ok { + t.Fatal("expected cache hit") + } + if got2[0] == 0xFF { + t.Fatal("cache returned same underlying slice, expected independent copy") + } + + // Stored slice must be an independent copy of the input. + dek[0] = 0xAA + got3, ok := c.get("ref1") + if !ok { + t.Fatal("expected cache hit") + } + if got3[0] == 0xAA { + t.Fatal("cache stored same underlying slice as input, expected independent copy") + } +} + +func TestDEKCache_Miss(t *testing.T) { + c := newDEKCache(10, time.Minute) + + _, ok := c.get("unknown") + if ok { + t.Fatal("expected cache miss for unknown key") + } +} + +func TestDEKCache_TTLExpiry(t *testing.T) { + c := newDEKCache(10, time.Millisecond) + dek := []byte("0123456789abcdef0123456789abcdef") + + c.put("ref1", dek) + + // Grab internal slice before expiry. + c.mu.Lock() + internalSlice := c.entries["ref1"].dek + c.mu.Unlock() + + time.Sleep(5 * time.Millisecond) + + _, ok := c.get("ref1") + if ok { + t.Fatal("expected cache miss after TTL expiry") + } + + // Expired entry should be removed from map. + c.mu.Lock() + _, exists := c.entries["ref1"] + c.mu.Unlock() + if exists { + t.Fatal("expired entry should have been removed from map") + } + + // Internal DEK should be zeroed. + for _, b := range internalSlice { + if b != 0 { + t.Fatal("expired DEK should have been zeroed") + } + } +} + +func TestDEKCache_LRUEviction(t *testing.T) { + c := newDEKCache(2, time.Minute) + + c.put("ref1", []byte("key1key1key1key1key1key1key1key1")) + c.put("ref2", []byte("key2key2key2key2key2key2key2key2")) + + // Grab internal slice of ref2 before eviction. + c.mu.Lock() + internalRef2 := c.entries["ref2"].dek + c.mu.Unlock() + + // Access ref1 to make it more recent than ref2. + _, _ = c.get("ref1") + + // Adding ref3 should evict ref2 (LRU). + c.put("ref3", []byte("key3key3key3key3key3key3key3key3")) + + if _, ok := c.get("ref2"); ok { + t.Fatal("expected ref2 to be evicted (LRU)") + } + if _, ok := c.get("ref1"); !ok { + t.Fatal("expected ref1 to still be cached") + } + if _, ok := c.get("ref3"); !ok { + t.Fatal("expected ref3 to still be cached") + } + + // Verify evicted internal DEK was zeroed. + for _, b := range internalRef2 { + if b != 0 { + t.Fatal("evicted DEK should have been zeroed") + } + } +} + +func TestDEKCache_PutUpdatesExisting(t *testing.T) { + c := newDEKCache(10, time.Minute) + dek1 := []byte("old_key_old_key_old_key_old_key_") + dek2 := []byte("new_key_new_key_new_key_new_key_") + + c.put("ref1", dek1) + + // Grab internal slice and expiry before update. + c.mu.Lock() + internalSlice := c.entries["ref1"].dek + oldExpiry := c.entries["ref1"].expiresAt + c.mu.Unlock() + + time.Sleep(time.Millisecond) // ensure time advances + c.put("ref1", dek2) + + got, ok := c.get("ref1") + if !ok { + t.Fatal("expected cache hit") + } + if string(got) != string(dek2) { + t.Fatalf("got %x, want %x", got, dek2) + } + + // Old internal slice should be zeroed. + for _, b := range internalSlice { + if b != 0 { + t.Fatal("old DEK slice should have been zeroed") + } + } + + // TTL should be refreshed. + c.mu.Lock() + newExpiry := c.entries["ref1"].expiresAt + c.mu.Unlock() + if !newExpiry.After(oldExpiry) { + t.Fatal("put on existing key should refresh TTL") + } +} + +func TestDEKCache_Delete(t *testing.T) { + c := newDEKCache(10, time.Minute) + dek := []byte("0123456789abcdef0123456789abcdef") + + c.put("ref1", dek) + + // Grab internal slice reference. + c.mu.Lock() + internalSlice := c.entries["ref1"].dek + c.mu.Unlock() + + c.delete("ref1") + + if _, ok := c.get("ref1"); ok { + t.Fatal("expected cache miss after delete") + } + + // Verify zeroed. + for _, b := range internalSlice { + if b != 0 { + t.Fatal("deleted DEK should have been zeroed") + } + } +} + +func TestDEKCache_DeleteMissing(t *testing.T) { + c := newDEKCache(10, time.Minute) + // Should not panic. + c.delete("nonexistent") +} + +func TestDEKCache_Clear(t *testing.T) { + c := newDEKCache(10, time.Minute) + + dek1 := []byte("key1key1key1key1key1key1key1key1") + dek2 := []byte("key2key2key2key2key2key2key2key2") + c.put("ref1", dek1) + c.put("ref2", dek2) + + // Grab internal slice references. + c.mu.Lock() + internal1 := c.entries["ref1"].dek + internal2 := c.entries["ref2"].dek + c.mu.Unlock() + + c.clear() + + if _, ok := c.get("ref1"); ok { + t.Fatal("expected miss after clear") + } + if _, ok := c.get("ref2"); ok { + t.Fatal("expected miss after clear") + } + + for _, b := range internal1 { + if b != 0 { + t.Fatal("cleared DEK 1 should be zeroed") + } + } + for _, b := range internal2 { + if b != 0 { + t.Fatal("cleared DEK 2 should be zeroed") + } + } +} + +func TestDEKCache_ClearThenReuse(t *testing.T) { + c := newDEKCache(10, time.Minute) + dek := []byte("0123456789abcdef0123456789abcdef") + + c.put("ref1", dek) + c.clear() + + // Cache should work normally after clear. + c.put("ref1", dek) + got, ok := c.get("ref1") + if !ok { + t.Fatal("expected cache hit after clear + put") + } + if string(got) != string(dek) { + t.Fatalf("got %x, want %x", got, dek) + } +} + +func TestDEKCache_Singleflight(t *testing.T) { + c := newDEKCache(10, time.Minute) + + dek := []byte("0123456789abcdef0123456789abcdef") + + // First caller starts the fetch. + started, _ := c.waitOrStart("ref1") + if !started { + t.Fatal("first caller should start") + } + + // Second and third callers should wait. + var wg sync.WaitGroup + results := make([][]byte, 2) + errs := make([]error, 2) + + for i := 0; i < 2; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + s, wait := c.waitOrStart("ref1") + if s { + t.Error("subsequent caller should not start") + return + } + results[idx], errs[idx] = wait() + }(i) + } + + // Simulate fetch completing. + time.Sleep(10 * time.Millisecond) // let goroutines reach wait() + c.finish("ref1", dek, nil) + + wg.Wait() + + for i := 0; i < 2; i++ { + if errs[i] != nil { + t.Fatalf("waiter %d got error: %v", i, errs[i]) + } + if string(results[i]) != string(dek) { + t.Fatalf("waiter %d got wrong dek", i) + } + } + + // Each waiter should have received an independent copy. + results[0][0] = 0xFF + if results[1][0] == 0xFF { + t.Fatal("waiters should receive independent copies") + } +} + +func TestDEKCache_SingleflightError(t *testing.T) { + c := newDEKCache(10, time.Minute) + fetchErr := &testError{msg: "kms failed"} + + started, _ := c.waitOrStart("ref1") + if !started { + t.Fatal("first caller should start") + } + + var wg sync.WaitGroup + waiterErrs := make([]error, 2) + + for i := 0; i < 2; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + _, wait := c.waitOrStart("ref1") + _, waiterErrs[idx] = wait() + }(i) + } + + time.Sleep(10 * time.Millisecond) + c.finish("ref1", nil, fetchErr) + + wg.Wait() + + for i := 0; i < 2; i++ { + if waiterErrs[i] == nil { + t.Fatalf("waiter %d should have received error", i) + } + if waiterErrs[i].Error() != "kms failed" { + t.Fatalf("waiter %d got error %q, want %q", i, waiterErrs[i].Error(), "kms failed") + } + } +} + +func TestDEKCache_SingleflightIndependentKeys(t *testing.T) { + c := newDEKCache(10, time.Minute) + + dek1 := []byte("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + dek2 := []byte("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb") + + // Start fetch for ref1. + started1, _ := c.waitOrStart("ref1") + if !started1 { + t.Fatal("first caller for ref1 should start") + } + + // Start fetch for ref2 — should NOT be blocked by ref1. + started2, _ := c.waitOrStart("ref2") + if !started2 { + t.Fatal("first caller for ref2 should start independently") + } + + c.finish("ref1", dek1, nil) + c.finish("ref2", dek2, nil) +} + +func TestDEKCache_ConcurrentAccess(t *testing.T) { + c := newDEKCache(10, time.Minute) + dek := []byte("0123456789abcdef0123456789abcdef") + + var wg sync.WaitGroup + for i := 0; i < 100; i++ { + wg.Add(3) + go func(n int) { + defer wg.Done() + c.put("ref1", dek) + }(i) + go func(n int) { + defer wg.Done() + c.get("ref1") + }(i) + go func(n int) { + defer wg.Done() + if n%10 == 0 { + c.delete("ref1") + } + }(i) + } + wg.Wait() +} + +type testError struct { + msg string +} + +func (e *testError) Error() string { return e.msg } diff --git a/encryption/pool.go b/encryption/pool.go index 0a5fe97..204ccbf 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -42,19 +42,38 @@ type Pool struct { keysTable KeysTable dataTables []EncryptedDataTable logger *slog.Logger + cache *dekCache // nil when caching disabled } -func NewPool(attester Attester, configs []*Config, keysTable KeysTable, dataTables []EncryptedDataTable, logger *slog.Logger) *Pool { +// PoolOption configures optional Pool behavior. +type PoolOption func(*Pool) + +// WithCache enables an in-memory LRU cache for decrypted data encryption keys, +// eliminating KMS round-trips on cache hits. The cache is local to this process +// and zeroes key material on eviction. +func WithCache(cfg CacheConfig) PoolOption { + return func(p *Pool) { + if cfg.MaxSize > 0 && cfg.TTL > 0 { + p.cache = newDEKCache(cfg.MaxSize, cfg.TTL) + } + } +} + +func NewPool(attester Attester, configs []*Config, keysTable KeysTable, dataTables []EncryptedDataTable, logger *slog.Logger, opts ...PoolOption) *Pool { if logger == nil { logger = slog.Default() } - return &Pool{ + p := &Pool{ attester: attester, configs: configs, keysTable: keysTable, dataTables: dataTables, logger: logger, } + for _, opt := range opts { + opt(p) + } + return p } // Encrypt encrypts the plaintext using a randomly selected cipher key from the Pool. It returns the key reference @@ -95,17 +114,32 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext if err != nil { return "", nil, fmt.Errorf("generate key: %w", err) } - } else if err := p.VerifyKey(ctx, att, key); err != nil { - return "", nil, fmt.Errorf("verify key: %w", err) - } - span.SetAnnotation("key_ref", key.KeyRef) - if privateKey == nil { - privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) - if err != nil { - return "", nil, fmt.Errorf("combine shares: %w", err) + if p.cache != nil { + p.cache.put(key.KeyRef, privateKey) + } + } else { + // Existing key — try cache before KMS. + if p.cache != nil { + if dek, ok := p.cache.get(key.KeyRef); ok { + privateKey = dek + } + } + if privateKey == nil { + if err := p.VerifyKey(ctx, att, key); err != nil { + return "", nil, fmt.Errorf("verify key: %w", err) + } + privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) + if err != nil { + return "", nil, fmt.Errorf("combine shares: %w", err) + } + + if p.cache != nil { + p.cache.put(key.KeyRef, privateKey) + } } } + span.SetAnnotation("key_ref", key.KeyRef) encrypted, err := aesgcm.Encrypt(att, privateKey, plaintext, additionalData) if err != nil { @@ -127,6 +161,7 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext // Decrypt decrypts the ciphertext using the latest cipher key from the Pool referenced by the keyRef. // // The key is verified against the attestation and migrated to the current generation if needed. +// If a DEK cache is configured, cached keys bypass DynamoDB and KMS on hit. func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef string, ciphertext []byte, additionalData []byte) (plaintext []byte, err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.Decrypt", tracing.WithAnnotation("key_ref", keyRef)) defer func() { @@ -139,6 +174,50 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, fmt.Errorf("decode ciphertext: %w", err) } + // Try cache. + var privateKey []byte + if p.cache != nil { + if dek, ok := p.cache.get(keyRef); ok { + privateKey = dek + } + } + + // Cache miss — full fetch with singleflight dedup. + if privateKey == nil { + privateKey, err = p.fetchDEK(ctx, att, keyRef) + if err != nil { + return nil, err + } + } + + // Decrypt data. + var decrypted []byte + switch decoded.Version { + case 1: + decrypted, err = aescbc.Decrypt(privateKey, decoded.EncryptedData) + case 2, 3: + decrypted, err = aesgcm.Decrypt(privateKey, decoded.EncryptedData, additionalData) + } + if err != nil { + return nil, fmt.Errorf("decrypt: %w", err) + } + + return decrypted, nil +} + +// fetchDEK retrieves a DEK through the full path: DynamoDB lookup, attestation +// verification, KMS share decryption, and Shamir combine. It uses singleflight +// to deduplicate concurrent fetches for the same keyRef, and populates the cache. +func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, err error) { + // Singleflight: if another goroutine is already fetching this keyRef, wait. + if p.cache != nil { + started, wait := p.cache.waitOrStart(keyRef) + if !started { + return wait() + } + defer func() { p.cache.finish(keyRef, privateKey, err) }() + } + key, found, err := p.keysTable.GetLatestByKeyRef(ctx, keyRef, false) if err != nil { return nil, fmt.Errorf("get latest key: %w", err) @@ -150,11 +229,6 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, fmt.Errorf("verify key: %w", err) } - span.SetAnnotation("generation", strconv.Itoa(key.Generation)) - if key.KeyIndex != nil { - span.SetAnnotation("key_index", strconv.Itoa(*key.KeyIndex)) - } - config, err := p.getConfig(key.Generation) if err != nil { return nil, fmt.Errorf("get config: %w", err) @@ -163,31 +237,24 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, fmt.Errorf("shares are invalid") } - privateKey, err := p.combineShares(ctx, att, config, key.EncryptedShares) + privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) if err != nil { return nil, fmt.Errorf("combine shares: %w", err) } - var decrypted []byte - switch decoded.Version { - case 1: - decrypted, err = aescbc.Decrypt(privateKey, decoded.EncryptedData) - case 2, 3: - decrypted, err = aesgcm.Decrypt(privateKey, decoded.EncryptedData, additionalData) - } - if err != nil { - return nil, fmt.Errorf("decrypt: %w", err) + if p.cache != nil { + p.cache.put(keyRef, privateKey) } + // Trigger migration if needed. Migration is synchronous but non-fatal: + // failure is logged and does not affect the returned DEK. if p.keyNeedsMigration(key) { - err := p.migrateKey(ctx, att, key, privateKey) - if err != nil { - // We don't want to fail the decryption if migration fails, log the error and continue + if err := p.migrateKey(ctx, att, key, privateKey); err != nil { p.logger.ErrorContext(ctx, "migrating key failed", "error", err, "key_ref", key.KeyRef, "generation", key.Generation, "key_index", key.KeyIndex) } } - return decrypted, nil + return privateKey, nil } // RotateKey marks a key as inactive by setting its KeyIndex to a negative value. It won't be used for encrypting @@ -232,6 +299,10 @@ func (p *Pool) RotateKey(ctx context.Context, att *enclave.Attestation, keyRef s return fmt.Errorf("deactivate key: %w", err) } + if p.cache != nil { + p.cache.delete(keyRef) + } + return nil } diff --git a/encryption/pool_test.go b/encryption/pool_test.go index ef1b65a..6ac4d40 100644 --- a/encryption/pool_test.go +++ b/encryption/pool_test.go @@ -5,6 +5,7 @@ import ( "crypto/x509" "encoding/pem" "errors" + "sync" "testing" "time" @@ -1062,3 +1063,538 @@ func TestPool_CleanupUnusedKeys(t *testing.T) { keysTable.AssertCalled(t, "Delete", mock.Anything, "old-key", 0) }) } + +func TestPool_DecryptCacheHit(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + // Encrypt to get a ciphertext (this also populates the cache via Encrypt path). + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKey, true, nil) + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + keyRef, ciphertext, err := pool.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + + callsAfterEncrypt := len(remoteKey1.Calls) + + // Decrypt — cache hit from Encrypt's put. No additional KMS calls. + plaintext, err := pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + require.Equal(t, callsAfterEncrypt, len(remoteKey1.Calls), "Decrypt should hit cache populated by Encrypt") + + // Second Decrypt — still a cache hit. + plaintext, err = pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + require.Equal(t, callsAfterEncrypt, len(remoteKey1.Calls), "repeated Decrypt should hit cache") +} + +// TestPool_DecryptPopulatesCache verifies the Decrypt path itself populates the +// cache (via fetchDEK), independent of Encrypt. +func TestPool_DecryptPopulatesCache(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKey, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + + // Encrypt WITHOUT cache to get a ciphertext. + poolNoCache := encryption.NewPool(enc, configs, keysTable, nil, nil) + keyRef, ciphertext, err := poolNoCache.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + + // Create a NEW pool with cache — cache is cold. + poolWithCache := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + // First Decrypt — cache miss, calls KMS. + plaintext, err := poolWithCache.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + + callsAfterFirst := len(remoteKey1.Calls) + require.Greater(t, callsAfterFirst, 0, "first Decrypt should have called KMS") + + // Second Decrypt — cache hit from fetchDEK's put. No additional KMS calls. + plaintext, err = poolWithCache.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + require.Equal(t, callsAfterFirst, len(remoteKey1.Calls), "second Decrypt should hit cache populated by fetchDEK") +} + +func TestPool_EncryptCacheHit(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + // First encrypt — cache miss, calls KMS Decrypt to combine shares. + _, _, err = pool.Encrypt(context.Background(), att, []byte("test1"), []byte("aad")) + require.NoError(t, err) + + decrypt1Calls := len(remoteKey1.Calls) + + // Second encrypt — same key index (deterministic random), should be a cache hit. + _, _, err = pool.Encrypt(context.Background(), att, []byte("test2"), []byte("aad")) + require.NoError(t, err) + + require.Equal(t, decrypt1Calls, len(remoteKey1.Calls), "expected no additional KMS calls on encrypt cache hit") +} + +func TestPool_RotateInvalidatesCache(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKey, true, nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", true).Return(cipherKey, true, nil) + keysTable.On("Deactivate", mock.Anything, "cipherKey4", 0, mock.AnythingOfType("time.Time"), mock.Anything).Return(nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + // Encrypt to get ciphertext + populate cache. + keyRef, ciphertext, err := pool.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + + // Decrypt — should be cache hit. + plaintext, err := pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + + callsBefore := len(remoteKey1.Calls) + + // Rotate — should invalidate cache. + err = pool.RotateKey(context.Background(), att, "cipherKey4") + require.NoError(t, err) + + // Decrypt again — should be cache miss, hitting KMS again. + plaintext, err = pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + + require.Greater(t, len(remoteKey1.Calls), callsBefore, "expected additional KMS calls after cache invalidation") +} + +func TestPool_NoCacheByDefault(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKey, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + + // No WithCache option. + pool := encryption.NewPool(enc, configs, keysTable, nil, nil) + + keyRef, ciphertext, err := pool.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + + // First decrypt. + _, err = pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + callsAfterFirst := len(remoteKey1.Calls) + + // Second decrypt — no cache, so KMS is called again. + _, err = pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + + require.Greater(t, len(remoteKey1.Calls), callsAfterFirst, "without cache, every decrypt should call KMS") +} + +func TestPool_DecryptSingleflight(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + // First, encrypt to get a ciphertext. + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKey, true, nil) + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + _, ciphertext, err := pool.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + + // Clear mock call history so we count only decrypt-path calls. + remoteKey1.Calls = nil + remoteKey2.Calls = nil + remoteKey1.ExpectedCalls = nil + remoteKey2.ExpectedCalls = nil + + // Re-register expectations. + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) + + // Invalidate cache so all goroutines start with a cold cache. + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", true).Return(cipherKey, true, nil) + keysTable.On("Deactivate", mock.Anything, "cipherKey4", 0, mock.AnythingOfType("time.Time"), mock.Anything).Return(nil) + err = pool.RotateKey(context.Background(), att, "cipherKey4") + require.NoError(t, err) + + // Launch N concurrent decrypts. + const N = 10 + var wg sync.WaitGroup + errs := make([]error, N) + results := make([]string, N) + for i := 0; i < N; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + pt, err := pool.Decrypt(context.Background(), att, "cipherKey4", ciphertext, []byte("aad")) + errs[idx] = err + if pt != nil { + results[idx] = string(pt) + } + }(i) + } + wg.Wait() + + for i := 0; i < N; i++ { + require.NoError(t, errs[i], "goroutine %d failed", i) + require.Equal(t, "test", results[i], "goroutine %d got wrong result", i) + } + + // With singleflight, exactly one goroutine fetches. RemoteKey1.Decrypt + // should be called once (one share per remote key, one fetch total). + decryptCalls := len(remoteKey1.Calls) + require.Equal(t, 1, decryptCalls, "singleflight should deduplicate concurrent fetches (got %d calls to remoteKey1)", decryptCalls) +} + +func TestPool_EncryptNewKeyPopulatesCache(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + // Key does not exist — GenerateKey will be called. + keysTable.On("Get", mock.Anything, 0, 4).Return(nil, false, nil).Once() + remoteKey1.On("Encrypt", mock.Anything, att, mock.Anything).Return("encryptedShare1", nil) + remoteKey2.On("Encrypt", mock.Anything, att, mock.Anything).Return("encryptedShare2", nil) + keysTable.On("Create", mock.Anything, mock.Anything).Return(false, nil) + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + // First Encrypt — GenerateKey creates key and caches DEK. + keyRef, ciphertext, err := pool.Encrypt(context.Background(), att, []byte("test"), []byte("aad")) + require.NoError(t, err) + require.NotEmpty(t, keyRef) + + // Decrypt should hit cache — no KMS Decrypt calls needed. + cipherKey := &data.CipherKey{ + Generation: 0, + KeyIndex: intPtr(4), + KeyRef: keyRef, + EncryptedShares: map[string]string{ + "remoteKey1": "encryptedShare1", + "remoteKey2": "encryptedShare2", + }, + CreatedAt: time.Now(), + } + hash, err := cipherKey.Hash() + require.NoError(t, err) + cipherKeyAtt, err := enc.GetAttestation(context.Background(), nil, hash) + require.NoError(t, err) + cipherKey.Attestation = cipherKeyAtt.Document() + _ = cipherKeyAtt.Close() + + keysTable.On("GetLatestByKeyRef", mock.Anything, keyRef, false).Return(cipherKey, true, nil) + + // RemoteKey.Decrypt should NOT be called — DEK is cached from GenerateKey. + plaintext, err := pool.Decrypt(context.Background(), att, keyRef, ciphertext, []byte("aad")) + require.NoError(t, err) + require.Equal(t, "test", string(plaintext)) + + remoteKey1.AssertNotCalled(t, "Decrypt", mock.Anything, mock.Anything, mock.Anything) + remoteKey2.AssertNotCalled(t, "Decrypt", mock.Anything, mock.Anything, mock.Anything) +} + +func TestPool_MultipleKeyRefs(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + kmsClient := &MockKMS{} + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + random := &constantReader{value: 0x42} + enc, err := enclave.New(context.Background(), enclave.DummyProvider(random), kmsClient, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{ + { + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }, + } + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + defer func() { _ = att.Close() }() + + // Create two distinct cipher keys. + cipherKeyA, privateKeyA := newCipherKey(t, enc) + cipherKeyB, privateKeyB := newCipherKey(t, enc, func(key *data.CipherKey) { + key.KeyRef = "cipherKeyB" + }) + + sharesA, err := shamir.Split(privateKeyA, 2, 2) + require.NoError(t, err) + sharesB, err := shamir.Split(privateKeyB, 2, 2) + require.NoError(t, err) + + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKey4", false).Return(cipherKeyA, true, nil) + keysTable.On("GetLatestByKeyRef", mock.Anything, "cipherKeyB", false).Return(cipherKeyB, true, nil) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(sharesA[0], nil).Once() + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(sharesA[1], nil).Once() + + pool := encryption.NewPool(enc, configs, keysTable, nil, nil, + encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + // Decrypt keyA — cache miss, populates cache for keyA. + _, err = pool.Decrypt(context.Background(), att, "cipherKey4", legacyCiphertext55_v2, []byte("aad")) + require.NoError(t, err) + + // Decrypt keyA again — cache hit, no KMS. + callsAfterA := len(remoteKey1.Calls) + _, err = pool.Decrypt(context.Background(), att, "cipherKey4", legacyCiphertext55_v2, []byte("aad")) + require.NoError(t, err) + require.Equal(t, callsAfterA, len(remoteKey1.Calls), "keyA should be a cache hit") + + // Decrypt keyB — cache miss for keyB (keyA cached, keyB not). + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(sharesB[0], nil).Once() + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(sharesB[1], nil).Once() + _, err = pool.Decrypt(context.Background(), att, "cipherKeyB", legacyCiphertext55_v2, []byte("aad")) + require.NoError(t, err) + require.Greater(t, len(remoteKey1.Calls), callsAfterA, "keyB should be a cache miss") + + // Decrypt keyB again — now a cache hit. + callsAfterB := len(remoteKey1.Calls) + _, err = pool.Decrypt(context.Background(), att, "cipherKeyB", legacyCiphertext55_v2, []byte("aad")) + require.NoError(t, err) + require.Equal(t, callsAfterB, len(remoteKey1.Calls), "keyB should now be a cache hit") +} + +func intPtr(v int) *int { return &v } From 4cf3d74172c9ef90646ef8ad8dc692b062c31fbc Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Tue, 22 Sep 2026 21:47:17 +0200 Subject: [PATCH 02/12] fix(encryption): verify cipher key row on DEK cache hit MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Encrypt consulted the DEK cache before VerifyKey, so a hit skipped the attestation check on a row read from DynamoDB — letting a tampered row choose which cached DEK encrypts the data. Verify first, then cache. Also: drop the unzeroed DEK copy the singleflight kept, make waiters respect context cancellation, evict deleted keys in CleanupUnusedKeys, and restore the generation/key_index span annotations. Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 57 +++++++++------------ encryption/cache_test.go | 107 ++++++++++++++++----------------------- encryption/pool.go | 102 ++++++++++++++++++++++++------------- encryption/pool_test.go | 74 ++++++++++++++++++++++++++- 4 files changed, 210 insertions(+), 130 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index 54ed107..f6ee93c 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -2,6 +2,7 @@ package encryption import ( "container/list" + "context" "sync" "time" ) @@ -17,8 +18,8 @@ type CacheConfig struct { // dekCacheEntry holds a cached DEK and its LRU/TTL metadata. type dekCacheEntry struct { - dek []byte // 32-byte AES-256 key (owned copy) - keyRef string // for reverse lookup from LRU list element + dek []byte // 32-byte AES-256 key (owned copy) + keyRef string // for reverse lookup from LRU list element expiresAt time.Time element *list.Element // back-pointer into LRU list } @@ -37,10 +38,11 @@ type dekCache struct { } // inflightEntry coordinates singleflight deduplication for concurrent -// cache misses on the same keyRef. +// cache misses on the same keyRef. It deliberately carries no key material: +// waiters read the result from the cache, so there is no second copy of a DEK +// to keep track of and zero. type inflightEntry struct { done chan struct{} - dek []byte err error } @@ -113,18 +115,6 @@ func (c *dekCache) delete(keyRef string) { } } -// clear removes and zeroes all entries. -func (c *dekCache) clear() { - c.mu.Lock() - defer c.mu.Unlock() - - for _, entry := range c.entries { - clear(entry.dek) - } - c.entries = make(map[string]*dekCacheEntry) - c.order.Init() -} - // evictLocked removes an entry, zeroing its DEK. Caller must hold c.mu. func (c *dekCache) evictLocked(entry *dekCacheEntry) { clear(entry.dek) @@ -133,19 +123,21 @@ func (c *dekCache) evictLocked(entry *dekCacheEntry) { } // waitOrStart implements singleflight deduplication. If another goroutine is -// already fetching the DEK for keyRef, started=false and wait blocks until the -// result is available. Otherwise started=true and the caller must call finish. -func (c *dekCache) waitOrStart(keyRef string) (started bool, wait func() ([]byte, error)) { +// already fetching the DEK for keyRef, started=false and wait blocks until that +// fetch finishes, reporting its error; the caller then reads the DEK from the +// cache. Otherwise started=true and the caller must call finish. +func (c *dekCache) waitOrStart(keyRef string) (started bool, wait func(ctx context.Context) error) { c.inflightMu.Lock() if entry, ok := c.inflight[keyRef]; ok { c.inflightMu.Unlock() - return false, func() ([]byte, error) { - <-entry.done - if entry.err != nil { - return nil, entry.err + return false, func(ctx context.Context) error { + select { + case <-entry.done: + return entry.err + case <-ctx.Done(): + return ctx.Err() } - return copyBytes(entry.dek), nil } } @@ -156,17 +148,18 @@ func (c *dekCache) waitOrStart(keyRef string) (started bool, wait func() ([]byte return true, nil } -// finish signals all waiters for keyRef with the fetch result. -func (c *dekCache) finish(keyRef string, dek []byte, err error) { +// finish signals all waiters for keyRef with the fetch outcome. +func (c *dekCache) finish(keyRef string, err error) { c.inflightMu.Lock() + defer c.inflightMu.Unlock() + entry, ok := c.inflight[keyRef] - if ok { - entry.dek = copyBytes(dek) - entry.err = err - close(entry.done) - delete(c.inflight, keyRef) + if !ok { + return } - c.inflightMu.Unlock() + entry.err = err + close(entry.done) + delete(c.inflight, keyRef) } func copyBytes(b []byte) []byte { diff --git a/encryption/cache_test.go b/encryption/cache_test.go index 4de3f05..7102b43 100644 --- a/encryption/cache_test.go +++ b/encryption/cache_test.go @@ -1,6 +1,8 @@ package encryption import ( + "context" + "errors" "sync" "testing" "time" @@ -190,59 +192,6 @@ func TestDEKCache_DeleteMissing(t *testing.T) { c.delete("nonexistent") } -func TestDEKCache_Clear(t *testing.T) { - c := newDEKCache(10, time.Minute) - - dek1 := []byte("key1key1key1key1key1key1key1key1") - dek2 := []byte("key2key2key2key2key2key2key2key2") - c.put("ref1", dek1) - c.put("ref2", dek2) - - // Grab internal slice references. - c.mu.Lock() - internal1 := c.entries["ref1"].dek - internal2 := c.entries["ref2"].dek - c.mu.Unlock() - - c.clear() - - if _, ok := c.get("ref1"); ok { - t.Fatal("expected miss after clear") - } - if _, ok := c.get("ref2"); ok { - t.Fatal("expected miss after clear") - } - - for _, b := range internal1 { - if b != 0 { - t.Fatal("cleared DEK 1 should be zeroed") - } - } - for _, b := range internal2 { - if b != 0 { - t.Fatal("cleared DEK 2 should be zeroed") - } - } -} - -func TestDEKCache_ClearThenReuse(t *testing.T) { - c := newDEKCache(10, time.Minute) - dek := []byte("0123456789abcdef0123456789abcdef") - - c.put("ref1", dek) - c.clear() - - // Cache should work normally after clear. - c.put("ref1", dek) - got, ok := c.get("ref1") - if !ok { - t.Fatal("expected cache hit after clear + put") - } - if string(got) != string(dek) { - t.Fatalf("got %x, want %x", got, dek) - } -} - func TestDEKCache_Singleflight(t *testing.T) { c := newDEKCache(10, time.Minute) @@ -268,13 +217,17 @@ func TestDEKCache_Singleflight(t *testing.T) { t.Error("subsequent caller should not start") return } - results[idx], errs[idx] = wait() + if errs[idx] = wait(context.Background()); errs[idx] != nil { + return + } + results[idx], _ = c.get("ref1") }(i) } - // Simulate fetch completing. + // Simulate the winner completing: populate the cache, then release waiters. time.Sleep(10 * time.Millisecond) // let goroutines reach wait() - c.finish("ref1", dek, nil) + c.put("ref1", dek) + c.finish("ref1", nil) wg.Wait() @@ -294,6 +247,37 @@ func TestDEKCache_Singleflight(t *testing.T) { } } +func TestDEKCache_SingleflightContextCancel(t *testing.T) { + c := newDEKCache(10, time.Minute) + + started, _ := c.waitOrStart("ref1") + if !started { + t.Fatal("first caller should start") + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + _, wait := c.waitOrStart("ref1") + done <- wait(ctx) + }() + + time.Sleep(10 * time.Millisecond) + cancel() + + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("got %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("waiter ignored context cancellation") + } + + // The starter can still finish without affecting the cancelled waiter. + c.finish("ref1", nil) +} + func TestDEKCache_SingleflightError(t *testing.T) { c := newDEKCache(10, time.Minute) fetchErr := &testError{msg: "kms failed"} @@ -311,12 +295,12 @@ func TestDEKCache_SingleflightError(t *testing.T) { go func(idx int) { defer wg.Done() _, wait := c.waitOrStart("ref1") - _, waiterErrs[idx] = wait() + waiterErrs[idx] = wait(context.Background()) }(i) } time.Sleep(10 * time.Millisecond) - c.finish("ref1", nil, fetchErr) + c.finish("ref1", fetchErr) wg.Wait() @@ -333,9 +317,6 @@ func TestDEKCache_SingleflightError(t *testing.T) { func TestDEKCache_SingleflightIndependentKeys(t *testing.T) { c := newDEKCache(10, time.Minute) - dek1 := []byte("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") - dek2 := []byte("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb") - // Start fetch for ref1. started1, _ := c.waitOrStart("ref1") if !started1 { @@ -348,8 +329,8 @@ func TestDEKCache_SingleflightIndependentKeys(t *testing.T) { t.Fatal("first caller for ref2 should start independently") } - c.finish("ref1", dek1, nil) - c.finish("ref2", dek2, nil) + c.finish("ref1", nil) + c.finish("ref2", nil) } func TestDEKCache_ConcurrentAccess(t *testing.T) { diff --git a/encryption/pool.go b/encryption/pool.go index 204ccbf..b2bf8db 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -119,16 +119,21 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext p.cache.put(key.KeyRef, privateKey) } } else { - // Existing key — try cache before KMS. + // The row is an unverified read and its KeyRef selects the DEK this + // data is encrypted under, so the attestation is checked before the + // cache is consulted — a hit must not be able to skip it. Verification + // is local (COSE + cert chain), so the KMS round-trips are still what + // the cache saves. + if err := p.VerifyKey(ctx, att, key); err != nil { + return "", nil, fmt.Errorf("verify key: %w", err) + } + if p.cache != nil { if dek, ok := p.cache.get(key.KeyRef); ok { privateKey = dek } } if privateKey == nil { - if err := p.VerifyKey(ctx, att, key); err != nil { - return "", nil, fmt.Errorf("verify key: %w", err) - } privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) if err != nil { return "", nil, fmt.Errorf("combine shares: %w", err) @@ -174,20 +179,9 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, fmt.Errorf("decode ciphertext: %w", err) } - // Try cache. - var privateKey []byte - if p.cache != nil { - if dek, ok := p.cache.get(keyRef); ok { - privateKey = dek - } - } - - // Cache miss — full fetch with singleflight dedup. - if privateKey == nil { - privateKey, err = p.fetchDEK(ctx, att, keyRef) - if err != nil { - return nil, err - } + privateKey, err := p.fetchDEK(ctx, att, keyRef) + if err != nil { + return nil, err } // Decrypt data. @@ -205,19 +199,55 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return decrypted, nil } -// fetchDEK retrieves a DEK through the full path: DynamoDB lookup, attestation -// verification, KMS share decryption, and Shamir combine. It uses singleflight -// to deduplicate concurrent fetches for the same keyRef, and populates the cache. -func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, err error) { - // Singleflight: if another goroutine is already fetching this keyRef, wait. - if p.cache != nil { - started, wait := p.cache.waitOrStart(keyRef) - if !started { - return wait() +// fetchDEK returns the DEK for keyRef, serving it from the cache when possible +// and otherwise loading it through loadDEK. Concurrent misses on the same keyRef +// are deduplicated so only one of them performs the KMS round-trips. +func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) ([]byte, error) { + if p.cache == nil { + return p.loadDEK(ctx, att, keyRef) + } + + if dek, ok := p.cache.get(keyRef); ok { + return dek, nil + } + + started, wait := p.cache.waitOrStart(keyRef) + if !started { + if err := wait(ctx); err != nil { + return nil, err + } + if dek, ok := p.cache.get(keyRef); ok { + return dek, nil } - defer func() { p.cache.finish(keyRef, privateKey, err) }() + // Evicted between the winner's put and this read; load it ourselves + // rather than waiting again. + return p.loadDEK(ctx, att, keyRef) } + var ( + dek []byte + err error + ) + // Deferred so a panic in loadDEK still releases the waiters, and ordered + // after the put so the cache is populated by the time they re-read it. + defer func() { p.cache.finish(keyRef, err) }() + + dek, err = p.loadDEK(ctx, att, keyRef) + if err == nil { + p.cache.put(keyRef, dek) + } + return dek, err +} + +// loadDEK recovers a DEK through the full path: DynamoDB lookup, attestation +// verification, KMS share decryption, and Shamir combine. +func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, err error) { + ctx, span := tracing.Trace(ctx, "encryption.Pool.loadDEK", tracing.WithAnnotation("key_ref", keyRef)) + defer func() { + span.RecordError(err) + span.End() + }() + key, found, err := p.keysTable.GetLatestByKeyRef(ctx, keyRef, false) if err != nil { return nil, fmt.Errorf("get latest key: %w", err) @@ -229,6 +259,11 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st return nil, fmt.Errorf("verify key: %w", err) } + span.SetAnnotation("generation", strconv.Itoa(key.Generation)) + if key.KeyIndex != nil { + span.SetAnnotation("key_index", strconv.Itoa(*key.KeyIndex)) + } + config, err := p.getConfig(key.Generation) if err != nil { return nil, fmt.Errorf("get config: %w", err) @@ -242,12 +277,8 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st return nil, fmt.Errorf("combine shares: %w", err) } - if p.cache != nil { - p.cache.put(keyRef, privateKey) - } - - // Trigger migration if needed. Migration is synchronous but non-fatal: - // failure is logged and does not affect the returned DEK. + // Migration is synchronous but non-fatal: failure is logged, the DEK is + // still returned, and the next load past the cache TTL retries it. if p.keyNeedsMigration(key) { if err := p.migrateKey(ctx, att, key, privateKey); err != nil { p.logger.ErrorContext(ctx, "migrating key failed", "error", err, "key_ref", key.KeyRef, "generation", key.Generation, "key_index", key.KeyIndex) @@ -351,6 +382,9 @@ func (p *Pool) CleanupUnusedKeys(ctx context.Context) (deleted int, err error) { if err := p.keysTable.Delete(ctx, key.KeyRef, key.Generation); err != nil { return deleted, fmt.Errorf("delete cipher key by ref %q: %w", key.KeyRef, err) } + if p.cache != nil { + p.cache.delete(key.KeyRef) + } deleted++ } } diff --git a/encryption/pool_test.go b/encryption/pool_test.go index 6ac4d40..5ca4c13 100644 --- a/encryption/pool_test.go +++ b/encryption/pool_test.go @@ -52,7 +52,6 @@ var ( legacyCiphertext55_v2 = []byte("v2.QkJCQkJCQkJCQkJC9X30rkaY8XO5h_ujMLmLXiPzlYA") ) - func TestPool_Encrypt(t *testing.T) { block, _ := pem.Decode([]byte(dummyPrivKey)) privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) @@ -1598,3 +1597,76 @@ func TestPool_MultipleKeyRefs(t *testing.T) { } func intPtr(v int) *int { return &v } + +// A cache hit must not let an unverified cipher key row through: the row's +// KeyRef selects which DEK encrypts the data, so its attestation is checked on +// every Encrypt, hit or miss. +func TestPool_EncryptVerifiesKeyOnCacheHit(t *testing.T) { + block, _ := pem.Decode([]byte(dummyPrivKey)) + privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) + require.NoError(t, err) + + newPool := func(opts ...encryption.PoolOption) (*encryption.Pool, *enclave.Attestation) { + remoteKey1 := &MockRemoteKey{} + remoteKey2 := &MockRemoteKey{} + keysTable := &MockKeysTable{} + + enc, err := enclave.New(context.Background(), enclave.DummyProvider(&constantReader{value: 0x42}), &MockKMS{}, privKey) + require.NoError(t, err) + + configs := []*encryption.Config{{ + PoolSize: 10, + Threshold: 2, + RemoteKeys: map[string]encryption.RemoteKey{ + "remoteKey1": remoteKey1, + "remoteKey2": remoteKey2, + }, + }} + + att, err := enc.GetAttestation(context.Background(), nil, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = att.Close() }) + + cipherKey, privateKey := newCipherKey(t, enc) + shares, err := shamir.Split(privateKey, 2, 2) + require.NoError(t, err) + + // Same KeyRef, attacker-swapped shares. The stored attestation commits + // to the original content, so VerifyKey must reject it. + tampered := *cipherKey + tampered.EncryptedShares = map[string]string{ + "remoteKey1": "attackerShare1", + "remoteKey2": "attackerShare2", + } + + remoteKey1.On("Decrypt", mock.Anything, mock.Anything, "encryptedShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, mock.Anything, "encryptedShare2").Return(shares[1], nil) + remoteKey1.On("Decrypt", mock.Anything, mock.Anything, "attackerShare1").Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, mock.Anything, "attackerShare2").Return(shares[1], nil) + + keysTable.On("Get", mock.Anything, 0, 4).Return(cipherKey, true, nil).Once() + keysTable.On("Get", mock.Anything, 0, 4).Return(&tampered, true, nil) + + return encryption.NewPool(enc, configs, keysTable, nil, nil, opts...), att + } + + t.Run("without cache", func(t *testing.T) { + pool, att := newPool() + + _, _, err := pool.Encrypt(context.Background(), att, []byte("first"), []byte("aad")) + require.NoError(t, err) + + _, _, err = pool.Encrypt(context.Background(), att, []byte("second"), []byte("aad")) + require.ErrorContains(t, err, "verify key") + }) + + t.Run("with cache", func(t *testing.T) { + pool, att := newPool(encryption.WithCache(encryption.CacheConfig{MaxSize: 10, TTL: time.Minute})) + + _, _, err := pool.Encrypt(context.Background(), att, []byte("first"), []byte("aad")) + require.NoError(t, err) + + _, _, err = pool.Encrypt(context.Background(), att, []byte("second"), []byte("aad")) + require.ErrorContains(t, err, "verify key") + }) +} From 968420ff007e7af3f3063a7c81065f5d752e82c1 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 14:02:43 +0200 Subject: [PATCH 03/12] feat(encryption): annotate DEK cache outcome on the operation span Records hit/miss/coalesced on the Encrypt or Decrypt span, so a slow operation shows whether it paid for a KMS round-trip. Annotated on the caller's span rather than loadDEK's child, which only exists on a miss. Co-Authored-By: Claude Opus 5 --- encryption/pool.go | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/encryption/pool.go b/encryption/pool.go index b2bf8db..ab060aa 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -45,6 +45,12 @@ type Pool struct { cache *dekCache // nil when caching disabled } +// annotationDEKCache records how a DEK was obtained on the span of the +// operation that needed it: served from the cache, loaded through DynamoDB and +// KMS, or coalesced onto another goroutine's in-flight load. It is what tells +// you whether a slow Encrypt/Decrypt paid for a KMS round-trip. +const annotationDEKCache = "dek_cache" + // PoolOption configures optional Pool behavior. type PoolOption func(*Pool) @@ -131,9 +137,13 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext if p.cache != nil { if dek, ok := p.cache.get(key.KeyRef); ok { privateKey = dek + span.SetAnnotation(annotationDEKCache, "hit") } } if privateKey == nil { + if p.cache != nil { + span.SetAnnotation(annotationDEKCache, "miss") + } privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) if err != nil { return "", nil, fmt.Errorf("combine shares: %w", err) @@ -207,12 +217,18 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st return p.loadDEK(ctx, att, keyRef) } + // The caller's span, not loadDEK's child: the annotation belongs on the + // operation whose latency it explains. + span := tracing.GetSpan(ctx) + if dek, ok := p.cache.get(keyRef); ok { + span.SetAnnotation(annotationDEKCache, "hit") return dek, nil } started, wait := p.cache.waitOrStart(keyRef) if !started { + span.SetAnnotation(annotationDEKCache, "coalesced") if err := wait(ctx); err != nil { return nil, err } @@ -221,9 +237,12 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } // Evicted between the winner's put and this read; load it ourselves // rather than waiting again. + span.SetAnnotation(annotationDEKCache, "miss") return p.loadDEK(ctx, att, keyRef) } + span.SetAnnotation(annotationDEKCache, "miss") + var ( dek []byte err error From 0a06aebd04e76aed0874e6dfaf8e9b35bbde112b Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 15:54:56 +0200 Subject: [PATCH 04/12] refactor(encryption): use x/sync singleflight for DEK loads Replaces the hand-rolled inflight map with singleflight.Group. DoChan keeps the context cancellation, and Result.Shared is the coalesced annotation, so fetchDEK loses the wait-then-re-read-then-load-anyway path. copyBytes was slices.Clone. x/sync was already in the module graph; this only promotes it to direct. Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 81 ++++----------------- encryption/cache_test.go | 149 --------------------------------------- encryption/pool.go | 47 ++++++------ go.mod | 2 +- 4 files changed, 34 insertions(+), 245 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index f6ee93c..a19612d 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -2,9 +2,11 @@ package encryption import ( "container/list" - "context" + "slices" "sync" "time" + + "golang.org/x/sync/singleflight" ) // CacheConfig configures the optional DEK (data encryption key) cache. @@ -33,26 +35,16 @@ type dekCache struct { maxSize int ttl time.Duration - inflightMu sync.Mutex - inflight map[string]*inflightEntry -} - -// inflightEntry coordinates singleflight deduplication for concurrent -// cache misses on the same keyRef. It deliberately carries no key material: -// waiters read the result from the cache, so there is no second copy of a DEK -// to keep track of and zero. -type inflightEntry struct { - done chan struct{} - err error + // group collapses concurrent misses on the same keyRef into one load. + group singleflight.Group } func newDEKCache(maxSize int, ttl time.Duration) *dekCache { return &dekCache{ - entries: make(map[string]*dekCacheEntry), - order: list.New(), - maxSize: maxSize, - ttl: ttl, - inflight: make(map[string]*inflightEntry), + entries: make(map[string]*dekCacheEntry), + order: list.New(), + maxSize: maxSize, + ttl: ttl, } } @@ -72,7 +64,7 @@ func (c *dekCache) get(keyRef string) ([]byte, bool) { } c.order.MoveToFront(entry.element) - return copyBytes(entry.dek), true + return slices.Clone(entry.dek), true } // put stores a copy of dek in the cache, evicting the LRU entry if full. @@ -83,14 +75,14 @@ func (c *dekCache) put(keyRef string, dek []byte) { if entry, ok := c.entries[keyRef]; ok { // Update existing entry. clear(entry.dek) - entry.dek = copyBytes(dek) + entry.dek = slices.Clone(dek) entry.expiresAt = time.Now().Add(c.ttl) c.order.MoveToFront(entry.element) return } entry := &dekCacheEntry{ - dek: copyBytes(dek), + dek: slices.Clone(dek), keyRef: keyRef, expiresAt: time.Now().Add(c.ttl), } @@ -121,52 +113,3 @@ func (c *dekCache) evictLocked(entry *dekCacheEntry) { c.order.Remove(entry.element) delete(c.entries, entry.keyRef) } - -// waitOrStart implements singleflight deduplication. If another goroutine is -// already fetching the DEK for keyRef, started=false and wait blocks until that -// fetch finishes, reporting its error; the caller then reads the DEK from the -// cache. Otherwise started=true and the caller must call finish. -func (c *dekCache) waitOrStart(keyRef string) (started bool, wait func(ctx context.Context) error) { - c.inflightMu.Lock() - - if entry, ok := c.inflight[keyRef]; ok { - c.inflightMu.Unlock() - return false, func(ctx context.Context) error { - select { - case <-entry.done: - return entry.err - case <-ctx.Done(): - return ctx.Err() - } - } - } - - entry := &inflightEntry{done: make(chan struct{})} - c.inflight[keyRef] = entry - c.inflightMu.Unlock() - - return true, nil -} - -// finish signals all waiters for keyRef with the fetch outcome. -func (c *dekCache) finish(keyRef string, err error) { - c.inflightMu.Lock() - defer c.inflightMu.Unlock() - - entry, ok := c.inflight[keyRef] - if !ok { - return - } - entry.err = err - close(entry.done) - delete(c.inflight, keyRef) -} - -func copyBytes(b []byte) []byte { - if b == nil { - return nil - } - cp := make([]byte, len(b)) - copy(cp, b) - return cp -} diff --git a/encryption/cache_test.go b/encryption/cache_test.go index 7102b43..732481f 100644 --- a/encryption/cache_test.go +++ b/encryption/cache_test.go @@ -1,8 +1,6 @@ package encryption import ( - "context" - "errors" "sync" "testing" "time" @@ -192,147 +190,6 @@ func TestDEKCache_DeleteMissing(t *testing.T) { c.delete("nonexistent") } -func TestDEKCache_Singleflight(t *testing.T) { - c := newDEKCache(10, time.Minute) - - dek := []byte("0123456789abcdef0123456789abcdef") - - // First caller starts the fetch. - started, _ := c.waitOrStart("ref1") - if !started { - t.Fatal("first caller should start") - } - - // Second and third callers should wait. - var wg sync.WaitGroup - results := make([][]byte, 2) - errs := make([]error, 2) - - for i := 0; i < 2; i++ { - wg.Add(1) - go func(idx int) { - defer wg.Done() - s, wait := c.waitOrStart("ref1") - if s { - t.Error("subsequent caller should not start") - return - } - if errs[idx] = wait(context.Background()); errs[idx] != nil { - return - } - results[idx], _ = c.get("ref1") - }(i) - } - - // Simulate the winner completing: populate the cache, then release waiters. - time.Sleep(10 * time.Millisecond) // let goroutines reach wait() - c.put("ref1", dek) - c.finish("ref1", nil) - - wg.Wait() - - for i := 0; i < 2; i++ { - if errs[i] != nil { - t.Fatalf("waiter %d got error: %v", i, errs[i]) - } - if string(results[i]) != string(dek) { - t.Fatalf("waiter %d got wrong dek", i) - } - } - - // Each waiter should have received an independent copy. - results[0][0] = 0xFF - if results[1][0] == 0xFF { - t.Fatal("waiters should receive independent copies") - } -} - -func TestDEKCache_SingleflightContextCancel(t *testing.T) { - c := newDEKCache(10, time.Minute) - - started, _ := c.waitOrStart("ref1") - if !started { - t.Fatal("first caller should start") - } - - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan error, 1) - go func() { - _, wait := c.waitOrStart("ref1") - done <- wait(ctx) - }() - - time.Sleep(10 * time.Millisecond) - cancel() - - select { - case err := <-done: - if !errors.Is(err, context.Canceled) { - t.Fatalf("got %v, want context.Canceled", err) - } - case <-time.After(time.Second): - t.Fatal("waiter ignored context cancellation") - } - - // The starter can still finish without affecting the cancelled waiter. - c.finish("ref1", nil) -} - -func TestDEKCache_SingleflightError(t *testing.T) { - c := newDEKCache(10, time.Minute) - fetchErr := &testError{msg: "kms failed"} - - started, _ := c.waitOrStart("ref1") - if !started { - t.Fatal("first caller should start") - } - - var wg sync.WaitGroup - waiterErrs := make([]error, 2) - - for i := 0; i < 2; i++ { - wg.Add(1) - go func(idx int) { - defer wg.Done() - _, wait := c.waitOrStart("ref1") - waiterErrs[idx] = wait(context.Background()) - }(i) - } - - time.Sleep(10 * time.Millisecond) - c.finish("ref1", fetchErr) - - wg.Wait() - - for i := 0; i < 2; i++ { - if waiterErrs[i] == nil { - t.Fatalf("waiter %d should have received error", i) - } - if waiterErrs[i].Error() != "kms failed" { - t.Fatalf("waiter %d got error %q, want %q", i, waiterErrs[i].Error(), "kms failed") - } - } -} - -func TestDEKCache_SingleflightIndependentKeys(t *testing.T) { - c := newDEKCache(10, time.Minute) - - // Start fetch for ref1. - started1, _ := c.waitOrStart("ref1") - if !started1 { - t.Fatal("first caller for ref1 should start") - } - - // Start fetch for ref2 — should NOT be blocked by ref1. - started2, _ := c.waitOrStart("ref2") - if !started2 { - t.Fatal("first caller for ref2 should start independently") - } - - c.finish("ref1", nil) - c.finish("ref2", nil) -} - func TestDEKCache_ConcurrentAccess(t *testing.T) { c := newDEKCache(10, time.Minute) dek := []byte("0123456789abcdef0123456789abcdef") @@ -357,9 +214,3 @@ func TestDEKCache_ConcurrentAccess(t *testing.T) { } wg.Wait() } - -type testError struct { - msg string -} - -func (e *testError) Error() string { return e.msg } diff --git a/encryption/pool.go b/encryption/pool.go index ab060aa..988ab3f 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "log/slog" + "slices" "strconv" "time" @@ -226,36 +227,30 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st return dek, nil } - started, wait := p.cache.waitOrStart(keyRef) - if !started { - span.SetAnnotation(annotationDEKCache, "coalesced") - if err := wait(ctx); err != nil { + loaded := p.cache.group.DoChan(keyRef, func() (any, error) { + dek, err := p.loadDEK(ctx, att, keyRef) + if err != nil { return nil, err } - if dek, ok := p.cache.get(keyRef); ok { - return dek, nil - } - // Evicted between the winner's put and this read; load it ourselves - // rather than waiting again. - span.SetAnnotation(annotationDEKCache, "miss") - return p.loadDEK(ctx, att, keyRef) - } - - span.SetAnnotation(annotationDEKCache, "miss") - - var ( - dek []byte - err error - ) - // Deferred so a panic in loadDEK still releases the waiters, and ordered - // after the put so the cache is populated by the time they re-read it. - defer func() { p.cache.finish(keyRef, err) }() - - dek, err = p.loadDEK(ctx, att, keyRef) - if err == nil { p.cache.put(keyRef, dek) + return dek, nil + }) + + select { + case res := <-loaded: + if res.Err != nil { + return nil, res.Err + } + if res.Shared { + span.SetAnnotation(annotationDEKCache, "coalesced") + } else { + span.SetAnnotation(annotationDEKCache, "miss") + } + // Shared between everyone who coalesced onto this load. + return slices.Clone(res.Val.([]byte)), nil + case <-ctx.Done(): + return nil, ctx.Err() } - return dek, err } // loadDEK recovers a DEK through the full path: DynamoDB lookup, attestation diff --git a/go.mod b/go.mod index d517f51..758c100 100644 --- a/go.mod +++ b/go.mod @@ -13,6 +13,7 @@ require ( github.com/fxamacker/cbor/v2 v2.9.0 github.com/go-chi/chi/v5 v5.2.5 github.com/stretchr/testify v1.11.1 + golang.org/x/sync v0.20.0 golang.org/x/tools v0.43.0 ) @@ -30,6 +31,5 @@ require ( github.com/stretchr/objx v0.5.2 // indirect github.com/x448/float16 v0.8.4 // indirect golang.org/x/mod v0.34.0 // indirect - golang.org/x/sync v0.20.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) From 52c29ce3410c1a0ee9d7799fbfcedfd162608ec7 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 16:06:41 +0200 Subject: [PATCH 05/12] refactor(encryption): back the DEK cache with expirable.LRU Replaces the hand-rolled LRU, TTL and eviction bookkeeping with hashicorp/golang-lru's expirable.LRU, which waas already has in its module graph at the same version. Evicted keys are no longer zeroed. The LRU releases its lock before the caller finishes copying a value and its expiry sweeper runs on its own goroutine, so clearing an evicted slice can hand out a half-cleared key and fail a decryption. The enclave already keeps unscrubbed copies from Shamir recombination and the AES key schedule, and its memory is neither swappable nor readable from the parent instance. Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 97 +++++++---------------------------- encryption/cache_test.go | 107 +++++---------------------------------- go.mod | 1 + go.sum | 2 + 4 files changed, 34 insertions(+), 173 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index a19612d..c296ad0 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -1,11 +1,10 @@ package encryption import ( - "container/list" "slices" - "sync" "time" + "github.com/hashicorp/golang-lru/v2/expirable" "golang.org/x/sync/singleflight" ) @@ -18,98 +17,40 @@ type CacheConfig struct { TTL time.Duration } -// dekCacheEntry holds a cached DEK and its LRU/TTL metadata. -type dekCacheEntry struct { - dek []byte // 32-byte AES-256 key (owned copy) - keyRef string // for reverse lookup from LRU list element - expiresAt time.Time - element *list.Element // back-pointer into LRU list -} - -// dekCache is a thread-safe LRU cache for decrypted data encryption keys. -// It zeroes key material on every eviction path (TTL, LRU, delete, clear). +// dekCache is an LRU of decrypted data encryption keys. Entries expire TTL +// after they are stored, regardless of use. +// +// Evicted keys are left for the garbage collector rather than zeroed. Zeroing +// would mean writing to a slice a concurrent reader may still be copying — the +// LRU releases its lock before the caller is done with the value, and the +// expiry sweeper runs on its own goroutine — which risks handing out a +// half-cleared key and failing a decryption. It buys little in return: the +// enclave already holds unscrubbed copies of this key from Shamir recombination +// and the AES key schedule, and its memory is neither swappable nor readable +// from the parent instance. type dekCache struct { - mu sync.Mutex - entries map[string]*dekCacheEntry - order *list.List // front = most recently used - maxSize int - ttl time.Duration - + lru *expirable.LRU[string, []byte] // group collapses concurrent misses on the same keyRef into one load. group singleflight.Group } func newDEKCache(maxSize int, ttl time.Duration) *dekCache { - return &dekCache{ - entries: make(map[string]*dekCacheEntry), - order: list.New(), - maxSize: maxSize, - ttl: ttl, - } + return &dekCache{lru: expirable.NewLRU[string, []byte](maxSize, nil, ttl)} } -// get returns a copy of the cached DEK for keyRef, or ok=false on miss/expiry. +// The cache and the caller each own their copy. func (c *dekCache) get(keyRef string) ([]byte, bool) { - c.mu.Lock() - defer c.mu.Unlock() - - entry, ok := c.entries[keyRef] + dek, ok := c.lru.Get(keyRef) if !ok { return nil, false } - - if time.Now().After(entry.expiresAt) { - c.evictLocked(entry) - return nil, false - } - - c.order.MoveToFront(entry.element) - return slices.Clone(entry.dek), true + return slices.Clone(dek), true } -// put stores a copy of dek in the cache, evicting the LRU entry if full. func (c *dekCache) put(keyRef string, dek []byte) { - c.mu.Lock() - defer c.mu.Unlock() - - if entry, ok := c.entries[keyRef]; ok { - // Update existing entry. - clear(entry.dek) - entry.dek = slices.Clone(dek) - entry.expiresAt = time.Now().Add(c.ttl) - c.order.MoveToFront(entry.element) - return - } - - entry := &dekCacheEntry{ - dek: slices.Clone(dek), - keyRef: keyRef, - expiresAt: time.Now().Add(c.ttl), - } - entry.element = c.order.PushFront(entry) - c.entries[keyRef] = entry - - if len(c.entries) > c.maxSize { - back := c.order.Back() - if back != nil { - c.evictLocked(back.Value.(*dekCacheEntry)) - } - } + c.lru.Add(keyRef, slices.Clone(dek)) } -// delete removes and zeroes a specific entry. Called by RotateKey. func (c *dekCache) delete(keyRef string) { - c.mu.Lock() - defer c.mu.Unlock() - - if entry, ok := c.entries[keyRef]; ok { - c.evictLocked(entry) - } -} - -// evictLocked removes an entry, zeroing its DEK. Caller must hold c.mu. -func (c *dekCache) evictLocked(entry *dekCacheEntry) { - clear(entry.dek) - c.order.Remove(entry.element) - delete(c.entries, entry.keyRef) + c.lru.Remove(keyRef) } diff --git a/encryption/cache_test.go b/encryption/cache_test.go index 732481f..b809087 100644 --- a/encryption/cache_test.go +++ b/encryption/cache_test.go @@ -22,20 +22,14 @@ func TestDEKCache_GetPut(t *testing.T) { // Returned slice must be an independent copy. got[0] = 0xFF - got2, ok := c.get("ref1") - if !ok { - t.Fatal("expected cache hit") - } + got2, _ := c.get("ref1") if got2[0] == 0xFF { t.Fatal("cache returned same underlying slice, expected independent copy") } // Stored slice must be an independent copy of the input. dek[0] = 0xAA - got3, ok := c.get("ref1") - if !ok { - t.Fatal("expected cache hit") - } + got3, _ := c.get("ref1") if got3[0] == 0xAA { t.Fatal("cache stored same underlying slice as input, expected independent copy") } @@ -44,44 +38,20 @@ func TestDEKCache_GetPut(t *testing.T) { func TestDEKCache_Miss(t *testing.T) { c := newDEKCache(10, time.Minute) - _, ok := c.get("unknown") - if ok { + if _, ok := c.get("unknown"); ok { t.Fatal("expected cache miss for unknown key") } } func TestDEKCache_TTLExpiry(t *testing.T) { c := newDEKCache(10, time.Millisecond) - dek := []byte("0123456789abcdef0123456789abcdef") - - c.put("ref1", dek) - - // Grab internal slice before expiry. - c.mu.Lock() - internalSlice := c.entries["ref1"].dek - c.mu.Unlock() + c.put("ref1", []byte("0123456789abcdef0123456789abcdef")) time.Sleep(5 * time.Millisecond) - _, ok := c.get("ref1") - if ok { + if _, ok := c.get("ref1"); ok { t.Fatal("expected cache miss after TTL expiry") } - - // Expired entry should be removed from map. - c.mu.Lock() - _, exists := c.entries["ref1"] - c.mu.Unlock() - if exists { - t.Fatal("expired entry should have been removed from map") - } - - // Internal DEK should be zeroed. - for _, b := range internalSlice { - if b != 0 { - t.Fatal("expired DEK should have been zeroed") - } - } } func TestDEKCache_LRUEviction(t *testing.T) { @@ -90,15 +60,8 @@ func TestDEKCache_LRUEviction(t *testing.T) { c.put("ref1", []byte("key1key1key1key1key1key1key1key1")) c.put("ref2", []byte("key2key2key2key2key2key2key2key2")) - // Grab internal slice of ref2 before eviction. - c.mu.Lock() - internalRef2 := c.entries["ref2"].dek - c.mu.Unlock() - - // Access ref1 to make it more recent than ref2. + // Make ref1 more recent than ref2, so ref3 evicts ref2. _, _ = c.get("ref1") - - // Adding ref3 should evict ref2 (LRU). c.put("ref3", []byte("key3key3key3key3key3key3key3key3")) if _, ok := c.get("ref2"); ok { @@ -110,29 +73,13 @@ func TestDEKCache_LRUEviction(t *testing.T) { if _, ok := c.get("ref3"); !ok { t.Fatal("expected ref3 to still be cached") } - - // Verify evicted internal DEK was zeroed. - for _, b := range internalRef2 { - if b != 0 { - t.Fatal("evicted DEK should have been zeroed") - } - } } func TestDEKCache_PutUpdatesExisting(t *testing.T) { c := newDEKCache(10, time.Minute) - dek1 := []byte("old_key_old_key_old_key_old_key_") dek2 := []byte("new_key_new_key_new_key_new_key_") - c.put("ref1", dek1) - - // Grab internal slice and expiry before update. - c.mu.Lock() - internalSlice := c.entries["ref1"].dek - oldExpiry := c.entries["ref1"].expiresAt - c.mu.Unlock() - - time.Sleep(time.Millisecond) // ensure time advances + c.put("ref1", []byte("old_key_old_key_old_key_old_key_")) c.put("ref1", dek2) got, ok := c.get("ref1") @@ -142,51 +89,21 @@ func TestDEKCache_PutUpdatesExisting(t *testing.T) { if string(got) != string(dek2) { t.Fatalf("got %x, want %x", got, dek2) } - - // Old internal slice should be zeroed. - for _, b := range internalSlice { - if b != 0 { - t.Fatal("old DEK slice should have been zeroed") - } - } - - // TTL should be refreshed. - c.mu.Lock() - newExpiry := c.entries["ref1"].expiresAt - c.mu.Unlock() - if !newExpiry.After(oldExpiry) { - t.Fatal("put on existing key should refresh TTL") - } } func TestDEKCache_Delete(t *testing.T) { c := newDEKCache(10, time.Minute) - dek := []byte("0123456789abcdef0123456789abcdef") - - c.put("ref1", dek) - - // Grab internal slice reference. - c.mu.Lock() - internalSlice := c.entries["ref1"].dek - c.mu.Unlock() + c.put("ref1", []byte("0123456789abcdef0123456789abcdef")) c.delete("ref1") if _, ok := c.get("ref1"); ok { t.Fatal("expected cache miss after delete") } - - // Verify zeroed. - for _, b := range internalSlice { - if b != 0 { - t.Fatal("deleted DEK should have been zeroed") - } - } } func TestDEKCache_DeleteMissing(t *testing.T) { c := newDEKCache(10, time.Minute) - // Should not panic. c.delete("nonexistent") } @@ -197,14 +114,14 @@ func TestDEKCache_ConcurrentAccess(t *testing.T) { var wg sync.WaitGroup for i := 0; i < 100; i++ { wg.Add(3) - go func(n int) { + go func() { defer wg.Done() c.put("ref1", dek) - }(i) - go func(n int) { + }() + go func() { defer wg.Done() c.get("ref1") - }(i) + }() go func(n int) { defer wg.Done() if n%10 == 0 { diff --git a/go.mod b/go.mod index 758c100..b77c952 100644 --- a/go.mod +++ b/go.mod @@ -12,6 +12,7 @@ require ( github.com/aws/aws-sdk-go-v2/service/kms v1.46.2 github.com/fxamacker/cbor/v2 v2.9.0 github.com/go-chi/chi/v5 v5.2.5 + github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/stretchr/testify v1.11.1 golang.org/x/sync v0.20.0 golang.org/x/tools v0.43.0 diff --git a/go.sum b/go.sum index 7d6833e..4b5b934 100644 --- a/go.sum +++ b/go.sum @@ -36,6 +36,8 @@ github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= From 67fe8d474be6a563feca6d693fee1070e70c90f6 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 17:19:15 +0200 Subject: [PATCH 06/12] fix(encryption): don't cache a DEK whose migration failed Caching it suppressed the retry until the entry expired; previously every decrypt re-attempted. Also corrects Decrypt's doc comment, which still claimed verification and migration happen on every call. Co-Authored-By: Claude Opus 5 --- encryption/pool.go | 39 +++++++++++++++++++++++---------------- 1 file changed, 23 insertions(+), 16 deletions(-) diff --git a/encryption/pool.go b/encryption/pool.go index 988ab3f..01a62a6 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -176,8 +176,9 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext // Decrypt decrypts the ciphertext using the latest cipher key from the Pool referenced by the keyRef. // -// The key is verified against the attestation and migrated to the current generation if needed. -// If a DEK cache is configured, cached keys bypass DynamoDB and KMS on hit. +// The key is verified against the attestation and migrated to the current +// generation if needed. A cached DEK skips all of that: it was verified when it +// was loaded, and its key ref never maps to different key material. func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef string, ciphertext []byte, additionalData []byte) (plaintext []byte, err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.Decrypt", tracing.WithAnnotation("key_ref", keyRef)) defer func() { @@ -215,7 +216,8 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str // are deduplicated so only one of them performs the KMS round-trips. func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) ([]byte, error) { if p.cache == nil { - return p.loadDEK(ctx, att, keyRef) + dek, _, err := p.loadDEK(ctx, att, keyRef) + return dek, err } // The caller's span, not loadDEK's child: the annotation belongs on the @@ -228,11 +230,13 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } loaded := p.cache.group.DoChan(keyRef, func() (any, error) { - dek, err := p.loadDEK(ctx, att, keyRef) + dek, cacheable, err := p.loadDEK(ctx, att, keyRef) if err != nil { return nil, err } - p.cache.put(keyRef, dek) + if cacheable { + p.cache.put(keyRef, dek) + } return dek, nil }) @@ -254,8 +258,10 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } // loadDEK recovers a DEK through the full path: DynamoDB lookup, attestation -// verification, KMS share decryption, and Shamir combine. -func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, err error) { +// verification, KMS share decryption, and Shamir combine. It reports whether the +// result may be cached; a key whose migration failed must not be, or the retry +// is suppressed until the entry expires. +func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, cacheable bool, err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.loadDEK", tracing.WithAnnotation("key_ref", keyRef)) defer func() { span.RecordError(err) @@ -264,13 +270,13 @@ func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef str key, found, err := p.keysTable.GetLatestByKeyRef(ctx, keyRef, false) if err != nil { - return nil, fmt.Errorf("get latest key: %w", err) + return nil, false, fmt.Errorf("get latest key: %w", err) } if !found { - return nil, fmt.Errorf("key not found") + return nil, false, fmt.Errorf("key not found") } if err := p.VerifyKey(ctx, att, key); err != nil { - return nil, fmt.Errorf("verify key: %w", err) + return nil, false, fmt.Errorf("verify key: %w", err) } span.SetAnnotation("generation", strconv.Itoa(key.Generation)) @@ -280,26 +286,27 @@ func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef str config, err := p.getConfig(key.Generation) if err != nil { - return nil, fmt.Errorf("get config: %w", err) + return nil, false, fmt.Errorf("get config: %w", err) } if !config.areSharesValid(key.EncryptedShares) { - return nil, fmt.Errorf("shares are invalid") + return nil, false, fmt.Errorf("shares are invalid") } privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) if err != nil { - return nil, fmt.Errorf("combine shares: %w", err) + return nil, false, fmt.Errorf("combine shares: %w", err) } - // Migration is synchronous but non-fatal: failure is logged, the DEK is - // still returned, and the next load past the cache TTL retries it. + // Migration is non-fatal: the DEK is returned either way, but a failed + // migration leaves the key uncacheable so the next decrypt retries it. if p.keyNeedsMigration(key) { if err := p.migrateKey(ctx, att, key, privateKey); err != nil { p.logger.ErrorContext(ctx, "migrating key failed", "error", err, "key_ref", key.KeyRef, "generation", key.Generation, "key_index", key.KeyIndex) + return privateKey, false, nil } } - return privateKey, nil + return privateKey, true, nil } // RotateKey marks a key as inactive by setting its KeyIndex to a negative value. It won't be used for encrypting From 8a2a1daacb4e3e2c3f13009bf8737e8055bf1f62 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 17:49:41 +0200 Subject: [PATCH 07/12] style(encryption): trim DEK cache comments Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 20 +++++++------------- encryption/pool.go | 35 ++++++++++++----------------------- 2 files changed, 19 insertions(+), 36 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index c296ad0..835de47 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -11,23 +11,18 @@ import ( // CacheConfig configures the optional DEK (data encryption key) cache. // A zero-value config disables caching. type CacheConfig struct { - // MaxSize is the maximum number of DEKs to cache. Must be > 0 to enable caching. + // MaxSize is the maximum number of DEKs to cache. MaxSize int - // TTL is the time-to-live for each cache entry. Must be > 0 to enable caching. + // TTL is how long a cached DEK lives, counted from when it is stored. TTL time.Duration } -// dekCache is an LRU of decrypted data encryption keys. Entries expire TTL -// after they are stored, regardless of use. +// dekCache is an LRU of decrypted data encryption keys. // -// Evicted keys are left for the garbage collector rather than zeroed. Zeroing -// would mean writing to a slice a concurrent reader may still be copying — the -// LRU releases its lock before the caller is done with the value, and the -// expiry sweeper runs on its own goroutine — which risks handing out a -// half-cleared key and failing a decryption. It buys little in return: the -// enclave already holds unscrubbed copies of this key from Shamir recombination -// and the AES key schedule, and its memory is neither swappable nor readable -// from the parent instance. +// Evicted keys are not zeroed: the LRU releases its lock before the caller has +// finished copying a value, so clearing one risks handing out a half-cleared +// key. The enclave already holds unscrubbed copies from Shamir recombination +// and the AES key schedule. type dekCache struct { lru *expirable.LRU[string, []byte] // group collapses concurrent misses on the same keyRef into one load. @@ -38,7 +33,6 @@ func newDEKCache(maxSize int, ttl time.Duration) *dekCache { return &dekCache{lru: expirable.NewLRU[string, []byte](maxSize, nil, ttl)} } -// The cache and the caller each own their copy. func (c *dekCache) get(keyRef string) ([]byte, bool) { dek, ok := c.lru.Get(keyRef) if !ok { diff --git a/encryption/pool.go b/encryption/pool.go index 01a62a6..e5b4c5c 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -46,18 +46,15 @@ type Pool struct { cache *dekCache // nil when caching disabled } -// annotationDEKCache records how a DEK was obtained on the span of the -// operation that needed it: served from the cache, loaded through DynamoDB and -// KMS, or coalesced onto another goroutine's in-flight load. It is what tells -// you whether a slow Encrypt/Decrypt paid for a KMS round-trip. +// annotationDEKCache records whether an operation's DEK came from the cache, a +// full load, or another goroutine's in-flight load. const annotationDEKCache = "dek_cache" // PoolOption configures optional Pool behavior. type PoolOption func(*Pool) // WithCache enables an in-memory LRU cache for decrypted data encryption keys, -// eliminating KMS round-trips on cache hits. The cache is local to this process -// and zeroes key material on eviction. +// eliminating KMS round-trips on cache hits. The cache is local to this process. func WithCache(cfg CacheConfig) PoolOption { return func(p *Pool) { if cfg.MaxSize > 0 && cfg.TTL > 0 { @@ -126,11 +123,9 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext p.cache.put(key.KeyRef, privateKey) } } else { - // The row is an unverified read and its KeyRef selects the DEK this - // data is encrypted under, so the attestation is checked before the - // cache is consulted — a hit must not be able to skip it. Verification - // is local (COSE + cert chain), so the KMS round-trips are still what - // the cache saves. + // The row is an unverified read and its KeyRef picks the DEK, so the + // attestation is checked before the cache is consulted; a hit must not + // skip it. Verification is local, so the cache still saves the KMS calls. if err := p.VerifyKey(ctx, att, key); err != nil { return "", nil, fmt.Errorf("verify key: %w", err) } @@ -196,7 +191,6 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, err } - // Decrypt data. var decrypted []byte switch decoded.Version { case 1: @@ -211,17 +205,15 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return decrypted, nil } -// fetchDEK returns the DEK for keyRef, serving it from the cache when possible -// and otherwise loading it through loadDEK. Concurrent misses on the same keyRef -// are deduplicated so only one of them performs the KMS round-trips. +// fetchDEK serves keyRef from the cache, or loads it once on behalf of every +// caller that misses concurrently. func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) ([]byte, error) { if p.cache == nil { dek, _, err := p.loadDEK(ctx, att, keyRef) return dek, err } - // The caller's span, not loadDEK's child: the annotation belongs on the - // operation whose latency it explains. + // The caller's span, not loadDEK's child, which only exists on a miss. span := tracing.GetSpan(ctx) if dek, ok := p.cache.get(keyRef); ok { @@ -257,10 +249,9 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } } -// loadDEK recovers a DEK through the full path: DynamoDB lookup, attestation -// verification, KMS share decryption, and Shamir combine. It reports whether the -// result may be cached; a key whose migration failed must not be, or the retry -// is suppressed until the entry expires. +// loadDEK recovers a DEK through DynamoDB, attestation verification, KMS share +// decryption and Shamir combine. A key whose migration failed is not cacheable: +// caching it would suppress the retry until the entry expires. func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, cacheable bool, err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.loadDEK", tracing.WithAnnotation("key_ref", keyRef)) defer func() { @@ -297,8 +288,6 @@ func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef str return nil, false, fmt.Errorf("combine shares: %w", err) } - // Migration is non-fatal: the DEK is returned either way, but a failed - // migration leaves the key uncacheable so the next decrypt retries it. if p.keyNeedsMigration(key) { if err := p.migrateKey(ctx, att, key, privateKey); err != nil { p.logger.ErrorContext(ctx, "migrating key failed", "error", err, "key_ref", key.KeyRef, "generation", key.Generation, "key_index", key.KeyIndex) From e431744e302257a8dc83f510bdb3bdbdc87f968d Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Wed, 23 Sep 2026 18:54:16 +0200 Subject: [PATCH 08/12] style(encryption): drop unexported-func and backstory comments Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 9 +-------- encryption/pool.go | 14 ++++---------- 2 files changed, 5 insertions(+), 18 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index 835de47..756b3c0 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -17,15 +17,8 @@ type CacheConfig struct { TTL time.Duration } -// dekCache is an LRU of decrypted data encryption keys. -// -// Evicted keys are not zeroed: the LRU releases its lock before the caller has -// finished copying a value, so clearing one risks handing out a half-cleared -// key. The enclave already holds unscrubbed copies from Shamir recombination -// and the AES key schedule. type dekCache struct { - lru *expirable.LRU[string, []byte] - // group collapses concurrent misses on the same keyRef into one load. + lru *expirable.LRU[string, []byte] group singleflight.Group } diff --git a/encryption/pool.go b/encryption/pool.go index e5b4c5c..d4bac26 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -46,8 +46,6 @@ type Pool struct { cache *dekCache // nil when caching disabled } -// annotationDEKCache records whether an operation's DEK came from the cache, a -// full load, or another goroutine's in-flight load. const annotationDEKCache = "dek_cache" // PoolOption configures optional Pool behavior. @@ -123,9 +121,8 @@ func (p *Pool) Encrypt(ctx context.Context, att *enclave.Attestation, plaintext p.cache.put(key.KeyRef, privateKey) } } else { - // The row is an unverified read and its KeyRef picks the DEK, so the - // attestation is checked before the cache is consulted; a hit must not - // skip it. Verification is local, so the cache still saves the KMS calls. + // The row is unverified and its KeyRef picks the DEK, so a cache hit + // must not skip this. if err := p.VerifyKey(ctx, att, key); err != nil { return "", nil, fmt.Errorf("verify key: %w", err) } @@ -205,8 +202,6 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return decrypted, nil } -// fetchDEK serves keyRef from the cache, or loads it once on behalf of every -// caller that misses concurrently. func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef string) ([]byte, error) { if p.cache == nil { dek, _, err := p.loadDEK(ctx, att, keyRef) @@ -249,9 +244,8 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } } -// loadDEK recovers a DEK through DynamoDB, attestation verification, KMS share -// decryption and Shamir combine. A key whose migration failed is not cacheable: -// caching it would suppress the retry until the entry expires. +// A key whose migration failed is not cacheable: caching it would suppress the +// retry until the entry expires. func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef string) (privateKey []byte, cacheable bool, err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.loadDEK", tracing.WithAnnotation("key_ref", keyRef)) defer func() { From f922efdf7e3fb149fbf8ddfe00255e7a444d4402 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Thu, 24 Sep 2026 13:40:36 +0200 Subject: [PATCH 09/12] docs(encryption): note the cache's background eviction goroutine expirable.LRU starts a cleanup goroutine per instance and v2.0.7 has no way to stop it, so a Pool built WithCache holds it for the process. Co-Authored-By: Claude Opus 5 --- encryption/pool.go | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/encryption/pool.go b/encryption/pool.go index d4bac26..0d43f2b 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -52,7 +52,8 @@ const annotationDEKCache = "dek_cache" type PoolOption func(*Pool) // WithCache enables an in-memory LRU cache for decrypted data encryption keys, -// eliminating KMS round-trips on cache hits. The cache is local to this process. +// eliminating KMS round-trips on cache hits. The cache is local to this process +// and starts a background eviction goroutine that runs until the process exits. func WithCache(cfg CacheConfig) PoolOption { return func(p *Pool) { if cfg.MaxSize > 0 && cfg.TTL > 0 { From a3e3eaee513ed9098fadf2f24e4e2de1ed700ebe Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Thu, 24 Sep 2026 13:40:51 +0200 Subject: [PATCH 10/12] test(encryption): gate the loader so dedup is actually proven The concurrent decrypts could serialize, in which case later callers hit a warm cache and the test passed without any coalescing. The reverse order failed it: a caller that missed the cache before the first load's put could start a second load. Blocking the share decrypt until every caller is inside Decrypt keeps the cache empty for the whole window, so one load is the only possible outcome. Verified by failing with 9 loads when the singleflight key is made unique. Co-Authored-By: Claude Opus 5 --- encryption/pool_test.go | 25 ++++++++++++++++++------- 1 file changed, 18 insertions(+), 7 deletions(-) diff --git a/encryption/pool_test.go b/encryption/pool_test.go index 5ca4c13..f551c9e 100644 --- a/encryption/pool_test.go +++ b/encryption/pool_test.go @@ -1413,8 +1413,17 @@ func TestPool_DecryptSingleflight(t *testing.T) { remoteKey1.ExpectedCalls = nil remoteKey2.ExpectedCalls = nil - // Re-register expectations. - remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) + const N = 10 + + // Hold the load until every caller is inside Decrypt. The cache stays empty + // while it blocks, so a caller can only avoid its own load by coalescing. + started := make(chan struct{}, N) + release := make(chan struct{}) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1"). + Run(func(mock.Arguments) { + started <- struct{}{} + <-release + }).Return(shares[0], nil) remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) // Invalidate cache so all goroutines start with a cold cache. @@ -1423,15 +1432,15 @@ func TestPool_DecryptSingleflight(t *testing.T) { err = pool.RotateKey(context.Background(), att, "cipherKey4") require.NoError(t, err) - // Launch N concurrent decrypts. - const N = 10 - var wg sync.WaitGroup + var entered, wg sync.WaitGroup + entered.Add(N) errs := make([]error, N) results := make([]string, N) for i := 0; i < N; i++ { wg.Add(1) go func(idx int) { defer wg.Done() + entered.Done() pt, err := pool.Decrypt(context.Background(), att, "cipherKey4", ciphertext, []byte("aad")) errs[idx] = err if pt != nil { @@ -1439,6 +1448,10 @@ func TestPool_DecryptSingleflight(t *testing.T) { } }(i) } + + <-started + entered.Wait() + close(release) wg.Wait() for i := 0; i < N; i++ { @@ -1446,8 +1459,6 @@ func TestPool_DecryptSingleflight(t *testing.T) { require.Equal(t, "test", results[i], "goroutine %d got wrong result", i) } - // With singleflight, exactly one goroutine fetches. RemoteKey1.Decrypt - // should be called once (one share per remote key, one fetch total). decryptCalls := len(remoteKey1.Calls) require.Equal(t, 1, decryptCalls, "singleflight should deduplicate concurrent fetches (got %d calls to remoteKey1)", decryptCalls) } From 79233147c7952ce98519a19535151ad529469d87 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Thu, 24 Sep 2026 13:40:52 +0200 Subject: [PATCH 11/12] test(encryption): use testify in the DEK cache tests Matches pool_test.go in the same package. Co-Authored-By: Claude Opus 5 --- encryption/cache_test.go | 56 ++++++++++++++-------------------------- 1 file changed, 20 insertions(+), 36 deletions(-) diff --git a/encryption/cache_test.go b/encryption/cache_test.go index b809087..7804e5c 100644 --- a/encryption/cache_test.go +++ b/encryption/cache_test.go @@ -4,6 +4,8 @@ import ( "sync" "testing" "time" + + "github.com/stretchr/testify/require" ) func TestDEKCache_GetPut(t *testing.T) { @@ -13,34 +15,25 @@ func TestDEKCache_GetPut(t *testing.T) { c.put("ref1", dek) got, ok := c.get("ref1") - if !ok { - t.Fatal("expected cache hit") - } - if string(got) != string(dek) { - t.Fatalf("got %x, want %x", got, dek) - } + require.True(t, ok, "expected cache hit") + require.Equal(t, dek, got) // Returned slice must be an independent copy. got[0] = 0xFF got2, _ := c.get("ref1") - if got2[0] == 0xFF { - t.Fatal("cache returned same underlying slice, expected independent copy") - } + require.NotEqual(t, byte(0xFF), got2[0], "cache returned same underlying slice, expected independent copy") // Stored slice must be an independent copy of the input. dek[0] = 0xAA got3, _ := c.get("ref1") - if got3[0] == 0xAA { - t.Fatal("cache stored same underlying slice as input, expected independent copy") - } + require.NotEqual(t, byte(0xAA), got3[0], "cache stored same underlying slice as input, expected independent copy") } func TestDEKCache_Miss(t *testing.T) { c := newDEKCache(10, time.Minute) - if _, ok := c.get("unknown"); ok { - t.Fatal("expected cache miss for unknown key") - } + _, ok := c.get("unknown") + require.False(t, ok, "expected cache miss for unknown key") } func TestDEKCache_TTLExpiry(t *testing.T) { @@ -49,9 +42,8 @@ func TestDEKCache_TTLExpiry(t *testing.T) { time.Sleep(5 * time.Millisecond) - if _, ok := c.get("ref1"); ok { - t.Fatal("expected cache miss after TTL expiry") - } + _, ok := c.get("ref1") + require.False(t, ok, "expected cache miss after TTL expiry") } func TestDEKCache_LRUEviction(t *testing.T) { @@ -64,15 +56,12 @@ func TestDEKCache_LRUEviction(t *testing.T) { _, _ = c.get("ref1") c.put("ref3", []byte("key3key3key3key3key3key3key3key3")) - if _, ok := c.get("ref2"); ok { - t.Fatal("expected ref2 to be evicted (LRU)") - } - if _, ok := c.get("ref1"); !ok { - t.Fatal("expected ref1 to still be cached") - } - if _, ok := c.get("ref3"); !ok { - t.Fatal("expected ref3 to still be cached") - } + _, ok := c.get("ref2") + require.False(t, ok, "expected ref2 to be evicted (LRU)") + _, ok = c.get("ref1") + require.True(t, ok, "expected ref1 to still be cached") + _, ok = c.get("ref3") + require.True(t, ok, "expected ref3 to still be cached") } func TestDEKCache_PutUpdatesExisting(t *testing.T) { @@ -83,12 +72,8 @@ func TestDEKCache_PutUpdatesExisting(t *testing.T) { c.put("ref1", dek2) got, ok := c.get("ref1") - if !ok { - t.Fatal("expected cache hit") - } - if string(got) != string(dek2) { - t.Fatalf("got %x, want %x", got, dek2) - } + require.True(t, ok, "expected cache hit") + require.Equal(t, dek2, got) } func TestDEKCache_Delete(t *testing.T) { @@ -97,9 +82,8 @@ func TestDEKCache_Delete(t *testing.T) { c.put("ref1", []byte("0123456789abcdef0123456789abcdef")) c.delete("ref1") - if _, ok := c.get("ref1"); ok { - t.Fatal("expected cache miss after delete") - } + _, ok := c.get("ref1") + require.False(t, ok, "expected cache miss after delete") } func TestDEKCache_DeleteMissing(t *testing.T) { From fb21d65cd434da5762517fd2a93370eaf044f552 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Thu, 24 Sep 2026 13:47:28 +0200 Subject: [PATCH 12/12] fix(encryption): disable the DEK cache below a one-second TTL expirable.LRU runs its eviction ticker at TTL/100. Under 100ns that panics outright; between 100ns and a millisecond it silently spins the cleanup loop against the mutex every get and put takes, which is the worse of the two. A bare integer literal is a legal time.Duration, so `TTL: 600` compiles, reads as ten minutes, and lands in that range. Warn and run uncached instead. Co-Authored-By: Claude Opus 5 --- encryption/cache.go | 6 ++++++ encryption/cache_test.go | 18 ++++++++++++++++++ encryption/pool.go | 9 +++++++-- 3 files changed, 31 insertions(+), 2 deletions(-) diff --git a/encryption/cache.go b/encryption/cache.go index 756b3c0..f33f21e 100644 --- a/encryption/cache.go +++ b/encryption/cache.go @@ -14,9 +14,15 @@ type CacheConfig struct { // MaxSize is the maximum number of DEKs to cache. MaxSize int // TTL is how long a cached DEK lives, counted from when it is stored. + // Values below minCacheTTL disable the cache. TTL time.Duration } +// minCacheTTL keeps expirable.LRU's eviction ticker, which runs at TTL/100, +// above 10ms. Below 100ns it panics outright, and a bare integer literal is a +// legal time.Duration, so a TTL meant as seconds lands in that range. +const minCacheTTL = time.Second + type dekCache struct { lru *expirable.LRU[string, []byte] group singleflight.Group diff --git a/encryption/cache_test.go b/encryption/cache_test.go index 7804e5c..627aef0 100644 --- a/encryption/cache_test.go +++ b/encryption/cache_test.go @@ -1,6 +1,7 @@ package encryption import ( + "log/slog" "sync" "testing" "time" @@ -8,6 +9,23 @@ import ( "github.com/stretchr/testify/require" ) +func TestWithCache_DisablesCacheBelowMinTTL(t *testing.T) { + quiet := slog.New(slog.DiscardHandler) + + for _, cfg := range []CacheConfig{ + {MaxSize: 10, TTL: 0}, + {MaxSize: 0, TTL: time.Minute}, + {MaxSize: 10, TTL: 50}, // would panic expirable.LRU + {MaxSize: 10, TTL: 600}, // reads as seconds, means nanoseconds + } { + p := NewPool(nil, nil, nil, nil, quiet, WithCache(cfg)) + require.Nil(t, p.cache, "cache should be disabled for %+v", cfg) + } + + p := NewPool(nil, nil, nil, nil, quiet, WithCache(CacheConfig{MaxSize: 10, TTL: minCacheTTL})) + require.NotNil(t, p.cache) +} + func TestDEKCache_GetPut(t *testing.T) { c := newDEKCache(10, time.Minute) dek := []byte("0123456789abcdef0123456789abcdef") diff --git a/encryption/pool.go b/encryption/pool.go index 0d43f2b..a5153ad 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -56,9 +56,14 @@ type PoolOption func(*Pool) // and starts a background eviction goroutine that runs until the process exits. func WithCache(cfg CacheConfig) PoolOption { return func(p *Pool) { - if cfg.MaxSize > 0 && cfg.TTL > 0 { - p.cache = newDEKCache(cfg.MaxSize, cfg.TTL) + if cfg.MaxSize <= 0 || cfg.TTL <= 0 { + return } + if cfg.TTL < minCacheTTL { + p.logger.Warn("DEK cache disabled: TTL below minimum", "ttl", cfg.TTL, "minimum", minCacheTTL) + return + } + p.cache = newDEKCache(cfg.MaxSize, cfg.TTL) } }