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:
Mihai Doarna
2025-10-08 18:04:24 +03:00
committed by GitHub
parent 2f7d32ab55
commit 3f40947cd6
5 changed files with 360 additions and 0 deletions
+2
View File
@@ -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, &currentSettings)
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)