diff --git a/backend/internal/cmds/encryption_key_rotate.go b/backend/internal/cmds/encryption_key_rotate.go index dddd0a7a..36f65f83 100644 --- a/backend/internal/cmds/encryption_key_rotate.go +++ b/backend/internal/cmds/encryption_key_rotate.go @@ -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 diff --git a/backend/internal/cmds/encryption_key_rotate_test.go b/backend/internal/cmds/encryption_key_rotate_test.go index e551943b..176c1fa5 100644 --- a/backend/internal/cmds/encryption_key_rotate_test.go +++ b/backend/internal/cmds/encryption_key_rotate_test.go @@ -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"). diff --git a/backend/internal/cmds/key_rotate.go b/backend/internal/cmds/key_rotate.go index 3f8c8f2c..b9303dce 100644 --- a/backend/internal/cmds/key_rotate.go +++ b/backend/internal/cmds/key_rotate.go @@ -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 +} diff --git a/backend/internal/cmds/key_rotate_test.go b/backend/internal/cmds/key_rotate_test.go index 6e051f48..81494e8d 100644 --- a/backend/internal/cmds/key_rotate_test.go +++ b/backend/internal/cmds/key_rotate_test.go @@ -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 diff --git a/backend/internal/service/e2etest_service.go b/backend/internal/service/e2etest_service.go index 2009db63..5c98e164 100644 --- a/backend/internal/service/e2etest_service.go +++ b/backend/internal/service/e2etest_service.go @@ -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 { diff --git a/backend/internal/service/jwt_service.go b/backend/internal/service/jwt_service.go index 7bfb45a1..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" ) @@ -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), diff --git a/backend/internal/service/jwt_service_test.go b/backend/internal/service/jwt_service_test.go index 13d5ed19..74a3b6bb 100644 --- a/backend/internal/service/jwt_service_test.go +++ b/backend/internal/service/jwt_service_test.go @@ -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) } diff --git a/backend/internal/utils/jwk/key_provider.go b/backend/internal/utils/jwk/key_provider.go index f53b69f8..5a8faa70 100644 --- a/backend/internal/utils/jwk/key_provider.go +++ b/backend/internal/utils/jwk/key_provider.go @@ -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) diff --git a/backend/internal/utils/jwk/key_provider_database.go b/backend/internal/utils/jwk/key_provider_database.go index d73327e7..f40f73dd 100644 --- a/backend/internal/utils/jwk/key_provider_database.go +++ b/backend/internal/utils/jwk/key_provider_database.go @@ -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) diff --git a/backend/internal/utils/jwk/key_provider_database_test.go b/backend/internal/utils/jwk/key_provider_database_test.go index 8beb9da4..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,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 { diff --git a/backend/internal/utils/jwk/utils.go b/backend/internal/utils/jwk/utils.go index 3c0b7266..52a34c5d 100644 --- a/backend/internal/utils/jwk/utils.go +++ b/backend/internal/utils/jwk/utils.go @@ -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(), "") +} diff --git a/backend/internal/utils/jwk/utils_test.go b/backend/internal/utils/jwk/utils_test.go index cc2ef919..fd2765d3 100644 --- a/backend/internal/utils/jwk/utils_test.go +++ b/backend/internal/utils/jwk/utils_test.go @@ -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 { diff --git a/tests/resources/export/database.json b/tests/resources/export/database.json index b6f7c7ce..7e9f4b75 100644 --- a/tests/resources/export/database.json +++ b/tests/resources/export/database.json @@ -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": [