IAM: Add a local cache for the SSO settings service (#112152)
* add a local cache for the SSO settings service * add unit tests for GetForProviderFromCache() func * append sso settings on update if not found
This commit is contained in:
@@ -22,6 +22,8 @@ type Service interface {
|
||||
ListWithRedactedSecrets(ctx context.Context) ([]*models.SSOSettings, error)
|
||||
// GetForProvider returns the SSO settings for a given provider (DB or config file)
|
||||
GetForProvider(ctx context.Context, provider string) (*models.SSOSettings, error)
|
||||
// GetForProviderFromCache returns the SSO settings for a given provider from cache. It falls back to GetForProvider if the settings are not in the cache.
|
||||
GetForProviderFromCache(ctx context.Context, provider string) (*models.SSOSettings, error)
|
||||
// GetForProviderWithRedactedSecrets returns the SSO settings for a given provider (DB or config file) with secret values redacted
|
||||
GetForProviderWithRedactedSecrets(ctx context.Context, provider string) (*models.SSOSettings, error)
|
||||
// Upsert creates or updates the SSO settings for a given provider
|
||||
|
||||
@@ -5,7 +5,9 @@ import (
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
@@ -43,6 +45,8 @@ type Service struct {
|
||||
providersList []string
|
||||
configurableProviders map[string]bool
|
||||
reloadables map[string]ssosettings.Reloadable
|
||||
cachedSSOSettings []*models.SSOSettings
|
||||
cacheMutex sync.RWMutex
|
||||
}
|
||||
|
||||
func ProvideService(cfg *setting.Cfg, sqlStore db.DB, ac ac.AccessControl,
|
||||
@@ -86,6 +90,7 @@ func ProvideService(cfg *setting.Cfg, sqlStore db.DB, ac ac.AccessControl,
|
||||
configurableProviders: configurableProviders,
|
||||
reloadables: make(map[string]ssosettings.Reloadable),
|
||||
settingsProvider: settingsProvider,
|
||||
cachedSSOSettings: make([]*models.SSOSettings, 0),
|
||||
}
|
||||
|
||||
usageStats.RegisterMetricsFunc(svc.getUsageStats)
|
||||
@@ -120,6 +125,28 @@ func (s *Service) GetForProvider(ctx context.Context, provider string) (*models.
|
||||
return s.mergeSSOSettings(dbSettings, systemSettings), nil
|
||||
}
|
||||
|
||||
func (s *Service) GetForProviderFromCache(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
s.cacheMutex.RLock()
|
||||
defer s.cacheMutex.RUnlock()
|
||||
|
||||
for _, setting := range s.cachedSSOSettings {
|
||||
if setting.Provider == provider {
|
||||
return &models.SSOSettings{
|
||||
Provider: setting.Provider,
|
||||
Source: setting.Source,
|
||||
Settings: deepCopyMap(setting.Settings),
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
// If settings are not in the cache, we return them from the database if the provider is valid
|
||||
if slices.Contains(s.providersList, provider) {
|
||||
return s.GetForProvider(ctx, provider)
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (s *Service) GetForProviderWithRedactedSecrets(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
if !s.isProviderConfigurable(provider) {
|
||||
return nil, ssosettings.ErrNotConfigurable
|
||||
@@ -160,6 +187,8 @@ func (s *Service) List(ctx context.Context) ([]*models.SSOSettings, error) {
|
||||
result = append(result, s.mergeSSOSettings(dbSettings, fallbackSettings))
|
||||
}
|
||||
|
||||
s.setCachedSSOSettings(result)
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
@@ -272,6 +301,8 @@ func (s *Service) Delete(ctx context.Context, provider string) error {
|
||||
}
|
||||
|
||||
func (s *Service) reload(reloadable ssosettings.Reloadable, provider string, currentSettings models.SSOSettings) {
|
||||
s.updateCachedSSOSettings(provider, ¤tSettings)
|
||||
|
||||
err := reloadable.Reload(context.Background(), currentSettings)
|
||||
if err != nil {
|
||||
s.metrics.reloadFailures.WithLabelValues(provider).Inc()
|
||||
@@ -629,3 +660,25 @@ func deepCopySlice(s []any) []any {
|
||||
|
||||
return newSlice
|
||||
}
|
||||
|
||||
func (s *Service) setCachedSSOSettings(settings []*models.SSOSettings) {
|
||||
s.cacheMutex.Lock()
|
||||
defer s.cacheMutex.Unlock()
|
||||
|
||||
s.cachedSSOSettings = settings
|
||||
}
|
||||
|
||||
func (s *Service) updateCachedSSOSettings(provider string, settings *models.SSOSettings) {
|
||||
s.cacheMutex.Lock()
|
||||
defer s.cacheMutex.Unlock()
|
||||
|
||||
for i := range s.cachedSSOSettings {
|
||||
if s.cachedSSOSettings[i].Provider == provider {
|
||||
s.cachedSSOSettings[i] = settings
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Provider not found, append new settings
|
||||
s.cachedSSOSettings = append(s.cachedSSOSettings, settings)
|
||||
}
|
||||
|
||||
@@ -367,6 +367,273 @@ func TestService_GetForProvider(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_GetForProviderFromCache(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
testCases := []struct {
|
||||
name string
|
||||
provider string
|
||||
setup func(env testEnv)
|
||||
want *models.SSOSettings
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "should return successfully from cache",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.service.cachedSSOSettings = []*models.SSOSettings{
|
||||
{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "client_id",
|
||||
"client_secret": "secret",
|
||||
},
|
||||
Source: models.DB,
|
||||
},
|
||||
}
|
||||
},
|
||||
want: &models.SSOSettings{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "client_id",
|
||||
"client_secret": "secret",
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should return successfully from database if not in cache",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.service.cachedSSOSettings = []*models.SSOSettings{
|
||||
{
|
||||
Provider: "google",
|
||||
Settings: map[string]any{"enabled": true},
|
||||
Source: models.DB,
|
||||
},
|
||||
}
|
||||
env.store.ExpectedSSOSetting = &models.SSOSettings{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "client_id",
|
||||
"client_secret": base64.RawStdEncoding.EncodeToString([]byte("client_secret")),
|
||||
},
|
||||
Source: models.DB,
|
||||
}
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_secret"), mock.Anything).Return([]byte("decrypted-client-secret"), nil).Once()
|
||||
},
|
||||
want: &models.SSOSettings{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "client_id",
|
||||
"client_secret": "decrypted-client-secret",
|
||||
},
|
||||
Source: models.DB,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should return nil if provider is not valid",
|
||||
provider: "invalid",
|
||||
setup: func(env testEnv) {
|
||||
env.service.cachedSSOSettings = []*models.SSOSettings{
|
||||
{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{"enabled": true},
|
||||
Source: models.DB,
|
||||
},
|
||||
}
|
||||
},
|
||||
want: nil,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should return error if store returns an error",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedError = fmt.Errorf("error")
|
||||
},
|
||||
want: nil,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
// create a local copy of "tc" to allow concurrent access within tests to the different items of testCases,
|
||||
// otherwise it would be like a moving pointer while tests run in parallel
|
||||
tc := tc
|
||||
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, true, false, true)
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
}
|
||||
|
||||
actual, err := env.service.GetForProviderFromCache(context.Background(), tc.provider)
|
||||
|
||||
if tc.wantErr {
|
||||
require.Error(t, err)
|
||||
return
|
||||
}
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tc.want, actual)
|
||||
|
||||
env.secrets.AssertExpectations(t)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("should return settings from cache after calling List", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, true, false, true)
|
||||
env.store.ExpectedSSOSettings = []*models.SSOSettings{
|
||||
{
|
||||
Provider: "github",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "github_client_id",
|
||||
"client_secret": base64.RawStdEncoding.EncodeToString([]byte("client_secret")),
|
||||
},
|
||||
Source: models.DB,
|
||||
},
|
||||
{
|
||||
Provider: "okta",
|
||||
Settings: map[string]any{
|
||||
"enabled": false,
|
||||
"client_id": "okta_client_id",
|
||||
"other_secret": base64.RawStdEncoding.EncodeToString([]byte("other_secret")),
|
||||
},
|
||||
Source: models.DB,
|
||||
},
|
||||
}
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_secret"), mock.Anything).Return([]byte("decrypted-client-secret"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("other_secret"), mock.Anything).Return([]byte("decrypted-other-secret"), nil).Once()
|
||||
|
||||
_, err := env.service.List(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
actual, err := env.service.GetForProviderFromCache(context.Background(), "github")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "github", actual.Provider)
|
||||
require.Equal(t, map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "github_client_id",
|
||||
"client_secret": "decrypted-client-secret",
|
||||
}, actual.Settings)
|
||||
require.Equal(t, env.store.ExpectedSSOSettings[0].Source, actual.Source)
|
||||
|
||||
actual, err = env.service.GetForProviderFromCache(context.Background(), "okta")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "okta", actual.Provider)
|
||||
require.Equal(t, map[string]any{
|
||||
"enabled": false,
|
||||
"client_id": "okta_client_id",
|
||||
"other_secret": "decrypted-other-secret",
|
||||
}, actual.Settings)
|
||||
require.Equal(t, env.store.ExpectedSSOSettings[1].Source, actual.Source)
|
||||
})
|
||||
|
||||
testCasesUpsert := []struct {
|
||||
name string
|
||||
provider string
|
||||
settings []*models.SSOSettings
|
||||
}{
|
||||
{
|
||||
name: "should return settings from cache after upsert if provider is already in cache",
|
||||
provider: "azuread",
|
||||
settings: []*models.SSOSettings{
|
||||
{
|
||||
Provider: "azuread",
|
||||
Settings: map[string]any{"enabled": true},
|
||||
Source: models.DB,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "should return settings from cache after upsert if provider is not in cache",
|
||||
provider: "github",
|
||||
settings: []*models.SSOSettings{},
|
||||
},
|
||||
}
|
||||
for _, tc := range testCasesUpsert {
|
||||
// create a local copy of "tc" to allow concurrent access within tests to the different items of testCases,
|
||||
// otherwise it would be like a moving pointer while tests run in parallel
|
||||
tc := tc
|
||||
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
|
||||
env.service.cachedSSOSettings = tc.settings
|
||||
|
||||
settings := models.SSOSettings{
|
||||
Provider: tc.provider,
|
||||
Settings: map[string]any{
|
||||
"client_id": "client-id",
|
||||
"client_secret": "client-secret",
|
||||
"enabled": true,
|
||||
},
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
reloadable := ssosettingstests.NewMockReloadable(t)
|
||||
reloadable.On("Validate", mock.Anything, settings, mock.Anything, mock.Anything).Return(nil)
|
||||
reloadable.On("Reload", mock.Anything, mock.MatchedBy(func(settings models.SSOSettings) bool {
|
||||
defer wg.Done()
|
||||
return true
|
||||
})).Return(nil).Maybe()
|
||||
env.reloadables[tc.provider] = reloadable
|
||||
|
||||
env.secrets.On("Encrypt", mock.Anything, []byte("client-secret"), mock.Anything).Return([]byte("encrypted-client-secret"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("encrypted-current-client-secret"), mock.Anything).Return([]byte("current-client-secret"), nil).Once()
|
||||
|
||||
env.store.UpsertFn = func(ctx context.Context, settings *models.SSOSettings) error {
|
||||
currentTime := time.Now()
|
||||
settings.ID = "someid"
|
||||
settings.Created = currentTime
|
||||
settings.Updated = currentTime
|
||||
|
||||
env.store.ActualSSOSettings = *settings
|
||||
return nil
|
||||
}
|
||||
|
||||
env.store.GetFn = func(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
return &models.SSOSettings{
|
||||
ID: "someid",
|
||||
Provider: provider,
|
||||
Settings: map[string]any{
|
||||
"client_secret": base64.RawStdEncoding.EncodeToString([]byte("encrypted-current-client-secret")),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
err := env.service.Upsert(context.Background(), &settings, &user.SignedInUser{})
|
||||
require.NoError(t, err)
|
||||
|
||||
wg.Wait()
|
||||
|
||||
actual, err := env.service.GetForProviderFromCache(context.Background(), tc.provider)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tc.provider, actual.Provider)
|
||||
require.Equal(t, map[string]any{
|
||||
"enabled": true,
|
||||
"client_id": "client-id",
|
||||
"client_secret": "client-secret",
|
||||
}, actual.Settings)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -60,6 +60,14 @@ func (f *FakeService) GetForProvider(ctx context.Context, provider string) (*mod
|
||||
return f.ExpectedSSOSetting, f.ExpectedError
|
||||
}
|
||||
|
||||
func (f *FakeService) GetForProviderFromCache(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
if f.GetForProviderFn != nil {
|
||||
return f.GetForProviderFn(ctx, provider)
|
||||
}
|
||||
f.ActualProvider = provider
|
||||
return f.ExpectedSSOSetting, f.ExpectedError
|
||||
}
|
||||
|
||||
func (f *FakeService) GetForProviderWithRedactedSecrets(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
if f.GetForProviderWithRedactedSecretsFn != nil {
|
||||
return f.GetForProviderWithRedactedSecretsFn(ctx, provider)
|
||||
|
||||
@@ -66,6 +66,36 @@ func (_m *MockService) GetForProvider(ctx context.Context, provider string) (*mo
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
// GetForProviderFromCache provides a mock function with given fields: ctx, provider
|
||||
func (_m *MockService) GetForProviderFromCache(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
ret := _m.Called(ctx, provider)
|
||||
|
||||
if len(ret) == 0 {
|
||||
panic("no return value specified for GetForProviderFromCache")
|
||||
}
|
||||
|
||||
var r0 *models.SSOSettings
|
||||
var r1 error
|
||||
if rf, ok := ret.Get(0).(func(context.Context, string) (*models.SSOSettings, error)); ok {
|
||||
return rf(ctx, provider)
|
||||
}
|
||||
if rf, ok := ret.Get(0).(func(context.Context, string) *models.SSOSettings); ok {
|
||||
r0 = rf(ctx, provider)
|
||||
} else {
|
||||
if ret.Get(0) != nil {
|
||||
r0 = ret.Get(0).(*models.SSOSettings)
|
||||
}
|
||||
}
|
||||
|
||||
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
|
||||
r1 = rf(ctx, provider)
|
||||
} else {
|
||||
r1 = ret.Error(1)
|
||||
}
|
||||
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
// GetForProviderWithRedactedSecrets provides a mock function with given fields: ctx, provider
|
||||
func (_m *MockService) GetForProviderWithRedactedSecrets(ctx context.Context, provider string) (*models.SSOSettings, error) {
|
||||
ret := _m.Called(ctx, provider)
|
||||
|
||||
Reference in New Issue
Block a user