Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions encryption/mocks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down
19 changes: 16 additions & 3 deletions encryption/pool.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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)
Expand All @@ -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 {
Expand Down
96 changes: 92 additions & 4 deletions encryption/pool_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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"))
Expand Down Expand Up @@ -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)
Expand Down
5 changes: 5 additions & 0 deletions tracing/span.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{}

Expand Down
Loading