diff --git a/encryption/cache.go b/encryption/cache.go new file mode 100644 index 0000000..f33f21e --- /dev/null +++ b/encryption/cache.go @@ -0,0 +1,49 @@ +package encryption + +import ( + "slices" + "time" + + "github.com/hashicorp/golang-lru/v2/expirable" + "golang.org/x/sync/singleflight" +) + +// 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. + 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 +} + +func newDEKCache(maxSize int, ttl time.Duration) *dekCache { + return &dekCache{lru: expirable.NewLRU[string, []byte](maxSize, nil, ttl)} +} + +func (c *dekCache) get(keyRef string) ([]byte, bool) { + dek, ok := c.lru.Get(keyRef) + if !ok { + return nil, false + } + return slices.Clone(dek), true +} + +func (c *dekCache) put(keyRef string, dek []byte) { + c.lru.Add(keyRef, slices.Clone(dek)) +} + +func (c *dekCache) delete(keyRef string) { + c.lru.Remove(keyRef) +} diff --git a/encryption/cache_test.go b/encryption/cache_test.go new file mode 100644 index 0000000..627aef0 --- /dev/null +++ b/encryption/cache_test.go @@ -0,0 +1,135 @@ +package encryption + +import ( + "log/slog" + "sync" + "testing" + "time" + + "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") + + c.put("ref1", dek) + + got, ok := c.get("ref1") + 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") + 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") + 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) + + _, ok := c.get("unknown") + require.False(t, ok, "expected cache miss for unknown key") +} + +func TestDEKCache_TTLExpiry(t *testing.T) { + c := newDEKCache(10, time.Millisecond) + c.put("ref1", []byte("0123456789abcdef0123456789abcdef")) + + time.Sleep(5 * time.Millisecond) + + _, ok := c.get("ref1") + require.False(t, ok, "expected cache miss after TTL expiry") +} + +func TestDEKCache_LRUEviction(t *testing.T) { + c := newDEKCache(2, time.Minute) + + c.put("ref1", []byte("key1key1key1key1key1key1key1key1")) + c.put("ref2", []byte("key2key2key2key2key2key2key2key2")) + + // Make ref1 more recent than ref2, so ref3 evicts ref2. + _, _ = c.get("ref1") + c.put("ref3", []byte("key3key3key3key3key3key3key3key3")) + + _, 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) { + c := newDEKCache(10, time.Minute) + dek2 := []byte("new_key_new_key_new_key_new_key_") + + c.put("ref1", []byte("old_key_old_key_old_key_old_key_")) + c.put("ref1", dek2) + + got, ok := c.get("ref1") + require.True(t, ok, "expected cache hit") + require.Equal(t, dek2, got) +} + +func TestDEKCache_Delete(t *testing.T) { + c := newDEKCache(10, time.Minute) + + c.put("ref1", []byte("0123456789abcdef0123456789abcdef")) + c.delete("ref1") + + _, ok := c.get("ref1") + require.False(t, ok, "expected cache miss after delete") +} + +func TestDEKCache_DeleteMissing(t *testing.T) { + c := newDEKCache(10, time.Minute) + c.delete("nonexistent") +} + +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() { + defer wg.Done() + c.put("ref1", dek) + }() + go func() { + defer wg.Done() + c.get("ref1") + }() + go func(n int) { + defer wg.Done() + if n%10 == 0 { + c.delete("ref1") + } + }(i) + } + wg.Wait() +} diff --git a/encryption/pool.go b/encryption/pool.go index 0a5fe97..a5153ad 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "log/slog" + "slices" "strconv" "time" @@ -42,19 +43,45 @@ 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 { +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 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 { + 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) + } +} + +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 +122,38 @@ 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 { + // 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) + } + + 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) + } + + 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 { @@ -126,7 +174,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. +// 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() { @@ -139,15 +189,85 @@ func (p *Pool) Decrypt(ctx context.Context, att *enclave.Attestation, keyRef str return nil, fmt.Errorf("decode ciphertext: %w", err) } + privateKey, err := p.fetchDEK(ctx, att, keyRef) + if err != nil { + return nil, 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) + } + + return decrypted, nil +} + +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, which only exists on a miss. + span := tracing.GetSpan(ctx) + + if dek, ok := p.cache.get(keyRef); ok { + span.SetAnnotation(annotationDEKCache, "hit") + return dek, nil + } + + loaded := p.cache.group.DoChan(keyRef, func() (any, error) { + dek, cacheable, err := p.loadDEK(ctx, att, keyRef) + if err != nil { + return nil, err + } + if cacheable { + 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() + } +} + +// 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() { + 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) + 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)) @@ -157,37 +277,25 @@ func (p *Pool) Decrypt(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") - } - - 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("shares are invalid") } - 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) - } + privateKey, err = p.combineShares(ctx, att, config, key.EncryptedShares) if err != nil { - return nil, fmt.Errorf("decrypt: %w", err) + return nil, false, fmt.Errorf("combine shares: %w", err) } 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 privateKey, false, nil } } - return decrypted, 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 @@ -232,6 +340,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 } @@ -280,6 +392,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 ef1b65a..f551c9e 100644 --- a/encryption/pool_test.go +++ b/encryption/pool_test.go @@ -5,6 +5,7 @@ import ( "crypto/x509" "encoding/pem" "errors" + "sync" "testing" "time" @@ -51,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) @@ -1062,3 +1062,622 @@ 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 + + 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. + 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) + + 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 { + results[idx] = string(pt) + } + }(i) + } + + <-started + entered.Wait() + close(release) + 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) + } + + 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 } + +// 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") + }) +} diff --git a/go.mod b/go.mod index d517f51..b77c952 100644 --- a/go.mod +++ b/go.mod @@ -12,7 +12,9 @@ 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 ) @@ -30,6 +32,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 ) 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=