diff --git a/backend/internal/cmds/encryption_key_rotate.go b/backend/internal/cmds/encryption_key_rotate.go index 8b9c0b3a..36f65f83 100644 --- a/backend/internal/cmds/encryption_key_rotate.go +++ b/backend/internal/cmds/encryption_key_rotate.go @@ -147,7 +147,7 @@ func rotateJWKEncryption(ctx context.Context, db *gorm.DB, oldKek []byte, newKek return fmt.Errorf("failed to init key provider for %q with new encryption key: %w", dbKey, err) } - err = newProvider.SaveKey(ctx, key) + err = newProvider.ReplaceKey(ctx, key) if err != nil { return fmt.Errorf("failed to store key %q with new encryption key: %w", dbKey, err) } diff --git a/backend/internal/cmds/key_rotate.go b/backend/internal/cmds/key_rotate.go index 52625b01..b9303dce 100644 --- a/backend/internal/cmds/key_rotate.go +++ b/backend/internal/cmds/key_rotate.go @@ -113,7 +113,7 @@ func keyRotate(ctx context.Context, flags keyRotateFlags, db *gorm.DB, instanceI } // Save the key - err = keyProvider.SaveKey(ctx, key) + err = keyProvider.ReplaceKey(ctx, key) if err != nil { return fmt.Errorf("failed to store new key: %w", err) } @@ -150,7 +150,7 @@ func sessionKeyRotate(ctx context.Context, flags keyRotateFlags, db *gorm.DB, in } // Save the key - err = keyProvider.SaveKey(ctx, key) + err = keyProvider.ReplaceKey(ctx, key) if err != nil { return fmt.Errorf("failed to store new session key: %w", err) } diff --git a/backend/internal/service/jwt_service.go b/backend/internal/service/jwt_service.go index 8a160913..5a07e90a 100644 --- a/backend/internal/service/jwt_service.go +++ b/backend/internal/service/jwt_service.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "time" "uuid" @@ -15,6 +16,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/model" + "github.com/pocket-id/pocket-id/backend/internal/utils" jwkutils "github.com/pocket-id/pocket-id/backend/internal/utils/jwk" ) @@ -77,13 +79,13 @@ func (s *JwtService) init(ctx context.Context, db *gorm.DB, instanceID string, e // LoadOrGenerateKey loads the signing keys from the database, generating and persisting them if they don't exist yet func (s *JwtService) LoadOrGenerateKey(ctx context.Context) error { // Load the key used for tokens that are consumed externally, such as ID tokens and access tokens for apps - err := s.loadOrGenerateSigningKey(ctx) + err := retryKeyStorage(ctx, s.loadOrGenerateSigningKey) if err != nil { return fmt.Errorf("error loading signing key: %w", err) } // Load the key used for Pocket ID's own sessions, which is symmetric - err = s.loadOrGenerateSessionKey(ctx) + err = retryKeyStorage(ctx, s.loadOrGenerateSessionKey) if err != nil { return fmt.Errorf("error loading session key: %w", err) } @@ -91,6 +93,28 @@ func (s *JwtService) LoadOrGenerateKey(ctx context.Context) error { return nil } +func retryKeyStorage(ctx context.Context, loadOrGenerate func(context.Context) error) error { + for retries := 0; ; retries++ { + // Run the full load and store sequence so a key written by another replica takes precedence + err := loadOrGenerate(ctx) + if !errors.Is(err, jwkutils.ErrRetryKeyStorage) { + return err + } + + // Return the last conflict after the configured number of retries has been exhausted + if retries == 3 { + return err + } + + // Wait briefly before reloading so the competing database transaction has time to finish + slog.WarnContext(ctx, "Failed to store key, retrying", slog.Int("retry", retries+1), slog.Any("error", err)) + err = utils.SleepWithContext(ctx, 200*time.Millisecond) + if err != nil { + return err + } + } +} + func (s *JwtService) loadOrGenerateSigningKey(ctx context.Context) error { // Get the key provider keyProvider, err := jwkutils.GetKeyProvider(s.db, s.envConfig, s.instanceID) diff --git a/backend/internal/service/jwt_service_test.go b/backend/internal/service/jwt_service_test.go index 147c700d..54611963 100644 --- a/backend/internal/service/jwt_service_test.go +++ b/backend/internal/service/jwt_service_test.go @@ -1,11 +1,13 @@ package service import ( + "context" "crypto/ecdsa" "crypto/ed25519" "crypto/elliptic" "crypto/rand" "crypto/rsa" + "errors" "sync" "testing" "time" @@ -78,7 +80,7 @@ func saveKeyToDatabase(t *testing.T, db *gorm.DB, instanceID string, envConfig * keyProvider, err := jwkutils.GetKeyProvider(db, envConfig, instanceID) require.NoError(t, err, "Failed to init key provider") - err = keyProvider.SaveKey(t.Context(), key) + err = keyProvider.ReplaceKey(t.Context(), key) require.NoError(t, err, "Failed to save key") kid, ok := key.KeyID() @@ -188,6 +190,51 @@ func TestJwtService_Init(t *testing.T) { assert.Equal(t, origKeyID, loadedKeyID, "Loaded key should have the same ID as the original") }) + for _, dbKey := range []string{jwkutils.PrivateKeyDBKey, jwkutils.SessionKeyDBKey} { + t.Run("should not retry a failed database write for "+dbKey, func(t *testing.T) { + // Configure the database to fail the first attempt for the selected key + db := testutils.NewDatabaseForTest(t) + mockEnvConfig := newTestEnvConfig() + instanceID := newInstanceID(t, db) + storeAttempts := 0 + storeErr := errors.New("test database error") + + err := db.Callback().Create().Before("gorm:create").Register("fail_first_key_storage", func(tx *gorm.DB) { + row, ok := tx.Statement.Dest.(*model.KV) + if !ok || row.Key != dbKey { + return + } + + storeAttempts++ + if storeAttempts == 1 { + _ = tx.AddError(storeErr) + } + }) + require.NoError(t, err) + + // Initialize the service and preserve the ordinary database failure + service := &JwtService{} + err = service.init(t.Context(), db, instanceID, mockEnvConfig) + require.ErrorIs(t, err, storeErr) + + // Verify the failed write was not retried + require.Equal(t, 1, storeAttempts) + }) + } + +} + +func TestRetryKeyStorageStopsAfterThreeRetries(t *testing.T) { + // Return a conflict on every attempt so the retry limit is reached + attempts := 0 + err := retryKeyStorage(t.Context(), func(_ context.Context) error { + attempts++ + return jwkutils.ErrRetryKeyStorage + }) + + // Verify the initial attempt was followed by exactly three retries + require.ErrorIs(t, err, jwkutils.ErrRetryKeyStorage) + require.Equal(t, 1+maxKeyStorageRetries, attempts) } func TestJwtService_SessionKey(t *testing.T) { @@ -304,7 +351,7 @@ func TestJwtService_SessionKey(t *testing.T) { require.NoError(t, err) keyProvider, err := jwkutils.GetSessionKeyProvider(db, mockEnvConfig, instanceID) require.NoError(t, err) - require.NoError(t, keyProvider.SaveKey(t.Context(), rotatedKey)) + require.NoError(t, keyProvider.ReplaceKey(t.Context(), rotatedKey)) require.NoError(t, service.LoadOrGenerateKey(t.Context())) // Tokens issued with the previous session key must no longer be accepted diff --git a/backend/internal/utils/jwk/key_provider.go b/backend/internal/utils/jwk/key_provider.go index fd81374c..5a8faa70 100644 --- a/backend/internal/utils/jwk/key_provider.go +++ b/backend/internal/utils/jwk/key_provider.go @@ -21,6 +21,7 @@ type KeyProvider interface { Init(opts KeyProviderOpts) error LoadKey(ctx context.Context) (jwk.Key, error) SaveKey(ctx context.Context, key jwk.Key) error + ReplaceKey(ctx context.Context, key jwk.Key) error } // GetKeyProvider returns the provider for the key used to sign tokens that are consumed externally, such as ID tokens and access tokens for apps diff --git a/backend/internal/utils/jwk/key_provider_database.go b/backend/internal/utils/jwk/key_provider_database.go index c436c1c7..f40f73dd 100644 --- a/backend/internal/utils/jwk/key_provider_database.go +++ b/backend/internal/utils/jwk/key_provider_database.go @@ -23,6 +23,9 @@ const ( SessionKeyDBKey = "session_key.json" ) +// ErrRetryKeyStorage signals that the caller should reload the key before trying to store it again +var ErrRetryKeyStorage = errors.New("retry key storage") + type KeyProviderDatabase struct { db *gorm.DB kek []byte @@ -88,23 +91,43 @@ func (f *KeyProviderDatabase) LoadKey(ctx context.Context) (key jwk.Key, err err } func (f *KeyProviderDatabase) SaveKey(ctx context.Context, key jwk.Key) error { - // Encode the key to JSON - data, err := EncodeJWKBytes(key) + // Prepare the encrypted database value before attempting the insert + row, err := f.prepareKeyRow(key) if err != nil { - return fmt.Errorf("failed to encode key to JSON: %w", err) + return err } - // Encrypt the key then encode to Base64 - enc, err := cryptoutils.Encrypt(f.kek, data, nil) - if err != nil { - return fmt.Errorf("failed to encrypt key: %w", err) - } - // Save to database - row := model.KV{ - Key: f.dbKey, - Value: new(base64.StdEncoding.EncodeToString(enc)), + // Insert only if the key doesn't exist yet + ctx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + result := f.db. + WithContext(ctx). + Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "key"}}, + DoNothing: true, + }). + Create(&row) + if result.Error != nil { + // Preserve ordinary database failures because only an existing row can be resolved by reloading + return fmt.Errorf("failed to store key in database: %w", result.Error) } + // Ask the caller to reload the winning key when another writer created the row first + if result.RowsAffected == 0 { + return ErrRetryKeyStorage + } + + return nil +} + +func (f *KeyProviderDatabase) ReplaceKey(ctx context.Context, key jwk.Key) error { + // Prepare the encrypted database value before attempting the replacement + row, err := f.prepareKeyRow(key) + if err != nil { + return err + } + + // Upsert explicitly because key rotation must replace an existing key and can also recover a missing row ctx, cancel := context.WithTimeout(ctx, 10*time.Second) defer cancel() err = f.db. @@ -116,13 +139,31 @@ func (f *KeyProviderDatabase) SaveKey(ctx context.Context, key jwk.Key) error { Create(&row). Error if err != nil { - // There's one scenario where if Pocket ID is started fresh with more than 1 replica, they both could be trying to create the key in the database at the same time - // In this case, only one of the replicas will succeed and the other one(s) will return an error here, which will cascade down and cause the replica(s) to crash and be restarted (at that point they'll load the then-existing key from the database) - return fmt.Errorf("failed to store key in database: %w", err) + return fmt.Errorf("failed to replace key in database: %w", err) } return nil } +func (f *KeyProviderDatabase) prepareKeyRow(key jwk.Key) (model.KV, error) { + // Encode the key to JSON + data, err := EncodeJWKBytes(key) + if err != nil { + return model.KV{}, fmt.Errorf("failed to encode key to JSON: %w", err) + } + + // Encrypt the key then encode to Base64 + enc, err := cryptoutils.Encrypt(f.kek, data, nil) + if err != nil { + return model.KV{}, fmt.Errorf("failed to encrypt key: %w", err) + } + + // Build the row once so inserts and explicit replacements encode keys identically + return model.KV{ + Key: f.dbKey, + Value: new(base64.StdEncoding.EncodeToString(enc)), + }, nil +} + // Compile-time interface check var _ KeyProvider = (*KeyProviderDatabase)(nil) diff --git a/backend/internal/utils/jwk/key_provider_database_test.go b/backend/internal/utils/jwk/key_provider_database_test.go index 2b7f5977..8a8e9fe2 100644 --- a/backend/internal/utils/jwk/key_provider_database_test.go +++ b/backend/internal/utils/jwk/key_provider_database_test.go @@ -5,11 +5,13 @@ import ( "crypto/elliptic" "crypto/rand" "encoding/base64" + "errors" "testing" "github.com/lestrrat-go/jwx/v4/jwk" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/model" cryptoutils "github.com/pocket-id/pocket-id/backend/internal/utils/crypto" @@ -257,6 +259,90 @@ func TestKeyProviderDatabase_SaveKey(t *testing.T) { assert.Equal(t, keyBytes, parsedKeyBytes, "Expected saved key to match original key") }) + + t.Run("SaveKey returns the database error without the retry sentinel", func(t *testing.T) { + db := testutils.NewDatabaseForTest(t) + kek := generateTestKEK(t) + storeErr := errors.New("test database error") + + err = db.Callback().Create().Before("gorm:create").Register("fail_key_storage", func(tx *gorm.DB) { + _ = tx.AddError(storeErr) + }) + require.NoError(t, err) + + provider := &KeyProviderDatabase{} + err = provider.Init(KeyProviderOpts{ + DB: db, + Kek: kek, + }) + require.NoError(t, err) + + err = provider.SaveKey(t.Context(), key) + require.ErrorIs(t, err, storeErr) + require.NotErrorIs(t, err, ErrRetryKeyStorage) + require.ErrorContains(t, err, "failed to store key in database") + }) + + t.Run("SaveKey keeps the first key when the row already exists", func(t *testing.T) { + // Initialize a provider backed by an empty database + db := testutils.NewDatabaseForTest(t) + provider := &KeyProviderDatabase{} + err := provider.Init(KeyProviderOpts{ + DB: db, + Kek: generateTestKEK(t), + }) + require.NoError(t, err) + + // Store the key that should win the conflict + err = provider.SaveKey(t.Context(), key) + require.NoError(t, err) + + // Try storing a second key and verify the caller is told to reload + replacementKey, err := GenerateKey("ES256", "") + require.NoError(t, err) + err = provider.SaveKey(t.Context(), replacementKey) + require.ErrorIs(t, err, ErrRetryKeyStorage) + + // Reload the row and verify the first key was not overwritten + loadedKey, err := provider.LoadKey(t.Context()) + require.NoError(t, err) + loadedKeyBytes, err := EncodeJWKBytes(loadedKey) + require.NoError(t, err) + keyBytes, err := EncodeJWKBytes(key) + require.NoError(t, err) + assert.Equal(t, keyBytes, loadedKeyBytes) + }) +} + +func TestKeyProviderDatabase_ReplaceKey(t *testing.T) { + // Initialize a provider and store the original key + db := testutils.NewDatabaseForTest(t) + provider := &KeyProviderDatabase{} + err := provider.Init(KeyProviderOpts{ + DB: db, + Kek: generateTestKEK(t), + }) + require.NoError(t, err) + + originalKey, err := GenerateKey("ES256", "") + require.NoError(t, err) + err = provider.SaveKey(t.Context(), originalKey) + require.NoError(t, err) + + // Explicitly replace the stored key + replacementKey, err := GenerateKey("ES256", "") + require.NoError(t, err) + err = provider.ReplaceKey(t.Context(), replacementKey) + require.NoError(t, err) + + // Reload the row and verify the replacement was persisted + loadedKey, err := provider.LoadKey(t.Context()) + require.NoError(t, err) + loadedKeyBytes, err := EncodeJWKBytes(loadedKey) + require.NoError(t, err) + replacementKeyBytes, err := EncodeJWKBytes(replacementKey) + require.NoError(t, err) + assert.Equal(t, replacementKeyBytes, loadedKeyBytes) } func TestKeyProviderDatabase_DBKey(t *testing.T) {