From 217881f2975d1cc63893714315b1d221c1ca8bb5 Mon Sep 17 00:00:00 2001 From: Patryk Kalinowski Date: Fri, 25 Sep 2026 11:56:35 +0200 Subject: [PATCH] fix(encryption): detach a coalesced DEK load from its leader DoChan ran the shared load on whichever caller arrived first. If that request was cancelled mid-load, every caller coalesced onto the same key ref got its context.Canceled, even though their own contexts were live. In waas that surfaces as sessions silently missing from a config reconcile, because the caller logs and skips a binding it cannot decrypt. The load now runs on a context carrying the leader's values but neither its cancellation nor its deadline, bounded by dekLoadTimeout so a load nothing is waiting for still terminates. Waiters keep leaving early on their own context. Two consequences of outliving the request, both handled here: migrateKey re-encrypts shares, which draws the AES-CBC IV from the NSM session; the leader's session is closed once its handler returns. It now opens its own attestation and no longer takes one from the caller. Span mutations are mutex-guarded but the middleware marshals the tree by reflection without the lock, so a still-running child would be read while it is written. tracing.Detach starts the load's span as a new root. Co-Authored-By: Claude Opus 5 --- encryption/mocks_test.go | 4 ++ encryption/pool.go | 19 ++++++-- encryption/pool_test.go | 96 ++++++++++++++++++++++++++++++++++++++-- tracing/span.go | 5 +++ 4 files changed, 117 insertions(+), 7 deletions(-) diff --git a/encryption/mocks_test.go b/encryption/mocks_test.go index 1b23b19..0f22101 100644 --- a/encryption/mocks_test.go +++ b/encryption/mocks_test.go @@ -113,6 +113,10 @@ func (m *MockRemoteKey) Encrypt(ctx context.Context, att *enclave.Attestation, p func (m *MockRemoteKey) Decrypt(ctx context.Context, att *enclave.Attestation, ciphertext string) ([]byte, error) { args := m.Called(ctx, att, ciphertext) + // A real KMS call observes cancellation while in flight. + if err := ctx.Err(); err != nil { + return nil, err + } if args.Get(0) == nil { return nil, args.Error(1) } diff --git a/encryption/pool.go b/encryption/pool.go index a5153ad..12a2b2f 100644 --- a/encryption/pool.go +++ b/encryption/pool.go @@ -48,6 +48,9 @@ type Pool struct { const annotationDEKCache = "dek_cache" +// dekLoadTimeout bounds a coalesced load that no longer answers to any caller. +const dekLoadTimeout = 30 * time.Second + // PoolOption configures optional Pool behavior. type PoolOption func(*Pool) @@ -223,7 +226,11 @@ func (p *Pool) fetchDEK(ctx context.Context, att *enclave.Attestation, keyRef st } loaded := p.cache.group.DoChan(keyRef, func() (any, error) { - dek, cacheable, err := p.loadDEK(ctx, att, keyRef) + // Shared by everyone who coalesces here, so no one caller's cancellation ends it. + loadCtx, cancel := context.WithTimeout(tracing.Detach(context.WithoutCancel(ctx)), dekLoadTimeout) + defer cancel() + + dek, cacheable, err := p.loadDEK(loadCtx, att, keyRef) if err != nil { return nil, err } @@ -289,7 +296,7 @@ func (p *Pool) loadDEK(ctx context.Context, att *enclave.Attestation, keyRef str } if p.keyNeedsMigration(key) { - if err := p.migrateKey(ctx, att, key, privateKey); err != nil { + if err := p.migrateKey(ctx, 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 } @@ -536,7 +543,7 @@ func (p *Pool) keyNeedsMigration(key *data.CipherKey) bool { return key.Generation < generation } -func (p *Pool) migrateKey(ctx context.Context, att *enclave.Attestation, key *data.CipherKey, privateKey []byte) (err error) { +func (p *Pool) migrateKey(ctx context.Context, key *data.CipherKey, privateKey []byte) (err error) { ctx, span := tracing.Trace(ctx, "encryption.Pool.migrateKey") defer func() { span.RecordError(err) @@ -553,6 +560,12 @@ func (p *Pool) migrateKey(ctx context.Context, att *enclave.Attestation, key *da return fmt.Errorf("split private key: %w", err) } + att, err := p.attester.GetAttestation(ctx, nil, nil) + if err != nil { + return fmt.Errorf("get attestation: %w", err) + } + defer func() { _ = att.Close() }() + i := 0 encryptedShares := make(map[string]string) for remoteKeyID, remoteKey := range config.RemoteKeys { diff --git a/encryption/pool_test.go b/encryption/pool_test.go index f551c9e..ad49f39 100644 --- a/encryption/pool_test.go +++ b/encryption/pool_test.go @@ -622,8 +622,8 @@ func TestPool_Decrypt(t *testing.T) { remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Once().Return(shares[0], nil) remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Once().Return(shares[1], nil) - remoteKey3.On("Encrypt", mock.Anything, att, mock.Anything).Once().Return("encryptedShare3", nil) - remoteKey4.On("Encrypt", mock.Anything, att, mock.Anything).Once().Return("encryptedShare4", nil) + remoteKey3.On("Encrypt", mock.Anything, mock.AnythingOfType("*enclave.Attestation"), mock.Anything).Once().Return("encryptedShare3", nil) + remoteKey4.On("Encrypt", mock.Anything, mock.AnythingOfType("*enclave.Attestation"), mock.Anything).Once().Return("encryptedShare4", nil) var migratedKey *data.CipherKey createMatcher := func(key *data.CipherKey) bool { @@ -700,8 +700,8 @@ func TestPool_Decrypt(t *testing.T) { remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1").Return(shares[0], nil) remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], nil) - remoteKey3.On("Encrypt", mock.Anything, att, mock.Anything).Return("encryptedShare3", nil) - remoteKey4.On("Encrypt", mock.Anything, att, mock.Anything).Return("", errors.New("mock error")) + remoteKey3.On("Encrypt", mock.Anything, mock.AnythingOfType("*enclave.Attestation"), mock.Anything).Return("encryptedShare3", nil) + remoteKey4.On("Encrypt", mock.Anything, mock.AnythingOfType("*enclave.Attestation"), mock.Anything).Return("", errors.New("mock error")) pool := encryption.NewPool(enc, configs, keysTable, nil, nil) plaintext, err := pool.Decrypt(context.Background(), att, "cipherKey4", legacyCiphertext55_v2, []byte("aad")) @@ -1463,6 +1463,94 @@ func TestPool_DecryptSingleflight(t *testing.T) { require.Equal(t, 1, decryptCalls, "singleflight should deduplicate concurrent fetches (got %d calls to remoteKey1)", decryptCalls) } +func TestPool_DecryptSurvivesLeaderCancellation(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) + + 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) + + remoteKey1.Calls, remoteKey1.ExpectedCalls = nil, nil + remoteKey2.Calls, remoteKey2.ExpectedCalls = nil, nil + + // Hold the load until the leader has given up, so the follower is provably + // waiting on a load started from a context that is already cancelled. + leaderGone := make(chan struct{}) + loadStarted := make(chan struct{}, 2) + remoteKey1.On("Decrypt", mock.Anything, att, "encryptedShare1"). + Run(func(mock.Arguments) { + loadStarted <- struct{}{} + <-leaderGone + }).Return(shares[0], nil) + remoteKey2.On("Decrypt", mock.Anything, att, "encryptedShare2").Return(shares[1], 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) + require.NoError(t, pool.RotateKey(context.Background(), att, "cipherKey4")) + + leaderCtx, cancelLeader := context.WithCancel(context.Background()) + leaderDone := make(chan struct{}) + go func() { + defer close(leaderDone) + _, _ = pool.Decrypt(leaderCtx, att, "cipherKey4", ciphertext, []byte("aad")) + }() + + // The leader owns the load once its share decrypt is in flight. + <-loadStarted + + followerDone := make(chan error, 1) + var plaintext []byte + go func() { + pt, err := pool.Decrypt(context.Background(), att, "cipherKey4", ciphertext, []byte("aad")) + plaintext = pt + followerDone <- err + }() + + cancelLeader() + <-leaderDone + close(leaderGone) + + require.NoError(t, <-followerDone, "a cancelled leader must not fail the coalesced load") + require.Equal(t, "test", string(plaintext)) +} + func TestPool_EncryptNewKeyPopulatesCache(t *testing.T) { block, _ := pem.Decode([]byte(dummyPrivKey)) privKey, err := x509.ParsePKCS1PrivateKey(block.Bytes) diff --git a/tracing/span.go b/tracing/span.go index f7c2455..e68163a 100644 --- a/tracing/span.go +++ b/tracing/span.go @@ -55,6 +55,11 @@ func Trace(ctx context.Context, name string, opts ...func(*Span)) (context.Conte return context.WithValue(ctx, spanKey{}, span), span } +// Detach returns ctx without its current span, so the next Trace starts a new root. +func Detach(ctx context.Context) context.Context { + return context.WithValue(ctx, spanKey{}, (*Span)(nil)) +} + type spanKey struct{} type configKey struct{}