mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-08 04:01:26 +02:00
feat: use HS256 for Pocket ID's own session tokens + fixes (#1733)
This commit is contained in:
committed by
GitHub
parent
8c76ba0714
commit
33090fa3e1
@@ -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
|
||||
|
||||
@@ -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").
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(), "")
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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": [
|
||||
|
||||
Reference in New Issue
Block a user