Files
Chris Lu 430d4b5394 kms: cache decrypted data keys so cache_enabled/cache_ttl take effect (#10173)
* kms: cache decrypted data keys so cache_enabled/cache_ttl take effect

cache_enabled/cache_ttl were parsed into KMSConfig but never consulted,
so every SSE-KMS read repeated a KMS Decrypt round-trip. Add a
CachedKMSProvider decorator keyed on (ciphertext, encryption context)
and wrap providers with it wherever one is created. Decrypt is
deterministic for that pair; GenerateDataKey/DescribeKey/GetKeyID pass
through untouched.

Cache hits and stores return private copies so the read path's
ClearSensitiveData can't wipe the cached key. TTL and max_cache_size
bound the cache.

* kms: zero superseded data keys and pool cache-key hash states

Wipe the old plaintext when a cache entry is overwritten so a key
material buffer left behind by two readers racing on the same miss does
not linger in memory. Reuse sha256 states via a sync.Pool so the read
hot path stops allocating a fresh hash on every Decrypt.

* kms: drop cache writes after Close

An in-flight Decrypt miss can call set after Close scrubbed the map,
repopulating it with a data key that would then never be cleared. Guard
set with a closed flag so post-Close stores wipe their copy and no-op.
2026-06-30 21:23:53 -07:00

246 lines
7.1 KiB
Go

package kms
import (
"context"
"crypto/sha256"
"encoding/binary"
"hash"
"sort"
"sync"
"time"
"github.com/seaweedfs/seaweedfs/weed/glog"
)
const (
// DefaultDataKeyCacheTTL is used when caching is enabled without an explicit TTL.
DefaultDataKeyCacheTTL = time.Hour
// DefaultDataKeyCacheSize bounds the number of cached decrypted data keys.
DefaultDataKeyCacheSize = 1000
)
// CachedKMSProvider wraps a KMSProvider with an in-memory cache of decrypted
// data keys. A KMS Decrypt call is deterministic for a given (ciphertext,
// encryption context) pair, so every read of the same SSE-KMS object would
// otherwise repeat an identical, network-bound KMS round-trip. Caching the
// plaintext data key removes that round-trip on the read hot path — this is
// what the cache_enabled / cache_ttl provider settings control.
//
// Only Decrypt is cached. GenerateDataKey must return a unique key per call,
// and DescribeKey / GetKeyID are cheap metadata lookups, so all three pass
// straight through to the wrapped provider via the embedded interface.
type CachedKMSProvider struct {
KMSProvider // wrapped provider; supplies the un-cached methods
ttl time.Duration
maxEntries int
mu sync.Mutex
entries map[string]*dataKeyCacheEntry
closed bool
}
type dataKeyCacheEntry struct {
keyID string
plaintext []byte
expiresAt time.Time
}
// NewCachedKMSProvider wraps provider with a decrypted-data-key cache. A
// non-positive ttl or maxEntries falls back to the package defaults.
func NewCachedKMSProvider(provider KMSProvider, ttl time.Duration, maxEntries int) *CachedKMSProvider {
if ttl <= 0 {
ttl = DefaultDataKeyCacheTTL
}
if maxEntries <= 0 {
maxEntries = DefaultDataKeyCacheSize
}
return &CachedKMSProvider{
KMSProvider: provider,
ttl: ttl,
maxEntries: maxEntries,
entries: make(map[string]*dataKeyCacheEntry),
}
}
// Decrypt returns the plaintext data key from cache when present, otherwise
// delegates to the wrapped provider and caches the result.
func (c *CachedKMSProvider) Decrypt(ctx context.Context, req *DecryptRequest) (*DecryptResponse, error) {
if req == nil || len(req.CiphertextBlob) == 0 {
// Let the wrapped provider produce its normal validation error.
return c.KMSProvider.Decrypt(ctx, req)
}
cacheKey := dataKeyCacheKey(req)
if resp := c.get(cacheKey); resp != nil {
glog.V(4).Infof("KMS data key cache hit")
return resp, nil
}
resp, err := c.KMSProvider.Decrypt(ctx, req)
if err != nil {
return nil, err
}
// Store our own copy, then hand the original back to the caller. The
// caller clears the plaintext it receives (ClearSensitiveData); keeping a
// private copy means later cache hits still return valid key material.
c.set(cacheKey, resp)
return resp, nil
}
// get returns a fresh copy of the cached response, or nil on miss/expiry. The
// copy is essential: callers clear the returned plaintext after use, which
// would otherwise wipe the cached key material for every subsequent reader.
func (c *CachedKMSProvider) get(cacheKey string) *DecryptResponse {
c.mu.Lock()
defer c.mu.Unlock()
entry, ok := c.entries[cacheKey]
if !ok {
return nil
}
if time.Now().After(entry.expiresAt) {
ClearSensitiveData(entry.plaintext)
delete(c.entries, cacheKey)
return nil
}
plaintext := make([]byte, len(entry.plaintext))
copy(plaintext, entry.plaintext)
return &DecryptResponse{KeyID: entry.keyID, Plaintext: plaintext}
}
// set stores a private copy of the plaintext data key.
func (c *CachedKMSProvider) set(cacheKey string, resp *DecryptResponse) {
if resp == nil {
return
}
plaintext := make([]byte, len(resp.Plaintext))
copy(plaintext, resp.Plaintext)
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
// A miss can finish its underlying Decrypt after Close scrubbed the
// map; don't repopulate it with key material that won't be cleared.
ClearSensitiveData(plaintext)
return
}
if old, exists := c.entries[cacheKey]; exists {
// Two readers can race on the same miss and both store; wipe the
// superseded key material instead of leaving it to linger in memory.
ClearSensitiveData(old.plaintext)
} else {
c.evictIfFullLocked()
}
c.entries[cacheKey] = &dataKeyCacheEntry{
keyID: resp.KeyID,
plaintext: plaintext,
expiresAt: time.Now().Add(c.ttl),
}
}
// evictIfFullLocked drops expired entries and, if still at capacity, the entry
// closest to expiry. Callers must hold c.mu.
func (c *CachedKMSProvider) evictIfFullLocked() {
if len(c.entries) < c.maxEntries {
return
}
now := time.Now()
for k, entry := range c.entries {
if now.After(entry.expiresAt) {
ClearSensitiveData(entry.plaintext)
delete(c.entries, k)
}
}
if len(c.entries) < c.maxEntries {
return
}
// Still full: evict the soonest-to-expire entry.
var oldestKey string
var oldestAt time.Time
for k, entry := range c.entries {
if oldestKey == "" || entry.expiresAt.Before(oldestAt) {
oldestKey = k
oldestAt = entry.expiresAt
}
}
if oldestKey != "" {
ClearSensitiveData(c.entries[oldestKey].plaintext)
delete(c.entries, oldestKey)
}
}
// Close clears cached key material and closes the wrapped provider.
func (c *CachedKMSProvider) Close() error {
c.mu.Lock()
c.closed = true
for k, entry := range c.entries {
ClearSensitiveData(entry.plaintext)
delete(c.entries, k)
}
c.mu.Unlock()
return c.KMSProvider.Close()
}
// sha256Pool reuses hash states across cache-key computations so the read hot
// path doesn't allocate a fresh one on every Decrypt.
var sha256Pool = sync.Pool{
New: func() interface{} { return sha256.New() },
}
// dataKeyCacheKey derives a stable cache key from the ciphertext blob and the
// encryption context. The context is part of the key because the same
// ciphertext decrypted under a different context is a different request, and
// caching by ciphertext alone would bypass the provider's context check.
func dataKeyCacheKey(req *DecryptRequest) string {
h := sha256Pool.Get().(hash.Hash)
// Length-prefix each field so distinct inputs cannot collide by concatenation.
writeLengthPrefixed(h, req.CiphertextBlob)
keys := make([]string, 0, len(req.EncryptionContext))
for k := range req.EncryptionContext {
keys = append(keys, k)
}
sort.Strings(keys)
for _, k := range keys {
writeLengthPrefixed(h, []byte(k))
writeLengthPrefixed(h, []byte(req.EncryptionContext[k]))
}
key := string(h.Sum(nil))
h.Reset()
sha256Pool.Put(h)
return key
}
func writeLengthPrefixed(h hash.Hash, b []byte) {
var lenBuf [8]byte
binary.BigEndian.PutUint64(lenBuf[:], uint64(len(b)))
h.Write(lenBuf[:])
h.Write(b)
}
// wrapWithCacheIfEnabled returns provider wrapped in a data-key cache when the
// configuration enables caching, otherwise provider unchanged.
func wrapWithCacheIfEnabled(provider KMSProvider, config *KMSConfig) KMSProvider {
if provider == nil || config == nil || !config.CacheEnabled {
return provider
}
// Avoid double-wrapping if a cached provider is somehow reused.
if _, ok := provider.(*CachedKMSProvider); ok {
return provider
}
cached := NewCachedKMSProvider(provider, config.CacheTTL, config.MaxCacheSize)
glog.V(1).Infof("KMS data key caching enabled (ttl=%s, max=%d)", cached.ttl, cached.maxEntries)
return cached
}