Fixed: initial key should be saved as first-write-wins

This commit is contained in:
ItalyPaleAle
2026-09-01 10:29:44 +02:00
parent f893b25e9a
commit 2ae67ce195
7 changed files with 221 additions and 22 deletions

View File

@@ -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)
}

View File

@@ -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)
}

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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)

View File

@@ -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) {