From f893b25e9aa6d4cc406e5a78d729f237472d869c Mon Sep 17 00:00:00 2001 From: ItalyPaleAle <43508+ItalyPaleAle@users.noreply.github.com> Date: Tue, 1 Sep 2026 10:09:30 +0200 Subject: [PATCH] feat: use HS256 for Pocket ID's own session tokens --- .../internal/cmds/encryption_key_rotate.go | 32 ++-- .../cmds/encryption_key_rotate_test.go | 30 ++++ backend/internal/cmds/key_rotate.go | 57 ++++++- backend/internal/cmds/key_rotate_test.go | 82 +++++++++ backend/internal/service/e2etest_service.go | 5 + backend/internal/service/jwt_service.go | 125 +++++++++++++- backend/internal/service/jwt_service_test.go | 157 ++++++++++++++++++ backend/internal/utils/jwk/key_provider.go | 15 +- .../utils/jwk/key_provider_database.go | 37 +++-- .../utils/jwk/key_provider_database_test.go | 74 +++++++++ backend/internal/utils/jwk/utils.go | 23 +++ backend/internal/utils/jwk/utils_test.go | 45 +++++ tests/resources/export/database.json | 4 + 13 files changed, 649 insertions(+), 37 deletions(-) diff --git a/backend/internal/cmds/encryption_key_rotate.go b/backend/internal/cmds/encryption_key_rotate.go index dddd0a7a..8b9c0b3a 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.SaveKey(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..52625b01 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(), @@ -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.SaveKey(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..8a160913 100644 --- a/backend/internal/service/jwt_service.go +++ b/backend/internal/service/jwt_service.go @@ -43,9 +43,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 +74,24 @@ 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 := s.loadOrGenerateSigningKey(ctx) + if err != nil { + return fmt.Errorf("error loading signing key: %w", err) + } + + // Load the key used for Pocket ID's own sessions, which is symmetric + err = s.loadOrGenerateSessionKey(ctx) + if err != nil { + return fmt.Errorf("error loading session key: %w", err) + } + + return nil +} + +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 +128,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 +210,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 +272,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 +322,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 +332,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..147c700d 100644 --- a/backend/internal/service/jwt_service_test.go +++ b/backend/internal/service/jwt_service_test.go @@ -12,6 +12,7 @@ import ( "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" @@ -189,6 +190,162 @@ func TestJwtService_Init(t *testing.T) { } +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.SaveKey(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") + assert.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) + assert.ErrorContains(t, err, "session key is not initialized") + + _, err = service.VerifyAccessToken("some-token") + require.Error(t, err) + assert.ErrorContains(t, err, "session key is not initialized") + }) +} + func TestJwtService_GetPublicJWK(t *testing.T) { mockConfig := appconfig.NewTestAppConfigService(nil) db := testutils.NewDatabaseForTest(t) diff --git a/backend/internal/utils/jwk/key_provider.go b/backend/internal/utils/jwk/key_provider.go index f53b69f8..fd81374c 100644 --- a/backend/internal/utils/jwk/key_provider.go +++ b/backend/internal/utils/jwk/key_provider.go @@ -14,6 +14,7 @@ type KeyProviderOpts struct { EnvConfig *common.EnvConfigSchema DB *gorm.DB Kek []byte + DBKey string } type KeyProvider interface { @@ -22,7 +23,18 @@ type KeyProvider interface { SaveKey(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 +46,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..c436c1c7 100644 --- a/backend/internal/utils/jwk/key_provider_database.go +++ b/backend/internal/utils/jwk/key_provider_database.go @@ -15,11 +15,18 @@ 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" +) type KeyProviderDatabase struct { - db *gorm.DB - kek []byte + db *gorm.DB + kek []byte + dbKey string } func (f *KeyProviderDatabase) Init(opts KeyProviderOpts) error { @@ -30,12 +37,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 +58,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,19 +69,19 @@ 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 @@ -88,7 +101,7 @@ func (f *KeyProviderDatabase) SaveKey(ctx context.Context, key jwk.Key) error { } // Save to database row := model.KV{ - Key: PrivateKeyDBKey, + Key: f.dbKey, Value: new(base64.StdEncoding.EncodeToString(enc)), } @@ -103,9 +116,9 @@ 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) + // There's one scenario where if Pocket ID is started fresh with more than 1 replica, they both could be trying to create the key in the database at the same time + // In this case, only one of the replicas will succeed and the other one(s) will return an error here, which will cascade down and cause the replica(s) to crash and be restarted (at that point they'll load the then-existing key from the database) + return fmt.Errorf("failed to store key in database: %w", err) } return 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..2b7f5977 100644 --- a/backend/internal/utils/jwk/key_provider_database_test.go +++ b/backend/internal/utils/jwk/key_provider_database_test.go @@ -259,6 +259,80 @@ func TestKeyProviderDatabase_SaveKey(t *testing.T) { }) } +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 { t.Helper() 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": [