feat: use HS256 for Pocket ID's own session tokens + fixes (#1733)

This commit is contained in:
Alessandro (Ale) Segala
2026-09-07 15:27:08 +02:00
committed by GitHub
parent 8c76ba0714
commit 33090fa3e1
13 changed files with 877 additions and 66 deletions

View File

@@ -93,9 +93,12 @@ func encryptionKeyRotate(ctx context.Context, flags encryptionKeyRotateFlags, db
}
err = db.Transaction(func(tx *gorm.DB) error {
err = rotateSigningKeyEncryption(ctx, tx, oldKek, newKek)
if err != nil {
return err
// Both the key used for tokens meant for external consumption and the session key are encrypted with the KEK
for _, dbKey := range []string{jwkutils.PrivateKeyDBKey, jwkutils.SessionKeyDBKey} {
err = rotateJWKEncryption(ctx, tx, oldKek, newKek, dbKey)
if err != nil {
return err
}
}
err = rotateScimTokens(tx, oldEncKey, newEncKey)
@@ -115,19 +118,20 @@ func encryptionKeyRotate(ctx context.Context, flags encryptionKeyRotateFlags, db
return nil
}
func rotateSigningKeyEncryption(ctx context.Context, db *gorm.DB, oldKek []byte, newKek []byte) error {
func rotateJWKEncryption(ctx context.Context, db *gorm.DB, oldKek []byte, newKek []byte, dbKey string) error {
oldProvider := &jwkutils.KeyProviderDatabase{}
err := oldProvider.Init(jwkutils.KeyProviderOpts{
DB: db,
Kek: oldKek,
DB: db,
Kek: oldKek,
DBKey: dbKey,
})
if err != nil {
return fmt.Errorf("failed to init key provider with old encryption key: %w", err)
return fmt.Errorf("failed to init key provider for %q with old encryption key: %w", dbKey, err)
}
key, err := oldProvider.LoadKey(ctx)
if err != nil {
return fmt.Errorf("failed to load signing key using old encryption key: %w", err)
return fmt.Errorf("failed to load key %q using old encryption key: %w", dbKey, err)
}
if key == nil {
return nil
@@ -135,15 +139,17 @@ func rotateSigningKeyEncryption(ctx context.Context, db *gorm.DB, oldKek []byte,
newProvider := &jwkutils.KeyProviderDatabase{}
err = newProvider.Init(jwkutils.KeyProviderOpts{
DB: db,
Kek: newKek,
DB: db,
Kek: newKek,
DBKey: dbKey,
})
if err != nil {
return fmt.Errorf("failed to init key provider with new encryption key: %w", err)
return fmt.Errorf("failed to init key provider for %q with new encryption key: %w", dbKey, err)
}
if err := newProvider.SaveKey(ctx, key); err != nil {
return fmt.Errorf("failed to store signing key with new encryption key: %w", err)
err = newProvider.ReplaceKey(ctx, key)
if err != nil {
return fmt.Errorf("failed to store key %q with new encryption key: %w", dbKey, err)
}
return nil

View File

@@ -41,6 +41,19 @@ func TestEncryptionKeyRotate(t *testing.T) {
require.NoError(t, err)
require.NoError(t, oldProvider.SaveKey(t.Context(), signingKey))
oldSessionKeyProvider := &jwkutils.KeyProviderDatabase{}
err = oldSessionKeyProvider.Init(jwkutils.KeyProviderOpts{
DB: db,
Kek: oldKek,
DBKey: jwkutils.SessionKeyDBKey,
})
require.NoError(t, err)
sessionKey, err := jwkutils.GenerateSessionKey()
require.NoError(t, err)
err = oldSessionKeyProvider.SaveKey(t.Context(), sessionKey)
require.NoError(t, err)
oldEncKey, err := datatype.DeriveEncryptedStringKey(oldKey)
require.NoError(t, err)
encToken, err := datatype.EncryptEncryptedStringWithKey(oldEncKey, []byte("scim-token-123"))
@@ -76,6 +89,23 @@ func TestEncryptionKeyRotate(t *testing.T) {
require.NoError(t, err)
require.NotNil(t, rotatedKey)
// The session key must be re-encrypted with the new encryption key too, so sessions survive the rotation
newSessionKeyProvider := &jwkutils.KeyProviderDatabase{}
err = newSessionKeyProvider.Init(jwkutils.KeyProviderOpts{
DB: db,
Kek: newKek,
DBKey: jwkutils.SessionKeyDBKey,
})
require.NoError(t, err)
rotatedSessionKey, err := newSessionKeyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, rotatedSessionKey)
sessionKeyID, _ := sessionKey.KeyID()
rotatedSessionKeyID, _ := rotatedSessionKey.KeyID()
assert.Equal(t, sessionKeyID, rotatedSessionKeyID, "The session key should be unchanged, only re-encrypted")
var storedToken string
err = db.Model(&scimsync.ServiceProvider{}).
Where("id = ?", "scim-1").

View File

@@ -19,9 +19,10 @@ import (
)
type keyRotateFlags struct {
Alg string
Crv string
Yes bool
Alg string
Crv string
SessionKey bool
Yes bool
}
func init() {
@@ -29,8 +30,13 @@ func init() {
keyRotateCmd := &cobra.Command{
Use: "key-rotate",
Short: "Generates a new token signing key and replaces the current one",
Short: "Generates a new signing key and replaces the current one",
RunE: func(cmd *cobra.Command, args []string) error {
// The session key is always a symmetric HS256 key, so the algorithm flags don't apply to it
if flags.SessionKey && (cmd.Flags().Changed("alg") || cmd.Flags().Changed("crv")) {
return errors.New("the --alg and --crv flags cannot be used together with --session-key")
}
db, _, err := bootstrap.NewDatabase(cmd.Context())
if err != nil {
return err
@@ -47,12 +53,18 @@ func init() {
keyRotateCmd.Flags().StringVarP(&flags.Alg, "alg", "a", "RS256", "Key algorithm. Supported values: RS256, RS384, RS512, ES256, ES384, ES512, EdDSA")
keyRotateCmd.Flags().StringVarP(&flags.Crv, "crv", "c", "", "Curve name when using EdDSA keys. Supported values: Ed25519")
keyRotateCmd.Flags().BoolVar(&flags.SessionKey, "session-key", false, "Rotate the key used to sign session tokens instead of the token signing key")
keyRotateCmd.Flags().BoolVarP(&flags.Yes, "yes", "y", false, "Do not prompt for confirmation")
rootCmd.AddCommand(keyRotateCmd)
}
func keyRotate(ctx context.Context, flags keyRotateFlags, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema) error {
// The session key is a separate key, generated with a fixed algorithm, so it's rotated on its own
if flags.SessionKey {
return sessionKeyRotate(ctx, flags, db, instanceID, envConfig)
}
// Validate the flags
switch strings.ToUpper(flags.Alg) {
case jwa.RS256().String(), jwa.RS384().String(), jwa.RS512().String(),
@@ -101,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)
}
@@ -111,3 +123,40 @@ func keyRotate(ctx context.Context, flags keyRotateFlags, db *gorm.DB, instanceI
return nil
}
func sessionKeyRotate(ctx context.Context, flags keyRotateFlags, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema) error {
if !flags.Yes {
fmt.Println("WARNING: Rotating the session key will invalidate all existing sessions, and all users will need to sign in again. Tokens issued to client applications are not affected.")
ok, err := utils.PromptForConfirmation("Confirm")
if err != nil {
return err
}
if !ok {
fmt.Println("Aborted")
os.Exit(1)
}
}
// Get the key provider for the session key
keyProvider, err := jwkutils.GetSessionKeyProvider(db, envConfig, instanceID)
if err != nil {
return fmt.Errorf("failed to get session key provider: %w", err)
}
// Generate a new key
key, err := jwkutils.GenerateSessionKey()
if err != nil {
return fmt.Errorf("failed to generate session key: %w", err)
}
// Save the key
err = keyProvider.ReplaceKey(ctx, key)
if err != nil {
return fmt.Errorf("failed to store new session key: %w", err)
}
fmt.Println("Session key rotated successfully")
fmt.Println("Note: if pocket-id is running, you will need to restart it for the new key to be loaded")
return nil
}

View File

@@ -121,6 +121,88 @@ func testKeyRotateWithDatabaseStorage(t *testing.T, flags keyRotateFlags, wantEr
}
}
func TestKeyRotateSessionKey(t *testing.T) {
envConfig := &common.EnvConfigSchema{
EncryptionKey: []byte("test-encryption-key-characters-long"),
}
db := testingutils.NewDatabaseForTest(t)
instanceID, err := instanceid.Load(t.Context(), db)
require.NoError(t, err)
// Seed both keys so we can check that only the session key is replaced
keyProvider, err := jwkutils.GetKeyProvider(db, envConfig, instanceID)
require.NoError(t, err)
signingKey, err := jwkutils.GenerateKey("ES256", "")
require.NoError(t, err)
err = keyProvider.SaveKey(t.Context(), signingKey)
require.NoError(t, err)
sessionKeyProvider, err := jwkutils.GetSessionKeyProvider(db, envConfig, instanceID)
require.NoError(t, err)
originalSessionKey, err := jwkutils.GenerateSessionKey()
require.NoError(t, err)
err = sessionKeyProvider.SaveKey(t.Context(), originalSessionKey)
require.NoError(t, err)
// Rotate the session key
err = keyRotate(t.Context(), keyRotateFlags{SessionKey: true, Yes: true}, db, instanceID, envConfig)
require.NoError(t, err)
// The session key must have been replaced with a new HS256 key
rotatedSessionKey, err := sessionKeyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, rotatedSessionKey)
alg, ok := rotatedSessionKey.Algorithm()
_ = assert.True(t, ok) &&
assert.Equal(t, "HS256", alg.String())
originalKeyID, _ := originalSessionKey.KeyID()
rotatedKeyID, _ := rotatedSessionKey.KeyID()
assert.NotEqual(t, originalKeyID, rotatedKeyID, "Session key should have been replaced")
// The token signing key must be left untouched
unchangedSigningKey, err := keyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, unchangedSigningKey)
signingKeyID, _ := signingKey.KeyID()
unchangedSigningKeyID, _ := unchangedSigningKey.KeyID()
assert.Equal(t, signingKeyID, unchangedSigningKeyID, "Token signing key should not have been rotated")
}
func TestKeyRotateDoesNotChangeSessionKey(t *testing.T) {
envConfig := &common.EnvConfigSchema{
EncryptionKey: []byte("test-encryption-key-characters-long"),
}
db := testingutils.NewDatabaseForTest(t)
instanceID, err := instanceid.Load(t.Context(), db)
require.NoError(t, err)
sessionKeyProvider, err := jwkutils.GetSessionKeyProvider(db, envConfig, instanceID)
require.NoError(t, err)
originalSessionKey, err := jwkutils.GenerateSessionKey()
require.NoError(t, err)
err = sessionKeyProvider.SaveKey(t.Context(), originalSessionKey)
require.NoError(t, err)
// Rotating the token signing key must leave existing sessions valid
err = keyRotate(t.Context(), keyRotateFlags{Alg: "ES256", Yes: true}, db, instanceID, envConfig)
require.NoError(t, err)
sessionKey, err := sessionKeyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, sessionKey)
originalKeyID, _ := originalSessionKey.KeyID()
sessionKeyID, _ := sessionKey.KeyID()
assert.Equal(t, originalKeyID, sessionKeyID, "Session key should not have been rotated")
}
func TestKeyRotateMultipleAlgorithms(t *testing.T) {
algorithms := []struct {
alg string

View File

@@ -534,6 +534,11 @@ func (s *TestService) SeedDatabase(baseURL string) error {
// {"alg":"RS256","d":"mvMDWSdPPvcum0c0iEHE2gbqtV2NKMmLwrl9E6K7g8lTV95SePLnW_bwyMPV7EGp7PQk3l17I5XRhFjze7GqTnFIOgKzMianPs7jv2ELtBMGK0xOPATgu1iGb70xZ6vcvuEfRyY3dJ0zr4jpUdVuXwKmx9rK4IdZn2dFCKfvSuspqIpz11RhF1ALrqDLkxGVv7ZwNh0_VhJZU9hcjG5l6xc7rQEKpPRkZp0IdjkGS8Z0FskoVaiRIWAbZuiVFB9WCW8k1czC4HQTPLpII01bUQx2ludbm0UlXRgVU9ptUUbU7GAImQqTOW8LfPGklEvcgzlIlR_oqw4P9yBxLi-yMQ","dp":"pvNCSnnhbo8Igw9psPR-DicxFnkXlu_ix4gpy6efTrxA-z1VDFDioJ814vKQNioYDzpyAP1gfMPhRkvG_q0hRZsJah3Sb9dfA-WkhSWY7lURQP4yIBTMU0PF_rEATuS7lRciYk1SOx5fqXZd3m_LP0vpBC4Ujlq6NAq6CIjCnms","dq":"TtUVGCCkPNgfOLmkYXu7dxxUCV5kB01-xAEK2OY0n0pG8vfDophH4_D_ZC7nvJ8J9uDhs_3JStexq1lIvaWtG99RNTChIEDzpdn6GH9yaVcb_eB4uJjrNm64FhF8PGCCwxA-xMCZMaARKwhMB2_IOMkxUbWboL3gnhJ2rDO_QO0","e":"AQAB","kid":"8uHDw3M6rf8","kty":"RSA","n":"yaeEL0VKoPBXIAaWXsUgmu05lAvEIIdJn0FX9lHh4JE5UY9B83C5sCNdhs9iSWzpeP11EVjWp8i3Yv2CF7c7u50BXnVBGtxpZpFC-585UXacoJ0chUmarL9GRFJcM1nPHBTFu68aRrn1rIKNHUkNaaxFo0NFGl_4EDDTO8HwawTjwkPoQlRzeByhlvGPVvwgB3Fn93B8QJ_cZhXKxJvjjrC_8Pk76heC_ntEMru71Ix77BoC3j2TuyiN7m9RNBW8BU5q6lKoIdvIeZfTFLzi37iufyfvMrJTixp9zhNB1NxlLCeOZl2MXegtiGqd2H3cbAyqoOiv9ihUWTfXj7SxJw","p":"_Yylc9e07CKdqNRD2EosMC2mrhrEa9j5oY_l00Qyy4-jmCA59Q9viyqvveRo0U7cRvFA5BWgWN6GGLh1DG3X-QBqVr0dnk3uzbobb55RYUXyPLuBZI2q6w2oasbiDwPdY7KpkVv_H-bpITQlyDvO8hhucA6rUV7F6KTQVz8M3Ms","q":"y5p3hch-7jJ21TkAhp_Vk1fLCAuD4tbErwQs2of9ja8sB4iJOs5Wn6HD3P7Mc8Plye7qaLHvzc8I5g0tPKWvC0DPd_FLPXiWwMVAzee3NUX_oGeJNOQp11y1w_KqdO9qZqHSEPZ3NcFL_SZMFgggxhM1uzRiPzsVN0lnD_6prZU","qi":"2Grt6uXHm61ji3xSdkBWNtUnj19vS1-7rFJp5SoYztVQVThf_W52BAiXKBdYZDRVoItC_VS2NvAOjeJjhYO_xQ_q3hK7MdtuXfEPpLnyXKkmWo3lrJ26wbeF6l05LexCkI7ShsOuSt-dsyaTJTszuKDIA6YOfWvfo3aVZmlWRaI","use":"sig"}
Value: new("7d/5hl7diJ2rnFL14hEAQf9tzpu29aqXQ8jpJ2iqqKUNFZpdOkEpud0CmRv4H3r8yyk2u/Gqqj9klSy58DJkYXGF5PAYgLyoBIb7L3JXWRbxg4cQ3QJCug13l2OTmpAKoVc+rmX8c3j3h1sNqyJ+7Ql5sS0jSeyiYgIsFNCdnK5alBDyvtcpe/QDpklmP4JCeVpvmf2rLGplk3g5UO5ydJ8UiDXxfDmi+gF6NKJvrGnnah8Ar3G/x88z+tTJtp0DIQFwxXwUM2XZqzEVGm8K2r0w5o9/Keh6bBBaiuH2C78ZOaijGV3DovhR+e9J0cYUYGwT42MZMx9fSWQ/lvWGGnf+Uq3MXJfjWSREfhkp8KTQwR9F7+dnVJWswOEk7jPR8I7hCWTMxJyvaFX3wgAXIVmhrgXZQQbYOqTt56IoqUl0xOJku8dA8opg2UcLlmmuOh6+hfkXKsiiS/H/9c1BVIGj1fCOiT6IePh4wKKSTbwJnPD5EKmdJpgTsUpjcDnXQKY4ReO0UpdRdKxwRDDLeQuG6j+ljGxR9GPudCU9Nmci6rFVI6n5LWYkQxBA1O73RpmXRZPDzntDfpXMEonkmSvOoxaCK2Id7CRKMdqvR0kEouwnhk5WSFtsfi3sA0pkXzPFxwZeWM8vFtbffZOZzXaOhxCOfcj1NClZohlZhyc4jvkxmrpY7PSaAzih0AmHI7y0LYFi6fZu/K4EheVa1+KF55nWZ8ARikHMWKAKkyExkTak7xyN884TDmzURRaPlQg4jzQte5WMNjAG/hlHibdMBNvgwiYd49ZxteJ8ABdbiXVRl+2JGbdjl2ubpQZwOn7bJKlqO56bIwsZ+e4+pXsuOGdBahkHrUjtMEmH3DZbGc6CJLbcmdhdpApLQRRcLAazxJhzAwJ47FRYsHsj57LnYNvmcKdIxw8rxCdLUuzz95uw0T3ankEO5J9sjem+HMEuKdwXK1UcuOn2rjR8Sd/BuvQmeso27dFbPXqXYNS90Ml45YyTvcKSiopD181oZR703TFUSpR7dsiqROMr+p/2jN9h6a8WbQ8xpksyclaQByY/M77AssbXnG6wfhRsntNIINCZLbBnjXOyz6ZHIC5K4tSTdcnWaiYPeRPQmnw9UUvHAcNU2yMWsy0eU377yDS0WstTxOdQutTdkczl8kv5Lo26JiEK7mSIuRK19ffF9Zz8FG8+eKv5zdyIPjyQRDYBysUoDv5huKe2eoxJu/MWS2Pql/ZtUGeD6Ozm3mCvh0vQ9ceagBkY6Ocm3du0ziAKP29Ri0mjg4DizVorbLzsh+EQH/s2Pi9MnjUZDlEmuLl2Xfp7/w4j/8u0N0tVR70VDFuGdKpTjFY3vS8EJrPtyMTM51x1D9rb8gIql8aR/rJw4YF+huxg1mv5n6+tGVqg5msbPmF12eJijP4lkmaRwIpLW5pJTtaDkUj7uOeu1mm4k+Dt5nh0/0jPHzrv6bcTCcbV7UjMHDoTXXqEpFAAJ66rHR7zdAJu+YKsnTIZyLmOpcowq7LL8G9qTvV0OSpyQWUIavRSgbDHFqEqRs+JU94jAzkq8nCY5MTd9m5sIv9InfdT3k+pwpsE/FKge8nghFLtbUrafGkzTky8SE2druvVcIvbfXMfLIKRUYjJgnWc0gQzF5J6pzXM7D2r/RG6JDzASqjlbURq6v9bhNerlOVdMujWKEEVcKWIzlbt4RkihRjM8AUqIZQOyicGQ+4yfIjAHw5viuABONYs3OIWULnFqJxdvS9rNKhfxSjIq9cfqyzevq2xrRoMXEonobh6M3bD2Vang8OAeVeD1OXWPERi4pepCYFS9RJ/Xa/UWxptsqSNuGcb3fAzQSmLpXLGdWRoKXvSe7EYgc0bGcLOjSTu5RURKo+EF9i4KT9EJauf6VXw5dTf/CCIJRXE1bWzXhSCFYntohYhX2ldOCDYpi/jFBC6Vtkw0ud3/xq8Nmhd5gUk+SpngByCZH3Pm3H+jvlbMpiqkDkm1v74hDX13Xhrcw2eWyuqKBVoRCCniUvwpYNbGvBfjC6Hcizv0Aybciwj+4nybt5EPoEUm6S6Gs7fG7QpPdvrzpAxX70MlmdkF/gwyuhbEeJhLK+WL7qAsN5CvHPzVbsIf90x+nGTtMJPgpxVr0tJMj+vprXV4WxutfARBiOnqe58MhA857sd+MzKBgKnoLOBRTiC3qc/0/ULwbG2HCCD7nmwzz7M4nUuMvo8rgS7z0BF68OClT8X3JwSXbL5Wg=="),
},
{
Key: jwkutils.SessionKeyDBKey,
// {"alg":"HS256","k":"5un_Rh6BPDVwGwRWPC_-w-HvT4BuUq5vYE4a2z4IL1k","kid":"YC7IX6YEFJc","kty":"oct","use":"sig"}
Value: new("6puTIBpn0u2Y8FJQ9C4gxuzORTgMmac9Tz9B2epw212hlaepET06ca/CPnwdirCNNg/tLG1wXd2MNSEMgZAnl1cPkF9hPabrRW+SUYpFDLu4yE9w5uc6ns//9pphedK5vS190oXcE7FaWBoso789JuQ0yoicNEnjAAjBWExnp0dXXufmXzUZSnKQ"),
},
}
for _, kv := range keyValues {

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"
)
@@ -43,9 +45,12 @@ const (
)
type JwtService struct {
db *gorm.DB
envConfig *common.EnvConfigSchema
privateKey jwk.Key
db *gorm.DB
envConfig *common.EnvConfigSchema
// privateKey signs tokens that are consumed externally, such as ID tokens and access tokens for apps
privateKey jwk.Key
// sessionKey is the symmetric key that signs Pocket ID's own session tokens
sessionKey jwk.Key
keyId string
instanceID string
jwksEncoded []byte
@@ -71,7 +76,46 @@ func (s *JwtService) init(ctx context.Context, db *gorm.DB, instanceID string, e
return s.LoadOrGenerateKey(ctx)
}
// 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 := 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 = retryKeyStorage(ctx, s.loadOrGenerateSessionKey)
if err != nil {
return fmt.Errorf("error loading session key: %w", err)
}
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)
if err != nil {
@@ -108,6 +152,49 @@ func (s *JwtService) LoadOrGenerateKey(ctx context.Context) error {
return nil
}
func (s *JwtService) loadOrGenerateSessionKey(ctx context.Context) error {
// Get the key provider for the session key
keyProvider, err := jwkutils.GetSessionKeyProvider(s.db, s.envConfig, s.instanceID)
if err != nil {
return fmt.Errorf("failed to get session key provider: %w", err)
}
// Try loading a key
key, err := keyProvider.LoadKey(ctx)
if err != nil {
return fmt.Errorf("failed to load session key: %w", err)
}
// If we have a key, store it in the object and we're done
if key != nil {
err = s.SetSessionKey(key)
if err != nil {
return fmt.Errorf("failed to set session key: %w", err)
}
return nil
}
// If we are here, we need to generate a new key
key, err = jwkutils.GenerateSessionKey()
if err != nil {
return fmt.Errorf("failed to generate session key: %w", err)
}
// Set the key in the object, which also validates it
err = s.SetSessionKey(key)
if err != nil {
return fmt.Errorf("failed to set session key: %w", err)
}
// Save the newly-generated key
err = keyProvider.SaveKey(ctx, key)
if err != nil {
return fmt.Errorf("failed to save session key: %w", err)
}
return nil
}
// generateKey generates a new key and stores it in the object
func (s *JwtService) generateKey() error {
// Default is to generate RS256 (RSA-2048) keys
@@ -147,6 +234,34 @@ func ValidateKey(privateKey jwk.Key) error {
return nil
}
// ValidateSessionKey validates the symmetric key used to sign session tokens
func ValidateSessionKey(sessionKey jwk.Key) error {
// Validate the loaded key
err := sessionKey.Validate()
if err != nil {
return fmt.Errorf("key object is invalid: %w", err)
}
if sessionKey.KeyType() != jwa.OctetSeq() {
return errors.New("key object is not a symmetric key")
}
keyID, ok := sessionKey.KeyID()
if !ok || keyID == "" {
return errors.New("key object does not contain a key ID")
}
usage, ok := sessionKey.KeyUsage()
if !ok || usage != KeyUsageSigning {
return errors.New("key object is not valid for signing")
}
// Session tokens are always signed with the same algorithm, so a key for anything else can't be used
alg, ok := sessionKey.Algorithm()
if !ok || alg == nil || alg.String() != jwkutils.SessionKeyAlg().String() {
return fmt.Errorf("key object is not valid for the %s algorithm", jwkutils.SessionKeyAlg().String())
}
return nil
}
func (s *JwtService) SetKey(privateKey jwk.Key) error {
// Validate the loaded key
err := ValidateKey(privateKey)
@@ -181,7 +296,24 @@ func (s *JwtService) SetKey(privateKey jwk.Key) error {
return nil
}
// SetSessionKey sets the symmetric key used to sign session tokens
// This key is never published in the JWKS, since it's a shared secret that only Pocket ID needs
func (s *JwtService) SetSessionKey(sessionKey jwk.Key) error {
err := ValidateSessionKey(sessionKey)
if err != nil {
return fmt.Errorf("session key is not valid: %w", err)
}
s.sessionKey = sessionKey
return nil
}
func (s *JwtService) GenerateAccessToken(user model.User, authenticationMethod string, sessionDuration time.Duration) (string, error) {
if s.sessionKey == nil {
return "", errors.New("session key is not initialized")
}
now := time.Now()
token, err := jwt.NewBuilder().
Subject(user.ID).
@@ -214,8 +346,8 @@ func (s *JwtService) GenerateAccessToken(user model.User, authenticationMethod s
return "", fmt.Errorf("failed to set '%s' claim in token: %w", common.AuthenticationMethodsClaim, err)
}
alg, _ := s.privateKey.Algorithm()
signed, err := jwt.Sign(token, jwt.WithKey(alg, s.privateKey))
// Session tokens are signed with the symmetric session key
signed, err := jwt.Sign(token, jwt.WithKey(jwkutils.SessionKeyAlg(), s.sessionKey))
if err != nil {
return "", fmt.Errorf("failed to sign token: %w", err)
}
@@ -224,11 +356,14 @@ func (s *JwtService) GenerateAccessToken(user model.User, authenticationMethod s
}
func (s *JwtService) VerifyAccessToken(tokenString string) (jwt.Token, error) {
alg, _ := s.privateKey.Algorithm()
if s.sessionKey == nil {
return nil, errors.New("session key is not initialized")
}
token, err := jwt.ParseString(
tokenString,
jwt.WithValidate(true),
jwt.WithKey(alg, s.privateKey),
jwt.WithKey(jwkutils.SessionKeyAlg(), s.sessionKey),
jwt.WithAcceptableSkew(clockSkew),
jwt.WithAudience(s.envConfig.AppURL),
jwt.WithIssuer(s.envConfig.AppURL),

View File

@@ -1,17 +1,20 @@
package service
import (
"context"
"crypto/ecdsa"
"crypto/ed25519"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"errors"
"sync"
"testing"
"time"
"github.com/lestrrat-go/jwx/v4/jwa"
"github.com/lestrrat-go/jwx/v4/jwk"
"github.com/lestrrat-go/jwx/v4/jws"
"github.com/lestrrat-go/jwx/v4/jwt"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -71,13 +74,13 @@ func newTestDbAndEnv(t *testing.T) (*gorm.DB, *common.EnvConfigSchema) {
return testutils.NewDatabaseForTest(t), newTestEnvConfig()
}
func saveKeyToDatabase(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, appConfig *appconfig.AppConfigService, key jwk.Key) string {
func saveKeyToDatabase(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, key jwk.Key) string {
t.Helper()
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()
@@ -145,7 +148,7 @@ func TestJwtService_Init(t *testing.T) {
instanceID := newInstanceID(t, db)
// Create a new JWK and save it to the database
origKeyID := createECDSAKeyJWK(t, db, instanceID, mockEnvConfig, mockConfig)
origKeyID := createECDSAKeyJWK(t, db, instanceID, mockEnvConfig)
// Now create a new service that should load the existing key
svc := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
@@ -167,7 +170,7 @@ func TestJwtService_Init(t *testing.T) {
instanceID := newInstanceID(t, db)
// Create a new JWK and save it to the database
origKeyID := createEdDSAKeyJWK(t, db, instanceID, mockEnvConfig, mockConfig)
origKeyID := createEdDSAKeyJWK(t, db, instanceID, mockEnvConfig)
// Now create a new service that should load the existing key
svc := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
@@ -187,6 +190,207 @@ 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, 4, attempts)
}
func TestJwtService_SessionKey(t *testing.T) {
mockConfig := appconfig.NewTestAppConfigService(nil)
t.Run("should generate a new session key when none exists", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
mockEnvConfig := newTestEnvConfig()
instanceID := newInstanceID(t, db)
// Initialize the JWT service
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
// Verify the session key was set and is a symmetric HS256 key
require.NotNil(t, service.sessionKey, "Session key should be set")
assert.Equal(t, jwa.OctetSeq(), service.sessionKey.KeyType(), "Session key should be a symmetric key")
alg, ok := service.sessionKey.Algorithm()
_ = assert.True(t, ok, "Session key should have an algorithm") &&
assert.Equal(t, jwa.HS256().String(), alg.String(), "Session key should use HS256")
// Verify the session key has been persisted in its own row in the database
keyProvider, err := jwkutils.GetSessionKeyProvider(db, mockEnvConfig, instanceID)
require.NoError(t, err, "Failed to init session key provider")
key, err := keyProvider.LoadKey(t.Context())
require.NoError(t, err, "Failed to load session key from provider")
require.NotNil(t, key, "Session key should be present in the database")
keyID, ok := key.KeyID()
_ = assert.True(t, ok, "Session key should have a key ID") &&
assert.NotEmpty(t, keyID)
})
t.Run("should load an existing session key", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
mockEnvConfig := newTestEnvConfig()
instanceID := newInstanceID(t, db)
// First create a service to generate a session key
firstService := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
origKeyID, ok := firstService.sessionKey.KeyID()
require.True(t, ok)
// Now create a new service that should load the existing session key
secondService := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
loadedKeyID, ok := secondService.sessionKey.KeyID()
require.True(t, ok)
assert.Equal(t, origKeyID, loadedKeyID, "Loaded session key should have the same ID as the original")
// A session token issued by the first service must be accepted by the second one
tokenString, err := firstService.GenerateAccessToken(model.User{Base: model.Base{ID: "user123"}}, "", time.Hour)
require.NoError(t, err)
_, err = secondService.VerifyAccessToken(tokenString)
require.NoError(t, err, "Session token should be verified by a service that loaded the same session key")
})
t.Run("session key is separate from the token signing key", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
mockEnvConfig := newTestEnvConfig()
instanceID := newInstanceID(t, db)
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
signingKeyID, ok := service.privateKey.KeyID()
require.True(t, ok)
sessionKeyID, ok := service.sessionKey.KeyID()
require.True(t, ok)
assert.NotEqual(t, signingKeyID, sessionKeyID, "Session key and token signing key should be different keys")
// The session key is a shared secret, so it must never be published in the JWKS
jwks, err := service.GetPublicJWKSAsJSON()
require.NoError(t, err)
assert.NotContains(t, string(jwks), sessionKeyID, "Session key must not be included in the JWKS")
assert.NotContains(t, string(jwks), jwa.OctetSeq().String(), "JWKS must not contain symmetric keys")
})
t.Run("session tokens are signed with HS256 and the session key", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
mockEnvConfig := newTestEnvConfig()
instanceID := newInstanceID(t, db)
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
tokenString, err := service.GenerateAccessToken(model.User{Base: model.Base{ID: "user123"}}, "", time.Hour)
require.NoError(t, err)
// Inspect the JWS header to confirm the algorithm and key used
msg, err := jws.ParseString(tokenString)
require.NoError(t, err)
require.Len(t, msg.Signatures(), 1)
headers := msg.Signatures()[0].ProtectedHeaders()
headerAlg, ok := headers.Algorithm()
_ = assert.True(t, ok, "Session token should declare an algorithm") &&
assert.Equal(t, jwa.HS256().String(), headerAlg.String(), "Session token should be signed with HS256")
sessionKeyID, _ := service.sessionKey.KeyID()
kid, ok := headers.KeyID()
_ = assert.True(t, ok, "Session token should reference a key ID") &&
assert.Equal(t, sessionKeyID, kid, "Session token should be signed with the session key")
})
t.Run("session tokens signed with a different session key are rejected", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
mockEnvConfig := newTestEnvConfig()
instanceID := newInstanceID(t, db)
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
tokenString, err := service.GenerateAccessToken(model.User{Base: model.Base{ID: "user123"}}, "", time.Hour)
require.NoError(t, err)
// Rotate the session key, as the key-rotate command does, then reload it
rotatedKey, err := jwkutils.GenerateSessionKey()
require.NoError(t, err)
keyProvider, err := jwkutils.GetSessionKeyProvider(db, mockEnvConfig, instanceID)
require.NoError(t, err)
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
_, err = service.VerifyAccessToken(tokenString)
require.Error(t, err, "Session token signed with the previous session key should be rejected")
})
t.Run("rejects an invalid session key", func(t *testing.T) {
service := &JwtService{}
// A key for tokens meant for external consumption is not valid as a session key
signingKey, err := jwkutils.GenerateKey(jwa.ES256().String(), "")
require.NoError(t, err)
err = service.SetSessionKey(signingKey)
require.Error(t, err, "An asymmetric key should not be accepted as a session key")
require.ErrorContains(t, err, "not a symmetric key")
// A symmetric key for another algorithm is not valid either
rawKey := make([]byte, 32)
_, err = rand.Read(rawKey)
require.NoError(t, err)
otherAlgKey, err := jwkutils.ImportRawKey(rawKey, jwa.HS512().String(), "")
require.NoError(t, err)
err = service.SetSessionKey(otherAlgKey)
require.Error(t, err, "A key for another algorithm should not be accepted as a session key")
require.ErrorContains(t, err, "not valid for the HS256 algorithm")
})
t.Run("returns an error when the session key is not initialized", func(t *testing.T) {
service := &JwtService{}
_, err := service.GenerateAccessToken(model.User{Base: model.Base{ID: "user123"}}, "", time.Hour)
require.Error(t, err)
require.ErrorContains(t, err, "session key is not initialized")
_, err = service.VerifyAccessToken("some-token")
require.Error(t, err)
require.ErrorContains(t, err, "session key is not initialized")
})
}
func TestJwtService_GetPublicJWK(t *testing.T) {
@@ -222,7 +426,7 @@ func TestJwtService_GetPublicJWK(t *testing.T) {
t.Run("returns public key when ECDSA private key is initialized", func(t *testing.T) {
// Create an ECDSA key and save it in the database
originalKeyID := createECDSAKeyJWK(t, db, instanceID, mockEnvConfig, mockConfig)
originalKeyID := createECDSAKeyJWK(t, db, instanceID, mockEnvConfig)
// Create a JWT service that loads the ECDSA key
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
@@ -258,7 +462,7 @@ func TestJwtService_GetPublicJWK(t *testing.T) {
mockEnvConfig := newTestEnvConfig()
// Create an EdDSA key and save it in the database
originalKeyID := createEdDSAKeyJWK(t, db, instanceID, mockEnvConfig, mockConfig)
originalKeyID := createEdDSAKeyJWK(t, db, instanceID, mockEnvConfig)
// Create a JWT service that loads the EdDSA key
service := initJwtService(t, db, instanceID, mockConfig, mockEnvConfig)
@@ -398,7 +602,7 @@ func TestGenerateVerifyAccessToken(t *testing.T) {
})
t.Run("works with Ed25519 keys", func(t *testing.T) {
origKeyID := createEdDSAKeyJWK(t, db, instanceID, envConfig, mockConfig)
origKeyID := createEdDSAKeyJWK(t, db, instanceID, envConfig)
service := initJwtService(t, db, instanceID, mockConfig, envConfig)
loadedKeyID, ok := service.privateKey.KeyID()
@@ -437,7 +641,7 @@ func TestGenerateVerifyAccessToken(t *testing.T) {
})
t.Run("works with P-256 keys", func(t *testing.T) {
origKeyID := createECDSAKeyJWK(t, db, instanceID, envConfig, mockConfig)
origKeyID := createECDSAKeyJWK(t, db, instanceID, envConfig)
service := initJwtService(t, db, instanceID, mockConfig, envConfig)
loadedKeyID, ok := service.privateKey.KeyID()
@@ -476,7 +680,7 @@ func TestGenerateVerifyAccessToken(t *testing.T) {
})
t.Run("works with RSA-4096 keys", func(t *testing.T) {
origKeyID := createRSA4096KeyJWK(t, db, instanceID, envConfig, mockConfig)
origKeyID := createRSA4096KeyJWK(t, db, instanceID, envConfig)
service := initJwtService(t, db, instanceID, mockConfig, envConfig)
loadedKeyID, ok := service.privateKey.KeyID()
@@ -559,13 +763,13 @@ func TestTokenTypeValidator(t *testing.T) {
})
}
func importKey(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, appConfig *appconfig.AppConfigService, privateKeyRaw any) string {
func importKey(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, privateKeyRaw any) string {
t.Helper()
privateKey, err := jwkutils.ImportRawKey(privateKeyRaw, "", "")
require.NoError(t, err, "Failed to import private key")
return saveKeyToDatabase(t, db, instanceID, envConfig, appConfig, privateKey)
return saveKeyToDatabase(t, db, instanceID, envConfig, privateKey)
}
// Because generating a RSA-406 key isn't immediate, we pre-compute one
@@ -574,7 +778,7 @@ var (
rsaKeyPrecomputeOnce sync.Once
)
func createRSA4096KeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, appConfig *appconfig.AppConfigService) string {
func createRSA4096KeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema) string {
t.Helper()
rsaKeyPrecomputeOnce.Do(func() {
@@ -586,10 +790,10 @@ func createRSA4096KeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig
})
// Import as JWK and save it
return importKey(t, db, instanceID, envConfig, appConfig, rsaKeyPrecomputed)
return importKey(t, db, instanceID, envConfig, rsaKeyPrecomputed)
}
func createECDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, appConfig *appconfig.AppConfigService) string {
func createECDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema) string {
t.Helper()
// Generate a new P-256 ECDSA key
@@ -597,11 +801,11 @@ func createECDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *
require.NoError(t, err, "Failed to generate ECDSA key")
// Import as JWK and save it
return importKey(t, db, instanceID, envConfig, appConfig, privateKeyRaw)
return importKey(t, db, instanceID, envConfig, privateKeyRaw)
}
// Helper function to create an Ed25519 key and save it as JWK
func createEdDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema, appConfig *appconfig.AppConfigService) string {
func createEdDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *common.EnvConfigSchema) string {
t.Helper()
// Generate a new Ed25519 key pair
@@ -609,5 +813,5 @@ func createEdDSAKeyJWK(t *testing.T, db *gorm.DB, instanceID string, envConfig *
require.NoError(t, err, "Failed to generate Ed25519 key")
// Import as JWK and save it
return importKey(t, db, instanceID, envConfig, appConfig, privateKeyRaw)
return importKey(t, db, instanceID, envConfig, privateKeyRaw)
}

View File

@@ -14,15 +14,28 @@ type KeyProviderOpts struct {
EnvConfig *common.EnvConfigSchema
DB *gorm.DB
Kek []byte
DBKey string
}
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
}
func GetKeyProvider(db *gorm.DB, envConfig *common.EnvConfigSchema, instanceID string) (keyProvider KeyProvider, err 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
func GetKeyProvider(db *gorm.DB, envConfig *common.EnvConfigSchema, instanceID string) (KeyProvider, error) {
return getKeyProvider(db, envConfig, instanceID, PrivateKeyDBKey)
}
// GetSessionKeyProvider returns the provider for the symmetric key used to sign session tokens
// This key is symmetric and separate from the one used to sign tokens for external consumption
func GetSessionKeyProvider(db *gorm.DB, envConfig *common.EnvConfigSchema, instanceID string) (KeyProvider, error) {
return getKeyProvider(db, envConfig, instanceID, SessionKeyDBKey)
}
func getKeyProvider(db *gorm.DB, envConfig *common.EnvConfigSchema, instanceID string, dbKey string) (keyProvider KeyProvider, err error) {
// Load the encryption key (KEK) if present
kek, err := LoadKeyEncryptionKey(envConfig, instanceID)
if err != nil {
@@ -34,6 +47,7 @@ func GetKeyProvider(db *gorm.DB, envConfig *common.EnvConfigSchema, instanceID s
DB: db,
EnvConfig: envConfig,
Kek: kek,
DBKey: dbKey,
})
if err != nil {
return nil, fmt.Errorf("failed to init key provider: %w", err)

View File

@@ -15,11 +15,21 @@ import (
cryptoutils "github.com/pocket-id/pocket-id/backend/internal/utils/crypto"
)
const PrivateKeyDBKey = "jwt_private_key.json"
const (
// PrivateKeyDBKey is the row in the "kv" table containing the key used to sign tokens that are consumed externally
PrivateKeyDBKey = "jwt_private_key.json"
// SessionKeyDBKey is the row in the "kv" table containing the symmetric key used to sign session tokens
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
db *gorm.DB
kek []byte
dbKey string
}
func (f *KeyProviderDatabase) Init(opts KeyProviderOpts) error {
@@ -30,12 +40,18 @@ func (f *KeyProviderDatabase) Init(opts KeyProviderOpts) error {
f.db = opts.DB
f.kek = opts.Kek
// Callers that don't ask for a specific row get the token signing key, which is the key most of the codebase deals with
f.dbKey = opts.DBKey
if f.dbKey == "" {
f.dbKey = PrivateKeyDBKey
}
return nil
}
func (f *KeyProviderDatabase) LoadKey(ctx context.Context) (key jwk.Key, err error) {
row := model.KV{
Key: PrivateKeyDBKey,
Key: f.dbKey,
}
ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
@@ -45,7 +61,7 @@ func (f *KeyProviderDatabase) LoadKey(ctx context.Context) (key jwk.Key, err err
// Key not present in the database - return nil so a new one can be generated
return nil, nil
} else if err != nil {
return nil, fmt.Errorf("failed to retrieve private key from the database: %w", err)
return nil, fmt.Errorf("failed to retrieve key from the database: %w", err)
}
if row.Value == nil || *row.Value == "" {
@@ -56,42 +72,62 @@ func (f *KeyProviderDatabase) LoadKey(ctx context.Context) (key jwk.Key, err err
// Decode from base64
enc, err := base64.StdEncoding.DecodeString(*row.Value)
if err != nil {
return nil, fmt.Errorf("failed to read encrypted private key: not a valid base64-encoded value: %w", err)
return nil, fmt.Errorf("failed to read encrypted key: not a valid base64-encoded value: %w", err)
}
// Decrypt the data
data, err := cryptoutils.Decrypt(f.kek, enc, nil)
if err != nil {
return nil, fmt.Errorf("failed to decrypt private key: %w", err)
return nil, fmt.Errorf("failed to decrypt key: %w", err)
}
// Parse the key
key, err = jwk.ParseKey(data)
if err != nil {
return nil, fmt.Errorf("failed to parse encrypted private key: %w", err)
return nil, fmt.Errorf("failed to parse encrypted key: %w", err)
}
return key, nil
}
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: PrivateKeyDBKey,
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.
@@ -103,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 private key in the database at the same time
// In this case, only one of the replicas will succeed; 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 private 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,164 @@ 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) {
t.Run("keys stored under different rows do not overwrite each other", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
kek := generateTestKEK(t)
signingKeyProvider := &KeyProviderDatabase{}
err := signingKeyProvider.Init(KeyProviderOpts{
DB: db,
Kek: kek,
DBKey: PrivateKeyDBKey,
})
require.NoError(t, err)
sessionKeyProvider := &KeyProviderDatabase{}
err = sessionKeyProvider.Init(KeyProviderOpts{
DB: db,
Kek: kek,
DBKey: SessionKeyDBKey,
})
require.NoError(t, err)
// Save a different key with each provider
signingKey, err := GenerateKey("ES256", "")
require.NoError(t, err)
err = signingKeyProvider.SaveKey(t.Context(), signingKey)
require.NoError(t, err)
sessionKey, err := GenerateSessionKey()
require.NoError(t, err)
err = sessionKeyProvider.SaveKey(t.Context(), sessionKey)
require.NoError(t, err)
// Each provider must load back the key it saved
loadedSigningKey, err := signingKeyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, loadedSigningKey)
signingKid, _ := signingKey.KeyID()
loadedSigningKid, _ := loadedSigningKey.KeyID()
assert.Equal(t, signingKid, loadedSigningKid)
loadedSessionKey, err := sessionKeyProvider.LoadKey(t.Context())
require.NoError(t, err)
require.NotNil(t, loadedSessionKey)
sessionKid, _ := sessionKey.KeyID()
loadedSessionKid, _ := loadedSessionKey.KeyID()
assert.Equal(t, sessionKid, loadedSessionKid)
// Both rows must exist in the database
var count int64
err = db.Model(&model.KV{}).Where("key IN ?", []string{PrivateKeyDBKey, SessionKeyDBKey}).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(2), count)
})
t.Run("defaults to the token signing key row when not set", func(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
kek := generateTestKEK(t)
provider := &KeyProviderDatabase{}
require.NoError(t, provider.Init(KeyProviderOpts{
DB: db,
Kek: kek,
}))
key, err := GenerateKey("ES256", "")
require.NoError(t, err)
require.NoError(t, provider.SaveKey(t.Context(), key))
var kv model.KV
err = db.Where("key = ?", PrivateKeyDBKey).First(&kv).Error
require.NoError(t, err, "Expected the key to be stored in the token signing key row")
})
}
func generateTestKEK(t *testing.T) []byte {

View File

@@ -7,6 +7,7 @@ import (
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
@@ -25,6 +26,12 @@ const (
KeyUsageSigning = "sig"
)
// SessionKeyAlg returns the algorithm used to sign session tokens
func SessionKeyAlg() jwa.SignatureAlgorithm {
// Session tokens are only consumed by Pocket ID itself, so use a faster symmetric key that's kept private
return jwa.HS256()
}
// EncodeJWK encodes a jwk.Key to a writable stream.
func EncodeJWK(w io.Writer, key jwk.Key) error {
enc := json.NewEncoder(w)
@@ -124,6 +131,9 @@ func EnsureAlgInKey(key jwk.Key, alg string, crv string) {
// Default to EdDSA and Ed25519 for OKP keys
_ = key.Set(jwk.AlgorithmKey, jwa.EdDSA())
_ = key.Set(jwk.OKPCrvKey, jwa.Ed25519())
case jwa.OctetSeq():
// Default to HS256 for symmetric keys
_ = key.Set(jwk.AlgorithmKey, jwa.HS256())
}
}
@@ -160,3 +170,16 @@ func GenerateKey(alg string, crv string) (key jwk.Key, err error) {
// Import the raw key
return ImportRawKey(rawKey, alg, crv)
}
// GenerateSessionKey generates a new symmetric key used to sign session tokens
func GenerateSessionKey() (jwk.Key, error) {
// Use HS256, which is based on SHA256
rawKey := make([]byte, sha256.Size)
_, err := io.ReadFull(rand.Reader, rawKey)
if err != nil {
return nil, fmt.Errorf("failed to generate session key: %w", err)
}
// Import the raw key
return ImportRawKey(rawKey, SessionKeyAlg().String(), "")
}

View File

@@ -6,6 +6,7 @@ import (
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/hex"
"testing"
@@ -154,6 +155,40 @@ func TestGenerateKey(t *testing.T) {
}
}
func TestGenerateSessionKey(t *testing.T) {
key, err := GenerateSessionKey()
require.NoError(t, err)
require.NotNil(t, key)
// The session key must be a symmetric key for HS256
assert.Equal(t, jwa.OctetSeq(), key.KeyType())
alg, ok := key.Algorithm()
_ = assert.True(t, ok, "algorithm should be set in the key") &&
assert.Equal(t, jwa.HS256().String(), alg.String())
// Verify other required fields are set
kid, ok := key.KeyID()
_ = assert.True(t, ok, "key ID should be set") &&
assert.NotEmpty(t, kid, "key ID should not be empty")
usage, ok := key.KeyUsage()
_ = assert.True(t, ok, "key usage should be set") &&
assert.Equal(t, KeyUsageSigning, usage)
// Verify the key material has the expected length
rawKey, err := jwk.Export[[]byte](key)
require.NoError(t, err)
assert.Len(t, rawKey, sha256.Size)
// Each invocation must return a different key
otherKey, err := GenerateSessionKey()
require.NoError(t, err)
otherRawKey, err := jwk.Export[[]byte](otherKey)
require.NoError(t, err)
assert.NotEqual(t, rawKey, otherRawKey, "each generated session key should be different")
}
func TestEnsureAlgInKey(t *testing.T) {
// Generate an RSA-2048 key
rsaKey, err := rsa.GenerateKey(rand.Reader, 2048)
@@ -282,6 +317,16 @@ func TestEnsureAlgInKey(t *testing.T) {
expectedAlg: jwa.EdDSA(),
expectedCrv: jwa.Ed25519().String(),
},
{
name: "Symmetric key defaults to HS256",
keyGen: func() (any, error) {
rawKey := make([]byte, sha256.Size)
_, err := rand.Read(rawKey)
return rawKey, err
},
expectedAlg: jwa.HS256(),
expectedCrv: "",
},
}
for _, tt := range tests {

View File

@@ -99,6 +99,10 @@
{
"key": "jwt_private_key.json",
"value": "7d/5hl7diJ2rnFL14hEAQf9tzpu29aqXQ8jpJ2iqqKUNFZpdOkEpud0CmRv4H3r8yyk2u/Gqqj9klSy58DJkYXGF5PAYgLyoBIb7L3JXWRbxg4cQ3QJCug13l2OTmpAKoVc+rmX8c3j3h1sNqyJ+7Ql5sS0jSeyiYgIsFNCdnK5alBDyvtcpe/QDpklmP4JCeVpvmf2rLGplk3g5UO5ydJ8UiDXxfDmi+gF6NKJvrGnnah8Ar3G/x88z+tTJtp0DIQFwxXwUM2XZqzEVGm8K2r0w5o9/Keh6bBBaiuH2C78ZOaijGV3DovhR+e9J0cYUYGwT42MZMx9fSWQ/lvWGGnf+Uq3MXJfjWSREfhkp8KTQwR9F7+dnVJWswOEk7jPR8I7hCWTMxJyvaFX3wgAXIVmhrgXZQQbYOqTt56IoqUl0xOJku8dA8opg2UcLlmmuOh6+hfkXKsiiS/H/9c1BVIGj1fCOiT6IePh4wKKSTbwJnPD5EKmdJpgTsUpjcDnXQKY4ReO0UpdRdKxwRDDLeQuG6j+ljGxR9GPudCU9Nmci6rFVI6n5LWYkQxBA1O73RpmXRZPDzntDfpXMEonkmSvOoxaCK2Id7CRKMdqvR0kEouwnhk5WSFtsfi3sA0pkXzPFxwZeWM8vFtbffZOZzXaOhxCOfcj1NClZohlZhyc4jvkxmrpY7PSaAzih0AmHI7y0LYFi6fZu/K4EheVa1+KF55nWZ8ARikHMWKAKkyExkTak7xyN884TDmzURRaPlQg4jzQte5WMNjAG/hlHibdMBNvgwiYd49ZxteJ8ABdbiXVRl+2JGbdjl2ubpQZwOn7bJKlqO56bIwsZ+e4+pXsuOGdBahkHrUjtMEmH3DZbGc6CJLbcmdhdpApLQRRcLAazxJhzAwJ47FRYsHsj57LnYNvmcKdIxw8rxCdLUuzz95uw0T3ankEO5J9sjem+HMEuKdwXK1UcuOn2rjR8Sd/BuvQmeso27dFbPXqXYNS90Ml45YyTvcKSiopD181oZR703TFUSpR7dsiqROMr+p/2jN9h6a8WbQ8xpksyclaQByY/M77AssbXnG6wfhRsntNIINCZLbBnjXOyz6ZHIC5K4tSTdcnWaiYPeRPQmnw9UUvHAcNU2yMWsy0eU377yDS0WstTxOdQutTdkczl8kv5Lo26JiEK7mSIuRK19ffF9Zz8FG8+eKv5zdyIPjyQRDYBysUoDv5huKe2eoxJu/MWS2Pql/ZtUGeD6Ozm3mCvh0vQ9ceagBkY6Ocm3du0ziAKP29Ri0mjg4DizVorbLzsh+EQH/s2Pi9MnjUZDlEmuLl2Xfp7/w4j/8u0N0tVR70VDFuGdKpTjFY3vS8EJrPtyMTM51x1D9rb8gIql8aR/rJw4YF+huxg1mv5n6+tGVqg5msbPmF12eJijP4lkmaRwIpLW5pJTtaDkUj7uOeu1mm4k+Dt5nh0/0jPHzrv6bcTCcbV7UjMHDoTXXqEpFAAJ66rHR7zdAJu+YKsnTIZyLmOpcowq7LL8G9qTvV0OSpyQWUIavRSgbDHFqEqRs+JU94jAzkq8nCY5MTd9m5sIv9InfdT3k+pwpsE/FKge8nghFLtbUrafGkzTky8SE2druvVcIvbfXMfLIKRUYjJgnWc0gQzF5J6pzXM7D2r/RG6JDzASqjlbURq6v9bhNerlOVdMujWKEEVcKWIzlbt4RkihRjM8AUqIZQOyicGQ+4yfIjAHw5viuABONYs3OIWULnFqJxdvS9rNKhfxSjIq9cfqyzevq2xrRoMXEonobh6M3bD2Vang8OAeVeD1OXWPERi4pepCYFS9RJ/Xa/UWxptsqSNuGcb3fAzQSmLpXLGdWRoKXvSe7EYgc0bGcLOjSTu5RURKo+EF9i4KT9EJauf6VXw5dTf/CCIJRXE1bWzXhSCFYntohYhX2ldOCDYpi/jFBC6Vtkw0ud3/xq8Nmhd5gUk+SpngByCZH3Pm3H+jvlbMpiqkDkm1v74hDX13Xhrcw2eWyuqKBVoRCCniUvwpYNbGvBfjC6Hcizv0Aybciwj+4nybt5EPoEUm6S6Gs7fG7QpPdvrzpAxX70MlmdkF/gwyuhbEeJhLK+WL7qAsN5CvHPzVbsIf90x+nGTtMJPgpxVr0tJMj+vprXV4WxutfARBiOnqe58MhA857sd+MzKBgKnoLOBRTiC3qc/0/ULwbG2HCCD7nmwzz7M4nUuMvo8rgS7z0BF68OClT8X3JwSXbL5Wg=="
},
{
"key": "session_key.json",
"value": "6puTIBpn0u2Y8FJQ9C4gxuzORTgMmac9Tz9B2epw212hlaepET06ca/CPnwdirCNNg/tLG1wXd2MNSEMgZAnl1cPkF9hPabrRW+SUYpFDLu4yE9w5uc6ns//9pphedK5vS190oXcE7FaWBoso789JuQ0yoicNEnjAAjBWExnp0dXXufmXzUZSnKQ"
}
],
"oidc_clients": [