mirror of
https://github.com/grafana/grafana.git
synced 2026-08-11 05:34:53 -05:00
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
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user