From ea741dda6badd73ca3b3025d31f8c66167025990 Mon Sep 17 00:00:00 2001 From: Karl Persson Date: Tue, 10 Oct 2023 14:17:16 +0200 Subject: [PATCH] Signingkeys: Add local cache (#76234) * IDForwarding: change audience to be prefixed by org and remove JTI * IDForwarding: Construct new signer each time we want to sign a token. * SigningKeys: Simplify storage layer and move logic to service * SigningKeys: Add private key to local cache --- pkg/services/auth/idimpl/service.go | 14 +- pkg/services/auth/idimpl/signer.go | 58 +++-- pkg/services/authn/clients/ext_jwt_test.go | 6 +- .../oauthserver/oasimpl/service_test.go | 8 +- pkg/services/signingkeys/signingkeys.go | 9 + .../signingkeys/signingkeysimpl/service.go | 176 ++++++++++++- .../signingkeysimpl/service_test.go | 69 +++-- .../signingkeys/signingkeystest/fake.go | 33 +-- .../signingkeys/signingkeystore/fake.go | 56 ++-- .../signingkeys/signingkeystore/store.go | 184 +++----------- .../signingkeys/signingkeystore/store_test.go | 239 +++++------------- 11 files changed, 386 insertions(+), 466 deletions(-) diff --git a/pkg/services/auth/idimpl/service.go b/pkg/services/auth/idimpl/service.go index 19a33033a81..08b65b4df41 100644 --- a/pkg/services/auth/idimpl/service.go +++ b/pkg/services/auth/idimpl/service.go @@ -3,7 +3,6 @@ package idimpl import ( "context" "fmt" - "strconv" "time" "github.com/go-jose/go-jose/v3/jwt" @@ -67,10 +66,9 @@ func (s *Service) SignIdentity(ctx context.Context, id identity.Requester) (stri now := time.Now() token, err := s.signer.SignIDToken(ctx, &auth.IDClaims{ Claims: jwt.Claims{ - ID: identifier, Issuer: s.cfg.AppURL, - Audience: jwt.Audience{strconv.FormatInt(id.GetOrgID(), 10)}, - Subject: fmt.Sprintf("%s:%s", namespace, identifier), + Audience: getAudience(id.GetOrgID()), + Subject: getSubject(namespace, identifier), Expiry: jwt.NewNumericDate(now.Add(tokenTTL)), IssuedAt: jwt.NewNumericDate(now), }, @@ -102,6 +100,14 @@ func (s *Service) hook(ctx context.Context, identity *authn.Identity, _ *authn.R return nil } +func getAudience(orgID int64) jwt.Audience { + return jwt.Audience{fmt.Sprintf("org:%d", orgID)} +} + +func getSubject(namespace, identifier string) string { + return fmt.Sprintf("%s:%s", namespace, identifier) +} + func prefixCacheKey(key string) string { return fmt.Sprintf("%s-%s", cachePrefix, key) } diff --git a/pkg/services/auth/idimpl/signer.go b/pkg/services/auth/idimpl/signer.go index 15300c0979e..4cd8a090503 100644 --- a/pkg/services/auth/idimpl/signer.go +++ b/pkg/services/auth/idimpl/signer.go @@ -11,40 +11,20 @@ import ( "github.com/grafana/grafana/pkg/services/signingkeys" ) -const idSignerKeyPrefix = "id" +const ( + keyPrefix = "id" + headerKeyID = "kid" +) var _ auth.IDSigner = (*LocalSigner)(nil) func ProvideLocalSigner(keyService signingkeys.Service, features featuremgmt.FeatureToggles) (*LocalSigner, error) { - if features.IsEnabled(featuremgmt.FlagIdForwarding) { - id, key, err := keyService.GetOrCreatePrivateKey(context.Background(), idSignerKeyPrefix, jose.ES256) - if err != nil { - return nil, err - } - - // FIXME: Handle key rotation - signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.ES256, Key: key}, &jose.SignerOptions{ - ExtraHeaders: map[jose.HeaderKey]interface{}{ - "kid": id, - }, - }) - - if err != nil { - return nil, err - } - - return &LocalSigner{ - features: features, - signer: signer, - }, nil - } - - return &LocalSigner{features: features}, nil + return &LocalSigner{features, keyService}, nil } type LocalSigner struct { - signer jose.Signer - features featuremgmt.FeatureToggles + features featuremgmt.FeatureToggles + keyService signingkeys.Service } func (s *LocalSigner) SignIDToken(ctx context.Context, claims *auth.IDClaims) (string, error) { @@ -52,7 +32,12 @@ func (s *LocalSigner) SignIDToken(ctx context.Context, claims *auth.IDClaims) (s return "", nil } - builder := jwt.Signed(s.signer).Claims(claims.Claims) + signer, err := s.getSigner(ctx) + if err != nil { + return "", err + } + + builder := jwt.Signed(signer).Claims(claims.Claims) token, err := builder.CompactSerialize() if err != nil { @@ -61,3 +46,20 @@ func (s *LocalSigner) SignIDToken(ctx context.Context, claims *auth.IDClaims) (s return token, nil } + +func (s *LocalSigner) getSigner(ctx context.Context) (jose.Signer, error) { + id, key, err := s.keyService.GetOrCreatePrivateKey(ctx, keyPrefix, jose.ES256) + if err != nil { + return nil, err + } + + signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.ES256, Key: key}, &jose.SignerOptions{ + ExtraHeaders: map[jose.HeaderKey]any{headerKeyID: id}, + }) + + if err != nil { + return nil, err + } + + return signer, nil +} diff --git a/pkg/services/authn/clients/ext_jwt_test.go b/pkg/services/authn/clients/ext_jwt_test.go index c3a6a82ab0c..7f2f3e91f7d 100644 --- a/pkg/services/authn/clients/ext_jwt_test.go +++ b/pkg/services/authn/clients/ext_jwt_test.go @@ -2,7 +2,6 @@ package clients import ( "context" - "crypto" "crypto/rand" "crypto/rsa" "fmt" @@ -516,8 +515,9 @@ func setupTestCtx(t *testing.T, cfg *setting.Cfg) *testEnv { } } - signingKeysSvc := &signingkeystest.FakeSigningKeysService{ExpectedKeys: map[string]crypto.Signer{ - signingkeys.ServerPrivateKeyID: pk}, + signingKeysSvc := &signingkeystest.FakeSigningKeysService{ + ExpectedSinger: pk, + ExpectedKeyID: signingkeys.ServerPrivateKeyID, } userSvc := &usertest.FakeUserService{} diff --git a/pkg/services/extsvcauth/oauthserver/oasimpl/service_test.go b/pkg/services/extsvcauth/oauthserver/oasimpl/service_test.go index 1188a6f800d..c97c5cc2c78 100644 --- a/pkg/services/extsvcauth/oauthserver/oasimpl/service_test.go +++ b/pkg/services/extsvcauth/oauthserver/oasimpl/service_test.go @@ -2,7 +2,6 @@ package oasimpl import ( "context" - "crypto" "crypto/rand" "crypto/rsa" "encoding/base64" @@ -93,10 +92,9 @@ func setupTestEnv(t *testing.T) *TestEnv { } env.S.oauthProvider = newProvider(config, env.S, &signingkeystest.FakeSigningKeysService{ - ExpectedKeys: map[string]crypto.Signer{ - "default": pk, - }, - ExpectedError: nil, + ExpectedSinger: pk, + ExpectedKeyID: "default", + ExpectedError: nil, }) return env diff --git a/pkg/services/signingkeys/signingkeys.go b/pkg/services/signingkeys/signingkeys.go index bbb7ceba529..f5678829886 100644 --- a/pkg/services/signingkeys/signingkeys.go +++ b/pkg/services/signingkeys/signingkeys.go @@ -10,6 +10,7 @@ package signingkeys import ( "context" "crypto" + "time" "github.com/go-jose/go-jose/v3" ) @@ -26,3 +27,11 @@ type Service interface { GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) GetOrCreatePrivateKey(ctx context.Context, keyPrefix string, alg jose.SignatureAlgorithm) (string, crypto.Signer, error) } + +type SigningKey struct { + KeyID string `xorm:"key_id"` + PrivateKey []byte `xorm:"private_key"` + AddedAt time.Time `xorm:"added_at"` + ExpiresAt *time.Time `xorm:"expires_at"` + Alg jose.SignatureAlgorithm `xorm:"alg"` +} diff --git a/pkg/services/signingkeys/signingkeysimpl/service.go b/pkg/services/signingkeys/signingkeysimpl/service.go index 5e314d7edd0..500f0606dd9 100644 --- a/pkg/services/signingkeys/signingkeysimpl/service.go +++ b/pkg/services/signingkeys/signingkeysimpl/service.go @@ -6,7 +6,10 @@ import ( "crypto/ecdsa" "crypto/elliptic" "crypto/rand" + "crypto/x509" + "encoding/base64" "encoding/json" + "encoding/pem" "errors" "net/http" "strings" @@ -17,6 +20,7 @@ import ( "github.com/grafana/grafana/pkg/api/response" "github.com/grafana/grafana/pkg/api/routing" "github.com/grafana/grafana/pkg/infra/db" + "github.com/grafana/grafana/pkg/infra/localcache" "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/infra/remotecache" contextmodel "github.com/grafana/grafana/pkg/services/contexthandler/model" @@ -31,9 +35,11 @@ func ProvideEmbeddedSigningKeysService(dbStore db.DB, secretsService secrets.Ser remoteCache remotecache.CacheStorage, routerRegister routing.RouteRegister, ) (*Service, error) { s := &Service{ - log: log.New("auth.key_service"), - store: signingkeystore.NewSigningKeyStore(dbStore, secretsService), - remoteCache: remoteCache, + log: log.New("auth.key_service"), + store: signingkeystore.NewSigningKeyStore(dbStore), + secretsService: secretsService, + remoteCache: remoteCache, + localCache: localcache.New(1*time.Hour, 1*time.Hour), } s.registerAPIEndpoints(routerRegister) @@ -46,14 +52,17 @@ func ProvideEmbeddedSigningKeysService(dbStore db.DB, secretsService secrets.Ser // // The service is under active development and is not yet ready for production use. type Service struct { - log log.Logger - store signingkeystore.SigningStore - remoteCache remotecache.CacheStorage + log log.Logger + store signingkeystore.SigningStore + secretsService secrets.Service + remoteCache remotecache.CacheStorage + localCache *localcache.CacheService } const ( jwksCacheKey = "signingkeys-jwks" - defaultExpiry = 12 * time.Hour + jwksTTL = 12 * time.Hour + privateKeyTTL = 60 * time.Second ) // GetJWKS returns the JSON Web Key Set (JWKS) with all the keys that can be used to verify tokens (public keys) @@ -66,15 +75,20 @@ func (s *Service) GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) { } } - jwks, err := s.store.GetJWKS(ctx) + keys, err := s.store.List(ctx) if err != nil { return jose.JSONWebKeySet{}, err } + jwks, err := s.buildJWKS(ctx, keys) + if err != nil { + return jwks, err + } + // cache jwks jwksBytes, err := json.Marshal(jwks) if err == nil { - if err := s.remoteCache.Set(ctx, jwksCacheKey, jwksBytes, defaultExpiry); err != nil { + if err := s.remoteCache.Set(ctx, jwksCacheKey, jwksBytes, jwksTTL); err != nil { s.log.Warn("Failed to cache JWKS", "err", err) } } @@ -82,6 +96,24 @@ func (s *Service) GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) { return jwks, err } +func (s *Service) buildJWKS(ctx context.Context, keys []signingkeys.SigningKey) (jose.JSONWebKeySet, error) { + var jwks jose.JSONWebKeySet + for _, key := range keys { + assertedKey, err := s.decodePrivateKey(ctx, key.PrivateKey) + if err != nil { + return jwks, err + } + + jwks.Keys = append(jwks.Keys, jose.JSONWebKey{ + Key: assertedKey.Public(), + Algorithm: string(key.Alg), + KeyID: key.KeyID, + Use: "sig", + }) + } + return jwks, nil +} + // GetOrCreatePrivateKey returns the private key with the specified key ID. If the key does not exist, it will be // created with the specified algorithm. // The key will be automatically rotated at the beginning of each month. The previous key will be kept for 30 days. @@ -93,30 +125,146 @@ func (s *Service) GetOrCreatePrivateKey(ctx context.Context, } keyID := keyMonthScopedID(keyPrefix, alg) - signer, err := s.store.GetPrivateKey(ctx, keyID) + signer, err := s.getPrivateKey(ctx, keyID) if err == nil { return keyID, signer, nil } + + // we only want to create a new signing key if none exits for keyID + if !errors.Is(err, signingkeys.ErrSigningKeyNotFound) { + return "", nil, err + } + s.log.Debug("Private key not found, generating new key", "keyID", keyID, "err", err) + signer, err = s.addPrivateKey(ctx, keyID, alg, false) + if err != nil { + return "", nil, err + } + + return keyID, signer, nil +} + +func (s *Service) getPrivateKey(ctx context.Context, keyID string) (crypto.Signer, error) { + if key, ok := s.localCache.Get(keyID); ok { + return key.(crypto.Signer), nil + } + + key, err := s.store.Get(ctx, keyID) + if err != nil { + return nil, err + } + + singer, err := s.decodePrivateKey(ctx, key.PrivateKey) + if err != nil { + return nil, err + } + + s.localCache.Set(keyID, singer, privateKeyTTL) + return singer, nil +} + +func (s *Service) addPrivateKey(ctx context.Context, keyID string, alg jose.SignatureAlgorithm, force bool) (crypto.Signer, error) { privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) if err != nil { s.log.Error("Error generating private key", "err", err) - return "", nil, signingkeys.ErrKeyGenerationFailed.Errorf("Error generating private key: %v", err) + return nil, signingkeys.ErrKeyGenerationFailed.Errorf("Error generating private key: %v", err) + } + + encoded, err := s.encodePrivateKey(ctx, privateKey) + if err != nil { + s.log.Error("Error encoding private key", "err", err) + return nil, err } expiry := time.Now().Add(30 * 24 * time.Hour) - if signer, err = s.store.AddPrivateKey(ctx, keyID, alg, privateKey, &expiry, false); err != nil && !errors.Is(err, signingkeys.ErrSigningKeyAlreadyExists) { - return "", nil, err + key, err := s.store.Add(ctx, &signingkeys.SigningKey{ + KeyID: keyID, + PrivateKey: encoded, + ExpiresAt: &expiry, + Alg: alg, + }, force) + + if err != nil && !errors.Is(err, signingkeys.ErrSigningKeyAlreadyExists) { + return nil, err } + signer, err := s.decodePrivateKey(ctx, key.PrivateKey) + if err != nil { + return nil, err + } + + // invalidate local cache + s.localCache.Delete(keyID) + // invalidate cache if err := s.remoteCache.Delete(ctx, jwksCacheKey); err != nil { // not a critical error, key might not be in cache s.log.Debug("Failed to invalidate JWKS cache", "err", err) } - return keyID, signer, nil + return signer, nil +} + +func (s *Service) encodePrivateKey(ctx context.Context, privateKey crypto.Signer) ([]byte, error) { + // Encode private key to binary format + pKeyBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + return nil, err + } + + // Encode private key to PEM format + privateKeyPEM := pem.EncodeToMemory(&pem.Block{ + Type: "PRIVATE KEY", + Bytes: pKeyBytes, + }) + + encrypted, err := s.secretsService.Encrypt(ctx, privateKeyPEM, secrets.WithoutScope()) + if err != nil { + return nil, err + } + + encoded := make([]byte, base64.StdEncoding.EncodedLen(len(encrypted))) + base64.StdEncoding.Encode(encoded, encrypted) + return encoded, nil +} + +func (s *Service) decodePrivateKey(ctx context.Context, privateKey []byte) (crypto.Signer, error) { + // Bail out if empty string since it'll cause a segfault in Decrypt + if len(privateKey) == 0 { + return nil, errors.New("private key is empty") + } + + payload := make([]byte, base64.StdEncoding.DecodedLen(len(privateKey))) + _, err := base64.StdEncoding.Decode(payload, privateKey) + if err != nil { + return nil, err + } + + decrypted, err := s.secretsService.Decrypt(ctx, payload) + if err != nil { + return nil, err + } + + block, _ := pem.Decode(decrypted) + if block == nil { + return nil, errors.New("failed to decode private key PEM") + } + + if block.Type != "PRIVATE KEY" { + return nil, errors.New("invalid block type") + } + + parsedKey, err := x509.ParsePKCS8PrivateKey(block.Bytes) + if err != nil { + return nil, err + } + + assertedKey, ok := parsedKey.(crypto.Signer) + if !ok { + return nil, errors.New("failed to assert private key as crypto.Signer") + } + return assertedKey, nil } func keyMonthScopedID(keyPrefix string, alg jose.SignatureAlgorithm) string { diff --git a/pkg/services/signingkeys/signingkeysimpl/service_test.go b/pkg/services/signingkeys/signingkeysimpl/service_test.go index aac3953bc82..239c255f622 100644 --- a/pkg/services/signingkeys/signingkeysimpl/service_test.go +++ b/pkg/services/signingkeys/signingkeysimpl/service_test.go @@ -13,13 +13,16 @@ import ( "time" "github.com/go-jose/go-jose/v3" + "github.com/grafana/grafana/pkg/services/signingkeys" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/grafana/grafana/pkg/api/routing" + "github.com/grafana/grafana/pkg/infra/db" + "github.com/grafana/grafana/pkg/infra/localcache" "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/infra/remotecache" - "github.com/grafana/grafana/pkg/services/signingkeys" + secretstest "github.com/grafana/grafana/pkg/services/secrets/fakes" "github.com/grafana/grafana/pkg/services/signingkeys/signingkeystore" "github.com/grafana/grafana/pkg/services/user" "github.com/grafana/grafana/pkg/web/webtest" @@ -32,28 +35,33 @@ ielIkb6/Ys51o7KjHxtANhPesw== -----END PRIVATE KEY-----` ) -func getPrivateKey(t *testing.T) *ecdsa.PrivateKey { +func getPrivateKey(t *testing.T, svc *Service) []byte { pemBlock, _ := pem.Decode([]byte(privateKeyPem)) privateKey, err := x509.ParsePKCS8PrivateKey(pemBlock.Bytes) require.NoError(t, err) - return privateKey.(*ecdsa.PrivateKey) + + bytes, err := svc.encodePrivateKey(context.Background(), privateKey.(*ecdsa.PrivateKey)) + require.NoError(t, err) + return bytes } -func TestEmbeddedKeyService_GetJWKS_OnlyPublicKeyShared(t *testing.T) { - mockStore := signingkeystore.NewFakeStore() - cacheStorage := remotecache.NewFakeCacheStorage() - - _, err := mockStore.AddPrivateKey(context.Background(), signingkeys.ServerPrivateKeyID, jose.ES256, getPrivateKey(t), nil, false) - require.NoError(t, err) - - _, err = mockStore.AddPrivateKey(context.Background(), "other", jose.ES256, getPrivateKey(t), nil, false) - require.NoError(t, err) +func TestIntegrationEmbeddedKeyService_GetJWKS_OnlyPublicKeyShared(t *testing.T) { + if testing.Short() { + t.Skip("skipping integration test") + } svc := &Service{ - log: log.NewNopLogger(), - store: mockStore, - remoteCache: cacheStorage, + log: log.NewNopLogger(), + store: signingkeystore.NewSigningKeyStore(db.InitTestDB(t)), + secretsService: secretstest.NewFakeSecretsService(), + remoteCache: remotecache.NewFakeCacheStorage(), + localCache: localcache.New(privateKeyTTL, 10*time.Hour), } + + _, _, err := svc.GetOrCreatePrivateKey(context.Background(), "key-1", jose.ES256) + require.NoError(t, err) + _, _, err = svc.GetOrCreatePrivateKey(context.Background(), "key-2", jose.ES256) + require.NoError(t, err) jwks, err := svc.GetJWKS(context.Background()) require.NoError(t, err) @@ -79,14 +87,18 @@ func TestEmbeddedKeyService_GetJWKS_OnlyPublicKeyShared(t *testing.T) { } } -func TestEmbeddedKeyService_GetOrCreatePrivateKey(t *testing.T) { - mockStore := signingkeystore.NewFakeStore() +func TestIntegrationEmbeddedKeyService_GetOrCreatePrivateKey(t *testing.T) { + if testing.Short() { + t.Skip("skipping integration test") + } cacheStorage := remotecache.NewFakeCacheStorage() svc := &Service{ - log: log.NewNopLogger(), - store: mockStore, - remoteCache: cacheStorage, + log: log.NewNopLogger(), + store: signingkeystore.NewSigningKeyStore(db.InitTestDB(t)), + secretsService: secretstest.NewFakeSecretsService(), + remoteCache: cacheStorage, + localCache: localcache.New(privateKeyTTL, 10*time.Hour), } wantedKeyID := keyMonthScopedID("test", jose.ES256) @@ -110,7 +122,6 @@ func TestEmbeddedKeyService_GetOrCreatePrivateKey(t *testing.T) { // new key is generated, so jwks cache should be voided require.Len(t, cacheStorage.Storage, 0) - assert.Contains(t, mockStore.PrivateKeys, wantedKeyID) err = cacheStorage.Set(context.Background(), jwksCacheKey, []byte("invalid"), 0) require.NoError(t, err) @@ -122,7 +133,6 @@ func TestEmbeddedKeyService_GetOrCreatePrivateKey(t *testing.T) { require.Equal(t, key, key2) require.Equal(t, wantedKeyID, id) - assert.Len(t, mockStore.PrivateKeys, 1) // no new key is generated, so jwks cache should not be voided require.Len(t, cacheStorage.Storage, 1) } @@ -132,9 +142,11 @@ func TestExposeJWKS(t *testing.T) { mockStore := signingkeystore.NewFakeStore() cacheStorage := remotecache.NewFakeCacheStorage() svc := &Service{ - log: log.NewNopLogger(), - store: mockStore, - remoteCache: cacheStorage, + log: log.NewNopLogger(), + store: mockStore, + remoteCache: cacheStorage, + secretsService: secretstest.NewFakeSecretsService(), + localCache: localcache.New(privateKeyTTL, 10*time.Hour), } routerRegister := routing.NewRouteRegister() @@ -142,8 +154,13 @@ func TestExposeJWKS(t *testing.T) { svc.registerAPIEndpoints(routerRegister) server := webtest.NewServer(t, routerRegister) + _, err := mockStore.Add(context.Background(), &signingkeys.SigningKey{ + KeyID: "test-key", + PrivateKey: getPrivateKey(t, svc), + AddedAt: time.Now(), + Alg: jose.ES256, + }, false) - _, err := mockStore.AddPrivateKey(context.Background(), "test-key", jose.ES256, getPrivateKey(t), nil, false) require.NoError(t, err) // create a new request context diff --git a/pkg/services/signingkeys/signingkeystest/fake.go b/pkg/services/signingkeys/signingkeystest/fake.go index 48620bfc0b1..3f7e23634a4 100644 --- a/pkg/services/signingkeys/signingkeystest/fake.go +++ b/pkg/services/signingkeys/signingkeystest/fake.go @@ -3,15 +3,14 @@ package signingkeystest import ( "context" "crypto" - "time" "github.com/go-jose/go-jose/v3" ) type FakeSigningKeysService struct { ExpectedJSONWebKeySet jose.JSONWebKeySet - ExpectedJSONWebKey jose.JSONWebKey - ExpectedKeys map[string]crypto.Signer + ExpectedKeyID string + ExpectedSinger crypto.Signer ExpectedError error } @@ -19,30 +18,6 @@ func (s *FakeSigningKeysService) GetJWKS(ctx context.Context) (jose.JSONWebKeySe return s.ExpectedJSONWebKeySet, nil } -// GetPublicKey returns the public key with the specified key ID -func (s *FakeSigningKeysService) GetPublicKey(ctx context.Context, keyID string) (crypto.PublicKey, error) { - return s.ExpectedKeys[keyID].Public(), s.ExpectedError -} - -// GetPrivateKey returns the private key with the specified key ID -func (s *FakeSigningKeysService) GetPrivateKey(ctx context.Context, keyID string) (crypto.PrivateKey, error) { - return s.ExpectedKeys[keyID], s.ExpectedError -} - -// AddPrivateKey adds a private key to the service -func (s *FakeSigningKeysService) AddPrivateKey(ctx context.Context, keyID string, - privateKey crypto.Signer, alg jose.SignatureAlgorithm, expiresAt *time.Time, force bool) error { - if s.ExpectedError != nil { - return s.ExpectedError - } - s.ExpectedKeys[keyID] = privateKey - return nil -} - -func (s *FakeSigningKeysService) GetOrCreatePrivateKey(ctx context.Context, - keyPrefix string, alg jose.SignatureAlgorithm) (string, crypto.Signer, error) { - if s.ExpectedError != nil { - return "", nil, s.ExpectedError - } - return keyPrefix, s.ExpectedKeys[keyPrefix], nil +func (s *FakeSigningKeysService) GetOrCreatePrivateKey(ctx context.Context, keyPrefix string, alg jose.SignatureAlgorithm) (string, crypto.Signer, error) { + return s.ExpectedKeyID, s.ExpectedSinger, s.ExpectedError } diff --git a/pkg/services/signingkeys/signingkeystore/fake.go b/pkg/services/signingkeys/signingkeystore/fake.go index 63fed1841df..b78721353e7 100644 --- a/pkg/services/signingkeys/signingkeystore/fake.go +++ b/pkg/services/signingkeys/signingkeystore/fake.go @@ -3,13 +3,15 @@ package signingkeystore import ( "context" "crypto" - "fmt" - "time" "github.com/go-jose/go-jose/v3" + "github.com/grafana/grafana/pkg/services/signingkeys" ) +var _ SigningStore = (*FakeStore)(nil) + type FakeStore struct { + Keys map[string]signingkeys.SigningKey PrivateKeys map[string]crypto.Signer jwks jose.JSONWebKeySet } @@ -17,46 +19,34 @@ type FakeStore struct { func NewFakeStore() *FakeStore { return &FakeStore{ PrivateKeys: make(map[string]crypto.Signer), + Keys: make(map[string]signingkeys.SigningKey), jwks: jose.JSONWebKeySet{}, } } -func (s *FakeStore) GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) { - return s.jwks, nil -} - -func (s *FakeStore) AddPrivateKey(ctx context.Context, keyID string, alg jose.SignatureAlgorithm, - privateKey crypto.Signer, expiresAt *time.Time, force bool) (crypto.Signer, error) { +func (s *FakeStore) Add(ctx context.Context, key *signingkeys.SigningKey, force bool) (*signingkeys.SigningKey, error) { if !force { - if key, ok := s.PrivateKeys[keyID]; ok { - if !hasExpired(key) { - return nil, fmt.Errorf("key already exists and has not expired") - } + if _, ok := s.Keys[key.KeyID]; ok { + return nil, signingkeys.ErrSigningKeyAlreadyExists } } - s.PrivateKeys[keyID] = privateKey + s.Keys[key.KeyID] = *key + return key, nil +} - jwk := jose.JSONWebKey{ - Key: privateKey.Public(), - Algorithm: string(alg), - KeyID: keyID, - Use: "sig", +func (s *FakeStore) List(ctx context.Context) ([]signingkeys.SigningKey, error) { + out := make([]signingkeys.SigningKey, 0, len(s.Keys)) + for _, key := range s.Keys { + out = append(out, key) + } + return out, nil +} + +func (s *FakeStore) Get(ctx context.Context, keyID string) (*signingkeys.SigningKey, error) { + if key, ok := s.Keys[keyID]; ok { + return &key, nil } - s.jwks.Keys = append(s.jwks.Keys, jwk) - - return privateKey, nil -} - -func (s *FakeStore) GetPrivateKey(ctx context.Context, keyID string) (crypto.Signer, error) { - if key, ok := s.PrivateKeys[keyID]; ok { - return key, nil - } - - return nil, fmt.Errorf("key not found") -} - -func hasExpired(key crypto.Signer) bool { - return false + return nil, signingkeys.ErrSigningKeyNotFound } diff --git a/pkg/services/signingkeys/signingkeystore/store.go b/pkg/services/signingkeys/signingkeystore/store.go index 0f6f23b3d4f..0c01a2850a1 100644 --- a/pkg/services/signingkeys/signingkeystore/store.go +++ b/pkg/services/signingkeys/signingkeystore/store.go @@ -2,11 +2,7 @@ package signingkeystore import ( "context" - "crypto" - "crypto/x509" "database/sql" - "encoding/base64" - "encoding/pem" "errors" "time" @@ -15,21 +11,19 @@ import ( "github.com/grafana/grafana/pkg/infra/db" "github.com/grafana/grafana/pkg/infra/localcache" "github.com/grafana/grafana/pkg/infra/log" - "github.com/grafana/grafana/pkg/services/secrets" "github.com/grafana/grafana/pkg/services/signingkeys" "github.com/grafana/grafana/pkg/services/sqlstore" ) type SigningStore interface { - // GetJWKS returns the JSON Web Key Set for the service - GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) - // AddPrivateKey adds a private key to the service. If the key already exists, it will be updated if force is true. + // List returns all non expired keys + List(ctx context.Context) ([]signingkeys.SigningKey, error) + // Add adds a signing key to the database. If the key already exists, it will be updated if force is true. // If force is false, the key will only be updated if it has expired. If the key does not exist, it will be added. - // If expiresAt is nil, the key will not expire. Retrieve the result key with GetPrivateKey. - AddPrivateKey(ctx context.Context, keyID string, alg jose.SignatureAlgorithm, - privateKey crypto.Signer, expiresAt *time.Time, force bool) (crypto.Signer, error) - // GetPrivateKey returns the private key with the specified key ID - GetPrivateKey(ctx context.Context, keyID string) (crypto.Signer, error) + // If expiresAt is nil, the key will not expire. Retrieve the result key with Get. + Add(ctx context.Context, key *signingkeys.SigningKey, force bool) (*signingkeys.SigningKey, error) + // Get returns the signing key with the specified key ID + Get(ctx context.Context, keyID string) (*signingkeys.SigningKey, error) } var _ SigningStore = (*Store)(nil) @@ -37,10 +31,9 @@ var _ SigningStore = (*Store)(nil) const cleanupRateLimitKey = "signingkeys-cleanup" type Store struct { - dbStore db.DB - secretsService secrets.Service - log log.Logger - localCache *localcache.CacheService + dbStore db.DB + log log.Logger + localCache *localcache.CacheService } type SigningKey struct { @@ -52,90 +45,57 @@ type SigningKey struct { Alg jose.SignatureAlgorithm `json:"alg" xorm:"alg" db:"alg"` } -func NewSigningKeyStore(dbStore db.DB, secretsService secrets.Service) *Store { +func NewSigningKeyStore(dbStore db.DB) *Store { return &Store{ - dbStore: dbStore, - secretsService: secretsService, - log: log.New("signing.key_service"), - localCache: localcache.New(12*time.Hour, 4*time.Hour), + dbStore: dbStore, + log: log.New("signing.key_service"), + localCache: localcache.New(12*time.Hour, 4*time.Hour), } } -// GetJWKS returns the JSON Web Key Set (JWKS) for the service. Expired keys will not be returned. -func (s *Store) GetJWKS(ctx context.Context) (jose.JSONWebKeySet, error) { - keySet := jose.JSONWebKeySet{} +func (s *Store) List(ctx context.Context) ([]signingkeys.SigningKey, error) { + var keys []signingkeys.SigningKey - keys := []*SigningKey{} err := s.dbStore.WithDbSession(ctx, func(dbSession *sqlstore.DBSession) error { return dbSession.SQL("SELECT * FROM signing_key WHERE expires_at IS NULL OR expires_at > ?", time.Now()).Find(&keys) }) - if err != nil { - return keySet, err - } - - for _, key := range keys { - assertedKey, err := s.decodePrivateKey(ctx, key) - if err != nil { - return keySet, err - } - - keySet.Keys = append(keySet.Keys, jose.JSONWebKey{ - Key: assertedKey.Public(), - Algorithm: string(key.Alg), - KeyID: key.KeyID, - Use: "sig", - }) - } - - return keySet, nil -} - -// AddPrivateKey adds a private key to the service. -func (s *Store) AddPrivateKey(ctx context.Context, - keyID string, alg jose.SignatureAlgorithm, privateKey crypto.Signer, expiresAt *time.Time, force bool) (crypto.Signer, error) { - privateKeyPEM, err := s.encodePrivateKey(ctx, privateKey) if err != nil { return nil, err } - key := &SigningKey{ - KeyID: keyID, - PrivateKey: privateKeyPEM, - AddedAt: time.Now(), - Alg: alg, - ExpiresAt: expiresAt, - } + return keys, nil +} - var signer crypto.Signer - err = s.dbStore.WithTransactionalDbSession(ctx, func(tx *sqlstore.DBSession) error { - existingKey := SigningKey{} - _, err := tx.SQL("SELECT * FROM signing_key WHERE key_id = ?", keyID).Get(&existingKey) +// Add adds a private key to the service. +func (s *Store) Add(ctx context.Context, key *signingkeys.SigningKey, force bool) (*signingkeys.SigningKey, error) { + var result *signingkeys.SigningKey + + err := s.dbStore.WithTransactionalDbSession(ctx, func(tx *sqlstore.DBSession) error { + existingKey := &signingkeys.SigningKey{} + exists, err := tx.SQL("SELECT * FROM signing_key WHERE key_id = ?", key.KeyID).Get(existingKey) if err != nil && !errors.Is(err, sql.ErrNoRows) { return err } - if len(existingKey.PrivateKey) == 0 { + if !exists { _, err = tx.Exec("INSERT INTO signing_key (key_id, private_key, added_at, alg, expires_at) VALUES (?, ?, ?, ?, ?)", key.KeyID, key.PrivateKey, key.AddedAt, key.Alg, key.ExpiresAt, ) - signer = privateKey + result = key return err } if force || (existingKey.ExpiresAt != nil && existingKey.ExpiresAt.Before(time.Now())) { _, err = tx.Exec("UPDATE signing_key SET private_key = ?, added_at = ?, alg = ?, expires_at = ? WHERE key_id = ?", key.PrivateKey, key.AddedAt, key.Alg, key.ExpiresAt, key.KeyID) - signer = privateKey + + result = key return err } - signer, err = s.decodePrivateKey(ctx, &existingKey) - if err != nil { - return err - } - - return signingkeys.ErrSigningKeyAlreadyExists.Errorf("The specified key already exists: %s", keyID) + result = existingKey + return signingkeys.ErrSigningKeyAlreadyExists.Errorf("The specified key already exists: %s", existingKey.KeyID) }) if _, ok := s.localCache.Get(cleanupRateLimitKey); !ok { @@ -156,14 +116,18 @@ func (s *Store) AddPrivateKey(ctx context.Context, s.localCache.Set(cleanupRateLimitKey, true, 1*time.Hour) } - return signer, err + return result, err } -// GetPrivateKey returns the private key with the specified key ID. Expired keys will not be returned. -func (s *Store) GetPrivateKey(ctx context.Context, keyID string) (crypto.Signer, error) { - key := SigningKey{} +// Get implements SigningStore. +func (s *Store) Get(ctx context.Context, keyID string) (*signingkeys.SigningKey, error) { + key := signingkeys.SigningKey{} err := s.dbStore.WithDbSession(ctx, func(dbSession *sqlstore.DBSession) error { - _, err := dbSession.SQL("SELECT * FROM signing_key WHERE key_id = ?", keyID).Get(&key) + exists, err := dbSession.SQL("SELECT * FROM signing_key WHERE key_id = ?", keyID).Get(&key) + if !exists { + return signingkeys.ErrSigningKeyNotFound.Errorf("The specified key was not found: %s", keyID) + } + return err }) @@ -176,73 +140,7 @@ func (s *Store) GetPrivateKey(ctx context.Context, keyID string) (crypto.Signer, return nil, signingkeys.ErrSigningKeyNotFound.Errorf("The specified key was not found: %s", keyID) } - signKey, err := s.decodePrivateKey(ctx, &key) - if err != nil { - return nil, err - } - - return signKey, nil -} - -func (s *Store) encodePrivateKey(ctx context.Context, privateKey crypto.Signer) ([]byte, error) { - // Encode private key to binary format - pKeyBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) - if err != nil { - return nil, err - } - - // Encode private key to PEM format - privateKeyPEM := pem.EncodeToMemory(&pem.Block{ - Type: "PRIVATE KEY", - Bytes: pKeyBytes, - }) - - encrypted, err := s.secretsService.Encrypt(ctx, privateKeyPEM, secrets.WithoutScope()) - if err != nil { - return nil, err - } - - encoded := make([]byte, base64.StdEncoding.EncodedLen(len(encrypted))) - base64.StdEncoding.Encode(encoded, encrypted) - return encoded, nil -} - -func (s *Store) decodePrivateKey(ctx context.Context, signingKey *SigningKey) (crypto.Signer, error) { - // Bail out if empty string since it'll cause a segfault in Decrypt - if len(signingKey.PrivateKey) == 0 { - return nil, errors.New("private key is empty") - } - - payload := make([]byte, base64.StdEncoding.DecodedLen(len(signingKey.PrivateKey))) - _, err := base64.StdEncoding.Decode(payload, signingKey.PrivateKey) - if err != nil { - return nil, err - } - - decrypted, err := s.secretsService.Decrypt(ctx, payload) - if err != nil { - return nil, err - } - - block, _ := pem.Decode(decrypted) - if block == nil { - return nil, errors.New("failed to decode private key PEM") - } - - if block.Type != "PRIVATE KEY" { - return nil, errors.New("invalid block type") - } - - parsedKey, err := x509.ParsePKCS8PrivateKey(block.Bytes) - if err != nil { - return nil, err - } - - assertedKey, ok := parsedKey.(crypto.Signer) - if !ok { - return nil, errors.New("failed to assert private key as crypto.Signer") - } - return assertedKey, nil + return &key, nil } // cleanupExpiredKeys removes expired keys from the database that have expired more than 61 days ago diff --git a/pkg/services/signingkeys/signingkeystore/store_test.go b/pkg/services/signingkeys/signingkeystore/store_test.go index 01375fad4c7..54e587dcd40 100644 --- a/pkg/services/signingkeys/signingkeystore/store_test.go +++ b/pkg/services/signingkeys/signingkeystore/store_test.go @@ -2,204 +2,81 @@ package signingkeystore import ( "context" - "crypto" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rand" - "crypto/rsa" "testing" "time" - "github.com/go-jose/go-jose/v3" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/grafana/grafana/pkg/infra/db" - "github.com/grafana/grafana/pkg/services/secrets/fakes" "github.com/grafana/grafana/pkg/services/signingkeys" ) func TestIntegrationSigningKeyStore(t *testing.T) { - ctx := context.Background() - - testCases := []struct { - name string - keyFunc func() (crypto.Signer, error) - keyID string - alg jose.SignatureAlgorithm - expected jose.JSONWebKey - }{ - { - name: "RSA key", - keyFunc: func() (crypto.Signer, error) { - return rsa.GenerateKey(rand.Reader, 2048) - }, - keyID: "test-rsa-key", - alg: jose.RS256, - expected: jose.JSONWebKey{ - Key: &rsa.PublicKey{}, - Algorithm: "RS256", - KeyID: "test-rsa-key", - Use: "sig", - }, - }, - { - name: "Elliptic Curve key", - keyFunc: func() (crypto.Signer, error) { - return ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - }, - keyID: "test-ec-key", - alg: jose.ES256, - expected: jose.JSONWebKey{ - Key: &ecdsa.PublicKey{}, - Algorithm: "ES256", - KeyID: "test-ec-key", - Use: "sig", - }, - }, + if testing.Short() { + t.Skip("skipping integration test") } - for _, tc := range testCases { - dbStore := db.InitTestDB(t) - secretSvc := fakes.NewFakeSecretsService() - store := NewSigningKeyStore(dbStore, secretSvc) - - t.Run(tc.name, func(t *testing.T) { - key, err := tc.keyFunc() - assert.NoError(t, err) - - _, err = store.AddPrivateKey(ctx, tc.keyID, tc.alg, key, nil, true) - assert.NoError(t, err) - - retrievedKey, err := store.GetPrivateKey(ctx, tc.keyID) - require.NoError(t, err) - - assert.Equal(t, key.Public(), retrievedKey.Public()) - - jwks, err := store.GetJWKS(ctx) - assert.NoError(t, err) - - require.Len(t, jwks.Keys, 1) - assert.Equal(t, key.Public(), jwks.Keys[0].Key) - assert.Equal(t, tc.expected.Algorithm, jwks.Keys[0].Algorithm) - assert.Equal(t, tc.expected.KeyID, jwks.Keys[0].KeyID) - assert.Equal(t, tc.expected.Use, jwks.Keys[0].Use) - }) - } -} - -func TestIntegrationAddPrivateKey(t *testing.T) { - ctx := context.Background() - - dbStore := db.InitTestDB(t) - secretSvc := fakes.NewFakeSecretsService() - store := NewSigningKeyStore(dbStore, secretSvc) - - key1 := generateRSAKey(t) - key2 := generateECKey(t) - key3 := generateECKey(t) - - testCases := []struct { - name string - keyID string - alg jose.SignatureAlgorithm - privateKey crypto.Signer - expiresAt *time.Time - force bool - expectedErr error - expectedKey crypto.Signer - expectedGot crypto.Signer - }{ - { - name: "Add new private key", - keyID: "test-key-1", - alg: jose.RS256, - privateKey: key1, - force: false, - expectedKey: key1, - expectedGot: key1, - }, - { - name: "Add new private key with expiration", - keyID: "test-key-2", - alg: jose.ES256, - privateKey: key2, - expiresAt: &[]time.Time{time.Now().Add(24 * time.Hour)}[0], - force: false, - expectedKey: key2, - expectedGot: key2, - }, - { - name: "Fail to replace unexpired key", - keyID: "test-key-1", - alg: jose.RS256, - privateKey: key3, - expiresAt: &[]time.Time{time.Now().Add(-24 * time.Hour)}[0], - force: false, - expectedErr: signingkeys.ErrSigningKeyAlreadyExists, - expectedKey: key1, - expectedGot: key1, - }, - { - name: "Replace key1 private key with force, already expired", - keyID: "test-key-1", - alg: jose.ES256, - privateKey: key3, - expiresAt: &[]time.Time{time.Now().Add(-24 * time.Hour)}[0], - force: true, - expectedKey: nil, - expectedGot: key3, - }, - { - name: "Replace key1 private key with no force, is expired", - keyID: "test-key-1", - alg: jose.ES256, - privateKey: key1, - expiresAt: &[]time.Time{time.Now().Add(24 * time.Hour)}[0], - force: false, - expectedKey: nil, - expectedGot: key1, - }, + setup := func() (context.Context, *Store) { + return context.Background(), NewSigningKeyStore(db.InitTestDB(t)) } - _, exists := store.localCache.Get(cleanupRateLimitKey) - require.False(t, exists) + t.Run("Should successfully add new singing key", func(_ *testing.T) { + ctx, store := setup() + key, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", AddedAt: time.Now().UTC(), PrivateKey: []byte{}}, false) + require.NoError(t, err) + assert.Equal(t, "1", key.KeyID) + }) - for _, tc := range testCases { - t.Run(tc.name, func(t *testing.T) { - got, err := store.AddPrivateKey(ctx, tc.keyID, tc.alg, tc.privateKey, tc.expiresAt, tc.force) - if tc.expectedErr != nil { - assert.ErrorIs(t, err, tc.expectedErr) - } else { - assert.NoError(t, err) - } + t.Run("Should return old key if already exists", func(_ *testing.T) { + ctx, store := setup() + key, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().UTC()}, false) + require.NoError(t, err) + assert.Equal(t, "1", key.KeyID) - if tc.expectedGot != nil { - assert.Equal(t, tc.expectedGot.Public(), got.Public()) - } else { - assert.Nil(t, got) - } + // try to add the same key again with a different AddedAt + key2, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().Add(10 * time.Minute).UTC()}, false) + require.ErrorIs(t, err, signingkeys.ErrSigningKeyAlreadyExists) + assert.Equal(t, "1", key2.KeyID) + }) - if tc.expectedKey != nil { - retrievedKey, err := store.GetPrivateKey(ctx, tc.keyID) - assert.NoError(t, err) - assert.Equal(t, tc.expectedKey.Public(), retrievedKey.Public()) - } - }) - } + t.Run("Should update old key when force is true", func(t *testing.T) { + ctx, store := setup() + key, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().UTC()}, false) + require.NoError(t, err) + assert.Equal(t, "1", key.KeyID) - _, exists = store.localCache.Get(cleanupRateLimitKey) - require.True(t, exists) -} - -func generateRSAKey(t *testing.T) *rsa.PrivateKey { - key, err := rsa.GenerateKey(rand.Reader, 2048) - require.NoError(t, err) - return key -} - -func generateECKey(t *testing.T) *ecdsa.PrivateKey { - key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - require.NoError(t, err) - return key + // try to add the same key again with a different AddedAt and force is true + key2, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().Add(10 * time.Minute).UTC()}, true) + require.NoError(t, err) + assert.Equal(t, "1", key2.KeyID) + assert.NotEqual(t, key.AddedAt, key2.AddedAt) + }) + + t.Run("Should update old key when expired", func(t *testing.T) { + ctx, store := setup() + key, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().UTC(), ExpiresAt: &time.Time{}}, false) + require.NoError(t, err) + assert.Equal(t, "1", key.KeyID) + + // try to add the same key again with a different AddedAt and force is false + key2, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().Add(10 * time.Minute).UTC()}, false) + require.NoError(t, err) + assert.Equal(t, "1", key2.KeyID) + assert.NotEqual(t, key.AddedAt, key2.AddedAt) + }) + + t.Run("List should return all keys that are not expired", func(t *testing.T) { + ctx, store := setup() + _, err := store.Add(ctx, &signingkeys.SigningKey{KeyID: "1", PrivateKey: []byte{}, AddedAt: time.Now().UTC()}, false) + require.NoError(t, err) + _, err = store.Add(ctx, &signingkeys.SigningKey{KeyID: "2", PrivateKey: []byte{}, AddedAt: time.Now().UTC(), ExpiresAt: &time.Time{}}, false) + require.NoError(t, err) + _, err = store.Add(ctx, &signingkeys.SigningKey{KeyID: "3", PrivateKey: []byte{}, AddedAt: time.Now().UTC()}, false) + require.NoError(t, err) + + keys, err := store.List(ctx) + require.NoError(t, err) + require.Len(t, keys, 2) + }) }