Files
teleport/lib/cache/cache_test.go
T
Przemko RobakowskiandZac Bergquist 48e80466c9 Add LinuxDesktop gRPC and backend (#62974)
* Add LinuxDesktop gRPC and backend

* Remove CloneResource

* Review comments

* Fix logins

* Update lib/auth/linuxdesktop/linuxdesktopv1/service.go

Co-authored-by: Zac Bergquist <zac.bergquist@goteleport.com>

* Fix role

---------

Co-authored-by: Zac Bergquist <zac.bergquist@goteleport.com>
2026-04-21 20:51:35 +00:00

2955 lines
102 KiB
Go

/*
* Teleport
* Copyright (C) 2023 Gravitational, Inc.
*
* This program is free software: you can redistribute it and/or modify
* it under the terms of the GNU Affero General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* This program is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU Affero General Public License for more details.
*
* You should have received a copy of the GNU Affero General Public License
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package cache
import (
"context"
"fmt"
"iter"
"log/slog"
"os"
"slices"
"sync"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/google/go-cmp/cmp/cmpopts"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/protobuf/testing/protocmp"
"google.golang.org/protobuf/types/known/timestamppb"
"github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
accessmonitoringrulesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessmonitoringrules/v1"
appauthconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/appauthconfig/v1"
"github.com/gravitational/teleport/api/gen/proto/go/teleport/autoupdate/v1"
beamsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/beams/v1"
clusterconfigpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/clusterconfig/v1"
crownjewelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/crownjewel/v1"
dbobjectv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/dbobject/v1"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1"
identitycenterv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/identitycenter/v1"
kubewaitingcontainerpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/kubewaitingcontainer/v1"
labelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/label/v1"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
notificationsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/notifications/v1"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
provisioningv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/provisioning/v1"
recordingencryptionv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/recordingencryption/v1"
scopedaccessv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/access/v1"
subcav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/subca/v1"
summaryv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/summarizer/v1"
userprovisioningpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/userprovisioning/v2"
usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1"
workloadclusterv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadcluster/v1"
workloadidentityv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/accesslist"
update "github.com/gravitational/teleport/api/types/autoupdate"
"github.com/gravitational/teleport/api/types/clusterconfig"
"github.com/gravitational/teleport/api/types/discoveryconfig"
"github.com/gravitational/teleport/api/types/header"
"github.com/gravitational/teleport/api/types/kubewaitingcontainer"
"github.com/gravitational/teleport/api/types/secreports"
"github.com/gravitational/teleport/api/types/trait"
"github.com/gravitational/teleport/api/types/userloginstate"
"github.com/gravitational/teleport/api/types/userprovisioning"
"github.com/gravitational/teleport/api/types/usertasks"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/entitlements"
"github.com/gravitational/teleport/lib/auth/authcatest"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/backend/memory"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/modules/modulestest"
scopedaccess "github.com/gravitational/teleport/lib/scopes/access"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/local"
"github.com/gravitational/teleport/lib/srv/db/common/databaseobject"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/log/logtest"
)
const eventBufferSize = 1024
func TestMain(m *testing.M) {
modules.SetModules(&modulestest.Modules{
TestFeatures: modules.Features{
Entitlements: map[entitlements.EntitlementKind]modules.EntitlementInfo{
entitlements.Identity: {Enabled: true},
entitlements.AccessLists: {
Enabled: true,
Limit: 2000,
},
},
},
})
logtest.InitLogger(testing.Verbose)
os.Exit(m.Run())
}
// TestNodesDontCacheHighVolumeResources verifies that resources classified as "high volume" aren't
// cached by nodes.
func TestNodesDontCacheHighVolumeResources(t *testing.T) {
t.Parallel()
for _, kind := range ForNode(Config{}).Watches {
require.False(t, isHighVolumeResource(kind.Kind), "resource=%q", kind.Kind)
}
}
// testPack contains pack of
// services used for test run
type testPack struct {
dataDir string
backend *backend.Wrapper
eventsC chan Event
cache *Cache
eventsS *proxyEvents
trustS *local.CA
provisionerS *local.ProvisioningService
clusterConfigS *local.ClusterConfigurationService
usersS *local.IdentityService
accessS *local.AccessService
dynamicAccessS *local.DynamicAccessService
presenceS *local.PresenceService
appSessionS *local.IdentityService
snowflakeSessionS *local.IdentityService
restrictions *local.RestrictionsService
apps *local.AppService
kubernetes *local.KubernetesService
databases *local.DatabaseService
databaseServices *local.DatabaseServicesService
webSessionS types.WebSessionInterface
webTokenS *local.IdentityService
windowsDesktops *local.WindowsDesktopService
dynamicWindowsDesktops *local.DynamicWindowsDesktopService
linuxDesktops *local.LinuxDesktopService
samlIDPServiceProviders *local.SAMLIdPServiceProviderService
userGroups *local.UserGroupService
okta *local.OktaService
integrations *local.IntegrationsService
userTasks *local.UserTasksService
discoveryConfigs *local.DiscoveryConfigService
userLoginStates *local.UserLoginStateService
secReports *local.SecReportsService
accessLists *local.AccessListService
kubeWaitingContainers *local.KubeWaitingContainerService
notifications *local.NotificationsService
accessMonitoringRules *local.AccessMonitoringRulesService
crownJewels *local.CrownJewelsService
databaseObjects *local.DatabaseObjectService
spiffeFederations *local.SPIFFEFederationService
staticHostUsers *local.StaticHostUserService
autoUpdateService *local.AutoUpdateService
provisioningStates *local.ProvisioningStateService
identityCenter *local.IdentityCenterService
pluginStaticCredentials *local.PluginStaticCredentialsService
gitServers *local.GitServerService
workloadIdentity *local.WorkloadIdentityService
beams *local.BeamService
healthCheckConfig *local.HealthCheckConfigService
botInstanceService *local.BotInstanceService
recordingEncryption *local.RecordingEncryptionService
plugin *local.PluginsService
appAuthConfigs *local.AppAuthConfigService
workloadClusters *local.WorkloadClusterService
summarizer *local.SummarizerService
subCA *local.SubCAService
}
// resourceOps contains helpers to modify the state of either types.Resource or types.Resource153 which
// have a slightly different interface.
type resourceOps[T any] struct {
Name func(T) string
Setup func(T)
Modify func(T)
cmpOpts []cmp.Option
}
func defaultResourceOps[T types.Resource]() *resourceOps[T] {
return &resourceOps[T]{
Modify: func(t T) {
// types.Resource metadata is immutable, modify expiry only.
if t.Expiry().IsZero() {
t.SetExpiry(time.Now().Add(30 * time.Minute))
} else {
t.SetExpiry(t.Expiry().Add(30 * time.Minute))
}
},
Name: func(t T) string { return t.GetName() },
Setup: func(t T) {
// types.Resource metadata is immutable, modify expiry only.
if t.Expiry().IsZero() {
t.SetExpiry(time.Now().Add(30 * time.Minute))
} else {
t.SetExpiry(t.Expiry().Add(30 * time.Minute))
}
},
cmpOpts: []cmp.Option{
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
cmpopts.IgnoreFields(header.Metadata{}, "Revision"),
},
}
}
func defaultResource153Ops[T types.Resource153]() *resourceOps[T] {
return &resourceOps[T]{
Setup: func(t T) {
metadata := t.GetMetadata()
if metadata.Expires == nil {
metadata.Expires = timestamppb.New(time.Now().Add(30 * time.Minute))
} else {
expiry := metadata.Expires.AsTime()
metadata.Expires = timestamppb.New(expiry.Add(30 * time.Minute))
}
metadata.Labels = map[string]string{"label": "value1"}
},
Modify: func(t T) {
metadata := t.GetMetadata()
if metadata.Expires == nil {
metadata.Expires = timestamppb.New(time.Now().Add(30 * time.Minute))
} else {
expiry := metadata.Expires.AsTime()
metadata.Expires = timestamppb.New(expiry.Add(30 * time.Minute))
}
metadata.Labels["label"] = "value2"
},
Name: func(t T) string { return t.GetMetadata().GetName() },
cmpOpts: []cmp.Option{
protocmp.IgnoreFields(&headerv1.Metadata{}, "revision"),
protocmp.Transform(),
cmpopts.EquateEmpty(),
},
}
}
// getAllAdapter adapts collection getters that do not support pagination to conform to [testFuncs] interface
// TODO(okraport): delete this once all APIs are paginated.
func getAllAdapter[T any](fn func(context.Context) ([]T, error)) func(context.Context, int, string) ([]T, string, error) {
return func(ctx context.Context, _ int, _ string) ([]T, string, error) {
out, err := fn(ctx)
return out, "", trace.Wrap(err)
}
}
// singletonListAdapter adapts a singleton getter to conform to [testFuncs] interface
func singletonListAdapter[T any](fn func(context.Context) (T, error)) func(context.Context, int, string) ([]T, string, error) {
return func(ctx context.Context, _ int, _ string) ([]T, string, error) {
out, err := fn(ctx)
if err != nil {
if trace.IsNotFound(err) {
return nil, "", nil
}
return nil, "", trace.Wrap(err)
}
return []T{out}, "", nil
}
}
// testFuncs are functions to support testing an object in a cache.
type testFuncs[T any] struct {
newResource func(string) (T, error)
create func(context.Context, T) error
list func(context.Context, int, string) ([]T, string, error)
Range func(context.Context, string, string) iter.Seq2[T, error]
cacheGet func(context.Context, string) (T, error)
cacheList func(context.Context, int, string) ([]T, string, error)
cacheRange func(context.Context, string, string) iter.Seq2[T, error]
update func(context.Context, T) error
delete func(context.Context, string) error
deleteAll func(context.Context) error
resource *resourceOps[T]
}
func (f *testFuncs[T]) listAll(ctx context.Context) ([]T, error) {
return stream.Collect(clientutils.Resources(ctx, f.list))
}
func (f *testFuncs[T]) cacheListAll(ctx context.Context) ([]T, error) {
return stream.Collect(clientutils.Resources(ctx, f.cacheList))
}
func (t *testPack) Close() {
var errors []error
if t.backend != nil {
errors = append(errors, t.backend.Close())
}
if t.cache != nil {
errors = append(errors, t.cache.Close())
}
if err := trace.NewAggregate(errors...); err != nil {
slog.WarnContext(context.Background(), "Failed to close", "error", err)
}
}
func newPackForAuth(t *testing.T) *testPack {
return newTestPack(t, ForAuth)
}
func newTestPack(t *testing.T, setupConfig SetupConfigFn, opts ...packOption) *testPack {
pack, err := newPack(t, setupConfig, opts...)
require.NoError(t, err)
return pack
}
func NewTestPackWithoutCache(t *testing.T) *testPack {
pack, err := newPackWithoutCache(t.TempDir())
require.NoError(t, err)
return pack
}
type packCfg struct {
ignoreKinds []types.WatchKind
}
type packOption func(cfg *packCfg)
// ignoreKinds specifies the list of kinds that should be removed from the watch request by eventsProxy
// to simulate cache resource type rejection due to version incompatibility.
func ignoreKinds(kinds []types.WatchKind) packOption {
return func(cfg *packCfg) {
cfg.ignoreKinds = kinds
}
}
// newPackWithoutCache returns a new test pack without creating cache
func newPackWithoutCache(dir string, opts ...packOption) (*testPack, error) {
ctx := context.Background()
var cfg packCfg
for _, opt := range opts {
opt(&cfg)
}
p := &testPack{
dataDir: dir,
}
bk, err := memory.New(memory.Config{
Context: ctx,
Mirror: true,
})
if err != nil {
return nil, trace.Wrap(err)
}
p.backend = backend.NewWrapper(bk)
p.eventsC = make(chan Event, eventBufferSize)
clusterConfig, err := local.NewClusterConfigurationService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
idService, err := local.NewTestIdentityService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
dynamicWindowsDesktopService, err := local.NewDynamicWindowsDesktopService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.trustS = local.NewCAService(p.backend)
p.clusterConfigS = clusterConfig
p.provisionerS = local.NewProvisioningService(p.backend)
p.eventsS = newProxyEvents(local.NewEventsService(p.backend), cfg.ignoreKinds)
p.presenceS = local.NewPresenceService(p.backend)
p.usersS = idService
p.accessS = local.NewAccessService(p.backend)
p.dynamicAccessS = local.NewDynamicAccessService(p.backend)
p.appSessionS = idService
p.webSessionS = idService.WebSessions()
p.snowflakeSessionS = idService
p.webTokenS = idService
p.restrictions = local.NewRestrictionsService(p.backend)
p.apps = local.NewAppService(p.backend)
p.kubernetes = local.NewKubernetesService(p.backend)
p.databases = local.NewDatabasesService(p.backend)
p.databaseServices = local.NewDatabaseServicesService(p.backend)
p.windowsDesktops = local.NewWindowsDesktopService(p.backend)
p.linuxDesktops, err = local.NewLinuxDesktopService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.dynamicWindowsDesktops = dynamicWindowsDesktopService
p.samlIDPServiceProviders, err = local.NewSAMLIdPServiceProviderService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.userGroups, err = local.NewUserGroupService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
oktaSvc, err := local.NewOktaService(p.backend, p.backend.Clock())
if err != nil {
return nil, trace.Wrap(err)
}
p.okta = oktaSvc
igSvc, err := local.NewIntegrationsService(p.backend, local.WithIntegrationsServiceCacheMode(true))
if err != nil {
return nil, trace.Wrap(err)
}
p.integrations = igSvc
userTasksSvc, err := local.NewUserTasksService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.userTasks = userTasksSvc
dcSvc, err := local.NewDiscoveryConfigService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.discoveryConfigs = dcSvc
ulsSvc, err := local.NewUserLoginStateService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.userLoginStates = ulsSvc
secReportsSvc, err := local.NewSecReportsService(p.backend, p.backend.Clock())
if err != nil {
return nil, trace.Wrap(err)
}
p.secReports = secReportsSvc
accessListsSvc, err := local.NewAccessListServiceV2(local.AccessListServiceConfig{
Backend: p.backend,
Modules: modulestest.EnterpriseModules(),
})
if err != nil {
return nil, trace.Wrap(err)
}
p.accessLists = accessListsSvc
accessMonitoringRuleService, err := local.NewAccessMonitoringRulesService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.accessMonitoringRules = accessMonitoringRuleService
crownJewelsSvc, err := local.NewCrownJewelsService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.crownJewels = crownJewelsSvc
spiffeFederationsSvc, err := local.NewSPIFFEFederationService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.spiffeFederations = spiffeFederationsSvc
workloadIdentitySvc, err := local.NewWorkloadIdentityService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.workloadIdentity = workloadIdentitySvc
beamSvc, err := local.NewBeamService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.beams = beamSvc
databaseObjectsSvc, err := local.NewDatabaseObjectService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.databaseObjects = databaseObjectsSvc
kubeWaitingContSvc, err := local.NewKubeWaitingContainerService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.kubeWaitingContainers = kubeWaitingContSvc
notificationsSvc, err := local.NewNotificationsService(p.backend, p.backend.Clock())
if err != nil {
return nil, trace.Wrap(err)
}
p.notifications = notificationsSvc
staticHostUserService, err := local.NewStaticHostUserService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.staticHostUsers = staticHostUserService
p.autoUpdateService, err = local.NewAutoUpdateService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.provisioningStates, err = local.NewProvisioningStateService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.identityCenter, err = local.NewIdentityCenterService(local.IdentityCenterServiceConfig{
Backend: p.backend,
})
if err != nil {
return nil, trace.Wrap(err)
}
p.pluginStaticCredentials, err = local.NewPluginStaticCredentialsService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.gitServers, err = local.NewGitServerService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.healthCheckConfig, err = local.NewHealthCheckConfigService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.botInstanceService, err = local.NewBotInstanceService(p.backend, p.backend.Clock())
if err != nil {
return nil, trace.Wrap(err)
}
p.recordingEncryption, err = local.NewRecordingEncryptionService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.plugin = local.NewPluginsService(p.backend)
p.appAuthConfigs, err = local.NewAppAuthConfigService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
workloadClusterSvc, err := local.NewWorkloadClusterService(p.backend)
if err != nil {
return nil, trace.Wrap(err)
}
p.workloadClusters = workloadClusterSvc
p.summarizer, err = local.NewSummarizerService(local.SummarizerServiceConfig{
Backend: p.backend,
})
if err != nil {
return nil, trace.Wrap(err)
}
p.subCA, err = local.NewSubCAService(local.SubCAServiceParams{
Backend: p.backend,
})
if err != nil {
return nil, trace.Wrap(err)
}
return p, nil
}
// newPack returns a new test pack or fails the test on error
func newPack(t testing.TB, setupConfig func(c Config) Config, opts ...packOption) (*testPack, error) {
t.Helper()
ctx := t.Context()
p, err := newPackWithoutCache(t.TempDir(), opts...)
if err != nil {
return nil, trace.Wrap(err)
}
p.cache, err = New(setupConfig(Config{
Context: ctx,
Events: p.eventsS,
ClusterConfig: p.clusterConfigS,
Provisioner: p.provisionerS,
Trust: p.trustS,
Users: p.usersS,
Access: p.accessS,
DynamicAccess: p.dynamicAccessS,
Presence: p.presenceS,
AppSession: p.appSessionS,
WebSession: p.webSessionS,
WebToken: p.webTokenS,
Beams: p.beams,
SnowflakeSession: p.snowflakeSessionS,
Restrictions: p.restrictions,
Apps: p.apps,
Kubernetes: p.kubernetes,
DatabaseServices: p.databaseServices,
Databases: p.databases,
WindowsDesktops: p.windowsDesktops,
DynamicWindowsDesktops: p.dynamicWindowsDesktops,
LinuxDesktops: p.linuxDesktops,
SAMLIdPServiceProviders: p.samlIDPServiceProviders,
UserGroups: p.userGroups,
Okta: p.okta,
Integrations: p.integrations,
UserTasks: p.userTasks,
DiscoveryConfigs: p.discoveryConfigs,
UserLoginStates: p.userLoginStates,
SecReports: p.secReports,
AccessLists: p.accessLists,
KubeWaitingContainers: p.kubeWaitingContainers,
Notifications: p.notifications,
AccessMonitoringRules: p.accessMonitoringRules,
CrownJewels: p.crownJewels,
SPIFFEFederations: p.spiffeFederations,
DatabaseObjects: p.databaseObjects,
StaticHostUsers: p.staticHostUsers,
AutoUpdateService: p.autoUpdateService,
ProvisioningStates: p.provisioningStates,
IdentityCenter: p.identityCenter,
PluginStaticCredentials: p.pluginStaticCredentials,
GitServers: p.gitServers,
HealthCheckConfig: p.healthCheckConfig,
WorkloadIdentity: p.workloadIdentity,
BotInstanceService: p.botInstanceService,
RecordingEncryption: p.recordingEncryption,
Plugin: p.plugin,
MaxRetryPeriod: 200 * time.Millisecond,
EventsC: p.eventsC,
AppAuthConfig: p.appAuthConfigs,
StaticScopedToken: p.clusterConfigS,
WorkloadClusterService: p.workloadClusters,
Summarizer: p.summarizer,
SubCAService: p.subCA,
}))
if err != nil {
return nil, trace.Wrap(err)
}
// Wait for the watcher to start. Note that we do not enforce a timeout on this, as depending on the machine
// the calling test runs on and CPU/scheduler contention the expected time may vary. If the watcher fails to start we will
// timeout due to the top level harness timeout setting instead.
event := <-p.eventsC
if event.Type != WatcherStarted {
return nil, trace.CompareFailed("%q != %q %s", event.Type, WatcherStarted, event)
}
return p, nil
}
// TestWatchers tests watchers connected to the cache,
// verifies that all watchers of the cache will be closed
// if the underlying watcher to the target backend is closed
func TestWatchers(t *testing.T) {
t.Parallel()
ctx := context.Background()
p := newPackForAuth(t)
t.Cleanup(p.Close)
w, err := p.cache.NewWatcher(ctx, types.Watch{Kinds: []types.WatchKind{
{
Kind: types.KindCertAuthority,
Filter: types.CertAuthorityFilter{
types.HostCA: "example.com",
types.UserCA: types.Wildcard,
}.IntoMap(),
},
{
Kind: types.KindAccessRequest,
Filter: map[string]string{
"user": "alice",
},
},
}})
require.NoError(t, err)
t.Cleanup(func() {
require.NoError(t, w.Close())
})
select {
case e := <-w.Events():
require.Equal(t, types.OpInit, e.Type)
case <-time.After(100 * time.Millisecond):
t.Fatalf("Timeout waiting for event.")
}
ca, err := authcatest.NewCA(types.UserCA, "example.com")
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
select {
case e := <-w.Events():
require.Equal(t, types.OpPut, e.Type)
require.Equal(t, types.KindCertAuthority, e.Resource.GetKind())
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for event.")
}
// create an access request that matches the supplied filter
req, err := services.NewAccessRequest("alice", "dictator")
require.NoError(t, err)
req, err = p.dynamicAccessS.CreateAccessRequestV2(ctx, req)
require.NoError(t, err)
select {
case e := <-w.Events():
require.Equal(t, types.OpPut, e.Type)
require.Equal(t, types.KindAccessRequest, e.Resource.GetKind())
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for event.")
}
require.NoError(t, p.dynamicAccessS.DeleteAccessRequest(ctx, req.GetName()))
select {
case e := <-w.Events():
require.Equal(t, types.OpDelete, e.Type)
require.Equal(t, types.KindAccessRequest, e.Resource.GetKind())
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for event.")
}
// create an access request that does not match the supplied filter
req2, err := services.NewAccessRequest("bob", "dictator")
require.NoError(t, err)
// create and then delete the non-matching request.
req2, err = p.dynamicAccessS.CreateAccessRequestV2(ctx, req2)
require.NoError(t, err)
require.NoError(t, p.dynamicAccessS.DeleteAccessRequest(ctx, req2.GetName()))
// because our filter did not match the request, the create event should never
// have been created, meaning that the next event on the pipe is the delete
// event (which cannot be filtered out because username is not visible inside
// a delete event).
select {
case e := <-w.Events():
require.Equal(t, types.OpDelete, e.Type)
require.Equal(t, types.KindAccessRequest, e.Resource.GetKind())
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for event.")
}
// this ca will not be matched by our filter, so the same reasoning applies
// as we upsert it and delete it
filteredCa, err := authcatest.NewCA(types.HostCA, "example.net")
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, filteredCa))
require.NoError(t, p.trustS.DeleteCertAuthority(ctx, filteredCa.GetID()))
select {
case e := <-w.Events():
require.Equal(t, types.OpDelete, e.Type)
require.Equal(t, types.KindCertAuthority, e.Resource.GetKind())
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for event.")
}
// event has arrived, now close the watchers
p.backend.CloseWatchers()
// make sure watcher has been closed
select {
case <-w.Done():
case <-time.After(time.Second):
t.Fatalf("Timeout waiting for close event.")
}
}
func waitForRestart(t *testing.T, eventsC <-chan Event) {
expectEvent(t, eventsC, WatcherStarted)
}
func drainEvents(eventsC <-chan Event) {
for {
select {
case <-eventsC:
default:
return
}
}
}
func expectEvent(t *testing.T, eventsC <-chan Event, expectedEvent string) {
timeC := time.After(5 * time.Second)
for {
select {
case event := <-eventsC:
if event.Type == expectedEvent {
return
}
case <-timeC:
t.Fatalf("Timeout waiting for expected event: %s", expectedEvent)
}
}
}
func unexpectedEvent(t *testing.T, eventsC <-chan Event, unexpectedEvent string) {
timeC := time.After(time.Second)
for {
select {
case event := <-eventsC:
if event.Type == unexpectedEvent {
t.Fatalf("Received unexpected event: %s", unexpectedEvent)
}
case <-timeC:
return
}
}
}
func expectNextEvent(t *testing.T, eventsC <-chan Event, expectedEvent string, skipEvents ...string) {
timeC := time.After(5 * time.Second)
for {
// wait for watcher to restart
select {
case event := <-eventsC:
if slices.Contains(skipEvents, event.Type) {
continue
}
require.Equal(t, expectedEvent, event.Type)
return
case <-timeC:
t.Fatalf("Timeout waiting for expected event: %s", expectedEvent)
}
}
}
// TestCompletenessInit verifies that flaky backends don't cause
// the cache to return partial results during init.
func TestCompletenessInit(t *testing.T) {
t.Parallel()
ctx := context.Background()
const caCount = 100
const inits = 20
p := NewTestPackWithoutCache(t)
t.Cleanup(p.Close)
// put lots of CAs in the backend
for i := range caCount {
ca, err := authcatest.NewCA(types.UserCA, fmt.Sprintf("%d.example.com", i))
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
}
for range inits {
// simulate bad connection to auth server
p.backend.SetReadError(trace.ConnectionProblem(nil, "backend is unavailable"))
p.eventsS.closeWatchers()
var err error
p.cache, err = New(ForAuth(Config{
Context: ctx,
Events: p.eventsS,
ClusterConfig: p.clusterConfigS,
Provisioner: p.provisionerS,
Trust: p.trustS,
Users: p.usersS,
Access: p.accessS,
DynamicAccess: p.dynamicAccessS,
Presence: p.presenceS,
AppSession: p.appSessionS,
WebSession: p.webSessionS,
SnowflakeSession: p.snowflakeSessionS,
WebToken: p.webTokenS,
Beams: p.beams,
Restrictions: p.restrictions,
Apps: p.apps,
Kubernetes: p.kubernetes,
DatabaseServices: p.databaseServices,
Databases: p.databases,
WindowsDesktops: p.windowsDesktops,
DynamicWindowsDesktops: p.dynamicWindowsDesktops,
LinuxDesktops: p.linuxDesktops,
SAMLIdPServiceProviders: p.samlIDPServiceProviders,
UserGroups: p.userGroups,
Okta: p.okta,
Integrations: p.integrations,
UserTasks: p.userTasks,
DiscoveryConfigs: p.discoveryConfigs,
UserLoginStates: p.userLoginStates,
SecReports: p.secReports,
AccessLists: p.accessLists,
KubeWaitingContainers: p.kubeWaitingContainers,
Notifications: p.notifications,
AccessMonitoringRules: p.accessMonitoringRules,
CrownJewels: p.crownJewels,
DatabaseObjects: p.databaseObjects,
SPIFFEFederations: p.spiffeFederations,
StaticHostUsers: p.staticHostUsers,
AutoUpdateService: p.autoUpdateService,
ProvisioningStates: p.provisioningStates,
WorkloadIdentity: p.workloadIdentity,
RecordingEncryption: p.recordingEncryption,
MaxRetryPeriod: 200 * time.Millisecond,
IdentityCenter: p.identityCenter,
PluginStaticCredentials: p.pluginStaticCredentials,
EventsC: p.eventsC,
GitServers: p.gitServers,
HealthCheckConfig: p.healthCheckConfig,
BotInstanceService: p.botInstanceService,
Plugin: p.plugin,
AppAuthConfig: p.appAuthConfigs,
StaticScopedToken: p.clusterConfigS,
WorkloadClusterService: p.workloadClusters,
Summarizer: p.summarizer,
SubCAService: p.subCA,
}))
require.NoError(t, err)
p.backend.SetReadError(nil)
cas, err := p.cache.GetCertAuthorities(ctx, types.UserCA, false)
// we don't actually care whether the cache ever fully constructed
// the CA list. for the purposes of this test, we just care that it
// doesn't return the CA list *unless* it was successfully constructed.
if err == nil {
require.Len(t, cas, caCount)
} else {
require.True(t, trace.IsConnectionProblem(err))
}
require.NoError(t, p.cache.Close())
p.cache = nil
}
}
// TestCompletenessReset verifies that flaky backends don't cause
// the cache to return partial results during reset.
func TestCompletenessReset(t *testing.T) {
t.Parallel()
ctx := context.Background()
const caCount = 100
const resets = 20
p := NewTestPackWithoutCache(t)
t.Cleanup(p.Close)
// put lots of CAs in the backend
for i := range caCount {
ca, err := authcatest.NewCA(types.UserCA, fmt.Sprintf("%d.example.com", i))
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
}
var err error
p.cache, err = New(ForAuth(Config{
Context: ctx,
Events: p.eventsS,
ClusterConfig: p.clusterConfigS,
Provisioner: p.provisionerS,
Trust: p.trustS,
Users: p.usersS,
Access: p.accessS,
DynamicAccess: p.dynamicAccessS,
Presence: p.presenceS,
AppSession: p.appSessionS,
WebSession: p.webSessionS,
SnowflakeSession: p.snowflakeSessionS,
WebToken: p.webTokenS,
Beams: p.beams,
Restrictions: p.restrictions,
Apps: p.apps,
Kubernetes: p.kubernetes,
DatabaseServices: p.databaseServices,
Databases: p.databases,
WindowsDesktops: p.windowsDesktops,
DynamicWindowsDesktops: p.dynamicWindowsDesktops,
LinuxDesktops: p.linuxDesktops,
SAMLIdPServiceProviders: p.samlIDPServiceProviders,
UserGroups: p.userGroups,
Okta: p.okta,
Integrations: p.integrations,
UserTasks: p.userTasks,
DiscoveryConfigs: p.discoveryConfigs,
UserLoginStates: p.userLoginStates,
SecReports: p.secReports,
AccessLists: p.accessLists,
KubeWaitingContainers: p.kubeWaitingContainers,
Notifications: p.notifications,
AccessMonitoringRules: p.accessMonitoringRules,
CrownJewels: p.crownJewels,
DatabaseObjects: p.databaseObjects,
SPIFFEFederations: p.spiffeFederations,
StaticHostUsers: p.staticHostUsers,
AutoUpdateService: p.autoUpdateService,
ProvisioningStates: p.provisioningStates,
IdentityCenter: p.identityCenter,
PluginStaticCredentials: p.pluginStaticCredentials,
WorkloadIdentity: p.workloadIdentity,
RecordingEncryption: p.recordingEncryption,
MaxRetryPeriod: 200 * time.Millisecond,
EventsC: p.eventsC,
GitServers: p.gitServers,
HealthCheckConfig: p.healthCheckConfig,
BotInstanceService: p.botInstanceService,
Plugin: p.plugin,
AppAuthConfig: p.appAuthConfigs,
StaticScopedToken: p.clusterConfigS,
WorkloadClusterService: p.workloadClusters,
Summarizer: p.summarizer,
SubCAService: p.subCA,
}))
require.NoError(t, err)
// verify that CAs are immediately available
cas, err := p.cache.GetCertAuthorities(ctx, types.UserCA, false)
require.NoError(t, err)
require.Len(t, cas, caCount)
for range resets {
// simulate bad connection to auth server
p.backend.SetReadError(trace.ConnectionProblem(nil, "backend is unavailable"))
p.eventsS.closeWatchers()
p.backend.SetReadError(nil)
// load CAs while connection is bad
cas, err := p.cache.GetCertAuthorities(ctx, types.UserCA, false)
// we don't actually care whether the cache ever fully constructed
// the CA list. for the purposes of this test, we just care that it
// doesn't return the CA list *unless* it was successfully constructed.
if err == nil {
require.Len(t, cas, caCount)
} else {
require.True(t, trace.IsConnectionProblem(err))
}
}
}
// TestInitStrategy verifies that cache uses expected init strategy
// of serving backend state when init is taking too long.
func TestInitStrategy(t *testing.T) {
t.Parallel()
for range utils.GetIterations() {
initStrategy(t)
}
}
/*
goos: linux
goarch: amd64
pkg: github.com/gravitational/teleport/lib/cache
cpu: Intel(R) Core(TM) i7-8550U CPU @ 1.80GHz
BenchmarkListResourcesWithSort-8 1 2351035036 ns/op
*/
func BenchmarkListResourcesWithSort(b *testing.B) {
if testing.Short() {
b.Skip("skipping heavy benchmark")
}
p, err := newPack(b, ForAuth)
require.NoError(b, err)
defer p.Close()
ctx := context.Background()
count := 100000
for i := range count {
server := NewServer(types.KindNode, uuid.New().String(), "127.0.0.1:2022", apidefaults.Namespace)
// Set some static and dynamic labels.
server.Metadata.Labels = map[string]string{"os": "mac", "env": "prod", "country": "us", "tier": "frontend"}
server.Spec.CmdLabels = map[string]types.CommandLabelV2{
"version": {Result: "v8"},
"time": {Result: "now"},
}
_, err := p.presenceS.UpsertNode(ctx, server)
require.NoError(b, err)
select {
case event := <-p.eventsC:
require.Equal(b, EventProcessed, event.Type)
case <-time.After(200 * time.Millisecond):
b.Fatalf("timeout waiting for event, iteration=%d", i)
}
}
for _, limit := range []int32{100, 1_000, 10_000, 100_000} {
for _, totalCount := range []bool{true, false} {
b.Run(fmt.Sprintf("limit=%d,needTotal=%t", limit, totalCount), func(b *testing.B) {
for b.Loop() {
resp, err := p.cache.ListResources(ctx, proto.ListResourcesRequest{
ResourceType: types.KindNode,
Namespace: apidefaults.Namespace,
SortBy: types.SortBy{
IsDesc: true,
Field: types.ResourceSpecHostname,
},
// Predicate is the more expensive filter.
PredicateExpression: `search("mac", "frontend") && labels.version == "v8"`,
Limit: limit,
NeedTotalCount: totalCount,
})
require.NoError(b, err)
require.Len(b, resp.Resources, int(limit))
}
})
}
}
}
// TestListResources_NodesTTLVariant verifies that the custom ListNodes impl that we fallback to when
// using ttl-based caching works as expected.
func TestListResources_NodesTTLVariant(t *testing.T) {
t.Parallel()
const nodeCount = 100
const pageSize = 10
var err error
ctx := context.Background()
p, err := newPackWithoutCache(t.TempDir())
require.NoError(t, err)
t.Cleanup(p.Close)
p.cache, err = New(ForAuth(Config{
Context: ctx,
Events: p.eventsS,
ClusterConfig: p.clusterConfigS,
Provisioner: p.provisionerS,
Trust: p.trustS,
Users: p.usersS,
Access: p.accessS,
DynamicAccess: p.dynamicAccessS,
Presence: p.presenceS,
AppSession: p.appSessionS,
WebSession: p.webSessionS,
WebToken: p.webTokenS,
SnowflakeSession: p.snowflakeSessionS,
Beams: p.beams,
Restrictions: p.restrictions,
Apps: p.apps,
Kubernetes: p.kubernetes,
DatabaseServices: p.databaseServices,
Databases: p.databases,
WindowsDesktops: p.windowsDesktops,
DynamicWindowsDesktops: p.dynamicWindowsDesktops,
LinuxDesktops: p.linuxDesktops,
SAMLIdPServiceProviders: p.samlIDPServiceProviders,
UserGroups: p.userGroups,
Okta: p.okta,
Integrations: p.integrations,
UserTasks: p.userTasks,
DiscoveryConfigs: p.discoveryConfigs,
UserLoginStates: p.userLoginStates,
SecReports: p.secReports,
AccessLists: p.accessLists,
KubeWaitingContainers: p.kubeWaitingContainers,
Notifications: p.notifications,
AccessMonitoringRules: p.accessMonitoringRules,
CrownJewels: p.crownJewels,
DatabaseObjects: p.databaseObjects,
SPIFFEFederations: p.spiffeFederations,
StaticHostUsers: p.staticHostUsers,
AutoUpdateService: p.autoUpdateService,
ProvisioningStates: p.provisioningStates,
IdentityCenter: p.identityCenter,
PluginStaticCredentials: p.pluginStaticCredentials,
WorkloadIdentity: p.workloadIdentity,
RecordingEncryption: p.recordingEncryption,
MaxRetryPeriod: 200 * time.Millisecond,
EventsC: p.eventsC,
neverOK: true, // ensure reads are never healthy
GitServers: p.gitServers,
HealthCheckConfig: p.healthCheckConfig,
BotInstanceService: p.botInstanceService,
Plugin: p.plugin,
AppAuthConfig: p.appAuthConfigs,
StaticScopedToken: p.clusterConfigS,
WorkloadClusterService: p.workloadClusters,
Summarizer: p.summarizer,
SubCAService: p.subCA,
}))
require.NoError(t, err)
for range nodeCount {
server := NewServer(types.KindNode, uuid.New().String(), "127.0.0.1:2022", apidefaults.Namespace)
_, err := p.presenceS.UpsertNode(ctx, server)
require.NoError(t, err)
}
time.Sleep(time.Second * 2)
allNodes, err := p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.Len(t, allNodes, nodeCount)
var resources []types.ResourceWithLabels
var listResourcesStartKey string
sortBy := types.SortBy{
Field: types.ResourceMetadataName,
IsDesc: true,
}
require.EventuallyWithT(t, func(t *assert.CollectT) {
resp, err := p.cache.ListResources(ctx, proto.ListResourcesRequest{
Namespace: apidefaults.Namespace,
ResourceType: types.KindNode,
StartKey: listResourcesStartKey,
Limit: int32(pageSize),
SortBy: sortBy,
})
require.NoError(t, err)
resources = append(resources, resp.Resources...)
listResourcesStartKey = resp.NextKey
require.Len(t, resources, nodeCount)
}, 5*time.Second, 100*time.Millisecond)
servers, err := types.ResourcesWithLabels(resources).AsServers()
require.NoError(t, err)
fieldVals, err := types.Servers(servers).GetFieldVals(sortBy.Field)
require.NoError(t, err)
require.IsDecreasing(t, fieldVals)
}
func initStrategy(t *testing.T) {
ctx := context.Background()
p := NewTestPackWithoutCache(t)
t.Cleanup(p.Close)
p.backend.SetReadError(trace.ConnectionProblem(nil, "backend is out"))
var err error
p.cache, err = New(ForAuth(Config{
Context: ctx,
Events: p.eventsS,
ClusterConfig: p.clusterConfigS,
Provisioner: p.provisionerS,
Trust: p.trustS,
Users: p.usersS,
Access: p.accessS,
DynamicAccess: p.dynamicAccessS,
Presence: p.presenceS,
AppSession: p.appSessionS,
SnowflakeSession: p.snowflakeSessionS,
WebSession: p.webSessionS,
WebToken: p.webTokenS,
Beams: p.beams,
Restrictions: p.restrictions,
Apps: p.apps,
Kubernetes: p.kubernetes,
DatabaseServices: p.databaseServices,
Databases: p.databases,
WindowsDesktops: p.windowsDesktops,
DynamicWindowsDesktops: p.dynamicWindowsDesktops,
LinuxDesktops: p.linuxDesktops,
SAMLIdPServiceProviders: p.samlIDPServiceProviders,
UserGroups: p.userGroups,
Okta: p.okta,
Integrations: p.integrations,
UserTasks: p.userTasks,
DiscoveryConfigs: p.discoveryConfigs,
UserLoginStates: p.userLoginStates,
SecReports: p.secReports,
AccessLists: p.accessLists,
KubeWaitingContainers: p.kubeWaitingContainers,
Notifications: p.notifications,
AccessMonitoringRules: p.accessMonitoringRules,
CrownJewels: p.crownJewels,
DatabaseObjects: p.databaseObjects,
SPIFFEFederations: p.spiffeFederations,
StaticHostUsers: p.staticHostUsers,
AutoUpdateService: p.autoUpdateService,
ProvisioningStates: p.provisioningStates,
IdentityCenter: p.identityCenter,
PluginStaticCredentials: p.pluginStaticCredentials,
WorkloadIdentity: p.workloadIdentity,
RecordingEncryption: p.recordingEncryption,
MaxRetryPeriod: 200 * time.Millisecond,
EventsC: p.eventsC,
GitServers: p.gitServers,
HealthCheckConfig: p.healthCheckConfig,
BotInstanceService: p.botInstanceService,
Plugin: p.plugin,
AppAuthConfig: p.appAuthConfigs,
StaticScopedToken: p.clusterConfigS,
WorkloadClusterService: p.workloadClusters,
Summarizer: p.summarizer,
SubCAService: p.subCA,
}))
require.NoError(t, err)
_, err = p.cache.GetCertAuthorities(ctx, types.UserCA, false)
require.True(t, trace.IsConnectionProblem(err))
ca, err := authcatest.NewCA(types.UserCA, "example.com")
require.NoError(t, err)
// NOTE 1: this could produce event processed
// below, based on whether watcher restarts to get the event
// or not, which is normal, but has to be accounted below
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
p.backend.SetReadError(nil)
// wait for watcher to restart
waitForRestart(t, p.eventsC)
normalizeCA := func(ca types.CertAuthority) types.CertAuthority {
ca = ca.Clone()
ca.SetExpiry(time.Time{})
types.RemoveCASecrets(ca)
return ca
}
_ = normalizeCA
out, err := p.cache.GetCertAuthority(ctx, ca.GetID(), false)
require.NoError(t, err)
require.Empty(t, cmp.Diff(normalizeCA(ca), normalizeCA(out), cmpopts.IgnoreFields(types.Metadata{}, "Revision")))
// fail again, make sure last recent data is still served
// on errors
p.backend.SetReadError(trace.ConnectionProblem(nil, "backend is unavailable"))
p.eventsS.closeWatchers()
// wait for the watcher to fail
// there could be optional event processed event,
// see NOTE 1 above
expectNextEvent(t, p.eventsC, WatcherFailed, EventProcessed, Reloading)
// backend is out, but old value is available
out2, err := p.cache.GetCertAuthority(ctx, ca.GetID(), false)
require.NoError(t, err)
require.Empty(t, cmp.Diff(normalizeCA(ca), normalizeCA(out2), cmpopts.IgnoreFields(types.Metadata{}, "Revision")))
// add modification and expect the resource to recover
ca.SetRoleMap(types.RoleMap{types.RoleMapping{Remote: "test", Local: []string{"local-test"}}})
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
// now, recover the backend and make sure the
// service is back and the new value has propagated
p.backend.SetReadError(nil)
// wait for watcher to restart successfully; ignoring any failed
// attempts which occurred before backend became healthy again.
expectEvent(t, p.eventsC, WatcherStarted)
// new value is available now
out, err = p.cache.GetCertAuthority(ctx, ca.GetID(), false)
require.NoError(t, err)
require.Empty(t, cmp.Diff(normalizeCA(ca), normalizeCA(out), cmpopts.IgnoreFields(types.Metadata{}, "Revision")))
}
// TestRecovery tests error recovery scenario
func TestRecovery(t *testing.T) {
t.Parallel()
ctx := context.Background()
p := newPackForAuth(t)
t.Cleanup(p.Close)
ca, err := authcatest.NewCA(types.UserCA, "example.com")
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca))
select {
case event := <-p.eventsC:
require.Equal(t, EventProcessed, event.Type)
case <-time.After(time.Second):
t.Fatalf("timeout waiting for event")
}
// event has arrived, now close the watchers
watchers := p.eventsS.getWatchers()
require.Len(t, watchers, 1)
p.eventsS.closeWatchers()
// wait for watcher to restart
waitForRestart(t, p.eventsC)
// add modification and expect the resource to recover
ca2, err := authcatest.NewCA(types.UserCA, "example2.com")
require.NoError(t, err)
require.NoError(t, p.trustS.UpsertCertAuthority(ctx, ca2))
// wait for watcher to receive an event
select {
case event := <-p.eventsC:
require.Equal(t, EventProcessed, event.Type)
case <-time.After(time.Second):
t.Fatalf("timeout waiting for event")
}
out, err := p.cache.GetCertAuthority(context.Background(), ca2.GetID(), false)
require.NoError(t, err)
types.RemoveCASecrets(ca2)
require.Empty(t, cmp.Diff(ca2, out, cmpopts.IgnoreFields(types.Metadata{}, "Revision")))
}
func mustCreateDatabase(t *testing.T, name, protocol, uri string) *types.DatabaseV3 {
database, err := types.NewDatabaseV3(
types.Metadata{
Name: name,
},
types.DatabaseSpecV3{
Protocol: protocol,
URI: uri,
},
)
require.NoError(t, err)
return database
}
func newUserTasks(t *testing.T) *usertasksv1.UserTask {
t.Helper()
ut, err := usertasks.NewDiscoverEC2UserTask(&usertasksv1.UserTaskSpec{
Integration: "my-integration",
TaskType: usertasks.TaskTypeDiscoverEC2,
IssueType: "ec2-ssm-agent-not-registered",
State: "OPEN",
DiscoverEc2: &usertasksv1.DiscoverEC2{
AccountId: "123456789012",
Region: "us-east-1",
Instances: map[string]*usertasksv1.DiscoverEC2Instance{
"i-123": {
InstanceId: "i-123",
DiscoveryConfig: "dc01",
DiscoveryGroup: "dg01",
SyncTime: timestamppb.Now(),
},
},
},
})
require.NoError(t, err)
return ut
}
type testOptions struct {
skipPaginationTest bool
}
type optionsFunc func(*testOptions)
// TODO(okraport): remove this when all getters support pagination.
func withSkipPaginationTest() optionsFunc {
return func(opts *testOptions) {
opts.skipPaginationTest = true
}
}
// testResources is a wrapper for testing resources conforming to types.Resource
func testResources[T types.Resource](t *testing.T, p *testPack, funcs testFuncs[T], opts ...optionsFunc) {
funcs.resource = defaultResourceOps[T]()
testResourcesInternal(t, p, funcs, opts...)
}
// testResources153 is a wrapper for testing resources conforming to types.Resource153
func testResources153[T types.Resource153](t *testing.T, p *testPack, funcs testFuncs[T], opts ...optionsFunc) {
// TODO(rana): Add broader support for virtual resources in list operations.
// Virtual resources change the total count returned by list operations,
// and is unexpected for the current test. When updated, we can remove virtual
// resource filtering and paging from lib/cache/health_check_config_test.go.
opts = append(opts, withSkipPaginationTest())
funcs.resource = defaultResource153Ops[T]()
testResourcesInternal(t, p, funcs, opts...)
}
// testResourcesInternal is a generic tester for resources.
func testResourcesInternal[T any](t *testing.T, p *testPack, funcs testFuncs[T], opts ...optionsFunc) {
t.Helper()
require.NotNil(t, funcs.resource)
require.NotNil(t, funcs.resource.Name)
if funcs.update != nil {
require.NotNil(t, funcs.resource.Modify)
require.NotNil(t, funcs.resource.Setup)
}
var options testOptions
for _, opt := range opts {
opt(&options)
}
ctx := t.Context()
if !options.skipPaginationTest {
testResourcePagination(t, p, funcs)
}
// Create a resource.
r, err := funcs.newResource("test-resource-1")
require.NoError(t, err)
// update is optional as not every resource implements it
if funcs.update != nil {
funcs.resource.Setup(r)
}
err = funcs.create(ctx, r)
require.NoError(t, err)
cmpOpts := funcs.resource.cmpOpts
assertCacheContents := func(expected []T) {
require.EventuallyWithT(t, func(t *assert.CollectT) {
out, err := funcs.cacheListAll(ctx)
assert.NoError(t, err)
// If the cache is expected to be empty, then test explicitly for
// *that* rather than do an equality test. An equality test here
// would be overly-pedantic about a service returning `nil` rather
// than an empty slice.
if len(expected) == 0 {
require.Empty(t, out)
return
}
require.Empty(t, cmp.Diff(expected, out, cmpOpts...))
}, 2*time.Second, 10*time.Millisecond)
}
// Check that the resource is now in the backend.
out, err := funcs.listAll(ctx)
require.NoError(t, err)
require.Len(t, out, 1)
require.Empty(t, cmp.Diff([]T{r}, out, cmpOpts...))
// Wait until the information has been replicated to the cache.
assertCacheContents([]T{r})
// cacheGet is optional as not every resource implements it
if funcs.cacheGet != nil {
// Make sure a single cache get works.
getR, err := funcs.cacheGet(ctx, funcs.resource.Name(r))
require.NoError(t, err)
require.Empty(t, cmp.Diff(r, getR, cmpOpts...))
// Make sure we get a NotFoundError (and not a panic) when the resource
// is not found.
_, err = funcs.cacheGet(ctx, "no-such-resource")
require.ErrorAs(t, err, new(*trace.NotFoundError))
}
// update is optional as not every resource implements it
if funcs.update != nil {
// Not all create functions will result in resource being updated
// with the latest revision. To avoid any conditional update
// failures caused by mismatched revisions, an updated
// copy of the resource is loaded prior to updating.
if funcs.cacheGet != nil {
var err error
r, err = funcs.cacheGet(ctx, funcs.resource.Name(r))
require.NoError(t, err)
}
// Update the resource and upsert it into the backend again.
funcs.resource.Modify(r)
err = funcs.update(ctx, r)
require.NoError(t, err)
}
// Check that the resource is in the backend and only one exists (so an
// update occurred).
out, err = funcs.listAll(ctx)
require.NoError(t, err)
require.Empty(t, cmp.Diff([]T{r}, out, cmpOpts...))
// Check that information has been replicated to the cache.
assertCacheContents([]T{r})
if funcs.delete != nil {
// Add a second resource.
r2, err := funcs.newResource("test-resource-2")
require.NoError(t, err)
require.NoError(t, funcs.create(ctx, r2))
assertCacheContents([]T{r, r2})
// Check that only one resource is deleted.
require.NoError(t, funcs.delete(ctx, funcs.resource.Name(r2)))
assertCacheContents([]T{r})
}
// Remove all resources from the backend.
err = funcs.deleteAll(ctx)
require.NoError(t, err)
// Check that information has been replicated to the cache.
assertCacheContents([]T{})
}
func TestRelativeExpiry(t *testing.T) {
t.Parallel()
const checkInterval = time.Second
const nodeCount = int64(100)
// make sure the event buffer is much larger than node count
// so that we can batch create nodes without waiting on each event
require.Less(t, int(nodeCount*3), eventBufferSize)
ctx := context.Background()
clock := clockwork.NewFakeClockAt(time.Now().Add(time.Hour))
p := newTestPack(t, func(c Config) Config {
c.RelativeExpiryCheckInterval = checkInterval
c.Clock = clock
return ForAuth(c)
})
t.Cleanup(p.Close)
// add servers that expire at a range of times
now := clock.Now()
for i := range nodeCount {
exp := now.Add(time.Minute * time.Duration(i))
server := NewServer(types.KindNode, uuid.New().String(), "127.0.0.1:2022", apidefaults.Namespace)
server.SetExpiry(exp)
_, err := p.presenceS.UpsertNode(ctx, server)
require.NoError(t, err)
}
// wait for nodes to reach cache (we batch insert first for performance reasons)
for range nodeCount {
expectEvent(t, p.eventsC, EventProcessed)
}
nodes, err := p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.Len(t, nodes, 100)
clock.Advance(time.Minute * 25)
// get rid of events that were emitted before clock advanced
drainEvents(p.eventsC)
// wait for next relative expiry check to run
expectEvent(t, p.eventsC, RelativeExpiry)
// verify that roughly expected proportion of nodes was removed.
nodes, err = p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.True(t, len(nodes) < 100 && len(nodes) > 75, "node_count=%d", len(nodes))
clock.Advance(time.Minute * 25)
// get rid of events that were emitted before clock advanced
drainEvents(p.eventsC)
// wait for next relative expiry check to run
expectEvent(t, p.eventsC, RelativeExpiry)
// verify that roughly expected proportion of nodes was removed.
nodes, err = p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.True(t, len(nodes) < 75 && len(nodes) > 50, "node_count=%d", len(nodes))
// finally, we check the "sliding window" by verifying that we don't remove all nodes
// even if we advance well past the latest expiry time.
clock.Advance(time.Hour * 24)
// get rid of events that were emitted before clock advanced
drainEvents(p.eventsC)
// wait for next relative expiry check to run
expectEvent(t, p.eventsC, RelativeExpiry)
// verify that sliding window has preserved most recent nodes
nodes, err = p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.NotEmpty(t, nodes, "node_count=%d", len(nodes))
}
func TestRelativeExpiryLimit(t *testing.T) {
t.Parallel()
const (
checkInterval = time.Second
nodeCount = 100
expiryLimit = 10
)
// make sure the event buffer is much larger than node count
// so that we can batch create nodes without waiting on each event
require.Less(t, int(nodeCount*3), eventBufferSize)
ctx := context.Background()
clock := clockwork.NewFakeClockAt(time.Now().Add(time.Hour))
p := newTestPack(t, func(c Config) Config {
c.RelativeExpiryCheckInterval = checkInterval
c.RelativeExpiryLimit = expiryLimit
c.Clock = clock
return ForAuth(c)
})
t.Cleanup(p.Close)
// add servers that expire at a range of times
now := clock.Now()
for i := range nodeCount {
exp := now.Add(time.Minute * time.Duration(i))
server := NewServer(types.KindNode, uuid.New().String(), "127.0.0.1:2022", apidefaults.Namespace)
server.SetExpiry(exp)
_, err := p.presenceS.UpsertNode(ctx, server)
require.NoError(t, err)
}
// wait for nodes to reach cache (we batch insert first for performance reasons)
for range nodeCount {
expectEvent(t, p.eventsC, EventProcessed)
}
nodes, err := p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.Len(t, nodes, nodeCount)
clock.Advance(time.Hour * 24)
for expired := nodeCount - expiryLimit; expired > expiryLimit; expired -= expiryLimit {
// get rid of events that were emitted before clock advanced
drainEvents(p.eventsC)
// wait for next relative expiry check to run
expectEvent(t, p.eventsC, RelativeExpiry)
// verify that the limit is respected.
nodes, err = p.cache.GetNodes(ctx, apidefaults.Namespace)
require.NoError(t, err)
require.Len(t, nodes, expired)
// advance clock to trigger next relative expiry check
clock.Advance(time.Hour * 24)
}
}
func TestRelativeExpiryOnlyForAuth(t *testing.T) {
t.Parallel()
clock := clockwork.NewFakeClockAt(time.Now().Add(time.Hour))
p := newTestPack(t, func(c Config) Config {
c.RelativeExpiryCheckInterval = time.Second
c.Clock = clock
c.Watches = []types.WatchKind{{Kind: types.KindNode}}
return c
})
t.Cleanup(p.Close)
p2 := newTestPack(t, func(c Config) Config {
c.RelativeExpiryCheckInterval = time.Second
c.Clock = clock
c.target = "llama"
c.Watches = []types.WatchKind{
{Kind: types.KindNode},
{Kind: types.KindCertAuthority},
}
return c
})
t.Cleanup(p2.Close)
for range 2 {
clock.Advance(time.Hour * 24)
drainEvents(p.eventsC)
unexpectedEvent(t, p.eventsC, RelativeExpiry)
clock.Advance(time.Hour * 24)
drainEvents(p2.eventsC)
unexpectedEvent(t, p2.eventsC, RelativeExpiry)
}
}
func TestCache_Backoff(t *testing.T) {
t.Parallel()
clock := clockwork.NewFakeClock()
p := newTestPack(t, func(c Config) Config {
c.MaxRetryPeriod = defaults.MaxWatcherBackoff
c.Clock = clock
return ForNode(c)
})
t.Cleanup(p.Close)
// close watchers to trigger a reload event
watchers := p.eventsS.getWatchers()
require.Len(t, watchers, 1)
p.eventsS.closeWatchers()
p.backend.SetReadError(trace.ConnectionProblem(nil, "backend is unavailable"))
step := p.cache.Config.MaxRetryPeriod / 16.0
for i := range 5 {
// wait for cache to reload
select {
case event := <-p.eventsC:
require.Equal(t, Reloading, event.Type)
duration, err := time.ParseDuration(event.Event.Resource.GetKind())
require.NoError(t, err)
// emulate the logic of exponential backoff multiplier calc
var mul int64
if i == 0 {
mul = 0
} else {
mul = 1 << (i - 1)
}
stepMin := step * time.Duration(mul) / 2
stepMax := step * time.Duration(mul+1)
require.GreaterOrEqual(t, duration, stepMin)
require.LessOrEqual(t, duration, stepMax)
// wait for cache to get to retry.After
clock.BlockUntil(1)
// add some extra to the duration to ensure the retry occurs
clock.Advance(p.cache.MaxRetryPeriod)
case <-time.After(time.Minute):
t.Fatalf("timeout waiting for event")
}
// wait for cache to fail again - backend will still produce a ConnectionProblem error
select {
case event := <-p.eventsC:
require.Equal(t, WatcherFailed, event.Type)
case <-time.After(30 * time.Second):
t.Fatalf("timeout waiting for event")
}
}
}
// TestSetupConfigFns ensures that all WatchKinds used in setup config functions are present in ForAuth() as well.
func TestSetupConfigFns(t *testing.T) {
t.Parallel()
ctx := context.Background()
bk, err := memory.New(memory.Config{
Context: ctx,
Mirror: true,
})
require.NoError(t, err)
defer bk.Close()
clusterConfigCache, err := local.NewClusterConfigurationService(bk)
require.NoError(t, err)
clusterName, err := services.NewClusterNameWithRandomID(types.ClusterNameSpecV2{
ClusterName: "example.com",
})
require.NoError(t, err)
err = clusterConfigCache.UpsertClusterName(clusterName)
require.NoError(t, err)
setupFuncs := map[string]SetupConfigFn{
"ForProxy": ForProxy,
"ForRelay": ForRelay,
"ForRemoteProxy": ForRemoteProxy,
"ForNode": ForNode,
"ForKubernetes": ForKubernetes,
"ForApps": ForApps,
"ForDatabases": ForDatabases,
"ForWindowsDesktop": ForWindowsDesktop,
"ForDiscovery": ForDiscovery,
"ForOkta": ForOkta,
}
authKindMap := make(map[resourceKind]types.WatchKind)
for _, wk := range ForAuth(Config{ClusterConfig: clusterConfigCache}).Watches {
authKindMap[resourceKind{kind: wk.Kind, subkind: wk.SubKind}] = wk
}
for name, f := range setupFuncs {
t.Run(name, func(t *testing.T) {
for _, wk := range f(Config{ClusterConfig: clusterConfigCache}).Watches {
authWK, ok := authKindMap[resourceKind{kind: wk.Kind, subkind: wk.SubKind}]
if !ok || !authWK.Contains(wk) {
t.Errorf("%s includes WatchKind %s that is missing from ForAuth", name, wk.String())
}
if wk.Kind == types.KindCertAuthority {
require.NotEmpty(t, wk.Filter, "every setup fn except auth should have a CA filter")
}
}
})
}
authCAWatchKind, ok := authKindMap[resourceKind{kind: types.KindCertAuthority}]
require.True(t, ok)
require.Empty(t, authCAWatchKind.Filter, "auth should not use a CA filter")
}
type proxyEvents struct {
sync.Mutex
watchers []types.Watcher
events types.Events
ignoreKinds map[resourceKind]struct{}
}
func (p *proxyEvents) getWatchers() []types.Watcher {
p.Lock()
defer p.Unlock()
out := make([]types.Watcher, len(p.watchers))
copy(out, p.watchers)
return out
}
func (p *proxyEvents) closeWatchers() {
p.Lock()
defer p.Unlock()
for i := range p.watchers {
p.watchers[i].Close()
}
p.watchers = nil
}
func (p *proxyEvents) NewWatcher(ctx context.Context, watch types.Watch) (types.Watcher, error) {
var effectiveKinds []types.WatchKind
for _, requested := range watch.Kinds {
if _, ok := p.ignoreKinds[resourceKind{kind: requested.Kind, subkind: requested.SubKind}]; ok {
continue
}
effectiveKinds = append(effectiveKinds, requested)
}
if len(effectiveKinds) == 0 {
return nil, trace.BadParameter("all of the requested kinds were ignored")
}
watch.Kinds = effectiveKinds
w, err := p.events.NewWatcher(ctx, watch)
if err != nil {
return nil, trace.Wrap(err)
}
p.Lock()
defer p.Unlock()
p.watchers = append(p.watchers, w)
return w, nil
}
func newProxyEvents(events types.Events, ignoreKinds []types.WatchKind) *proxyEvents {
ignoreSet := make(map[resourceKind]struct{}, len(ignoreKinds))
for _, kind := range ignoreKinds {
ignoreSet[resourceKind{kind: kind.Kind, subkind: kind.SubKind}] = struct{}{}
}
return &proxyEvents{
events: events,
ignoreKinds: ignoreSet,
}
}
// TestCacheWatchKindExistsInEvents ensures that the watch kinds for each cache component are in sync
// with proto Events delivered via WatchEvents. If a watch kind is added to a cache component, but it
// doesn't exist in the proto Events definition, an error will cause the WatchEvents stream to be closed.
// This causes the cache to reinitialize every time that an unknown message is received and can lead to
// a permanently unhealthy cache.
//
// While this test will ensure that there are no issues for the current release, it does not guarantee
// that this issue won't arise across releases.
func TestCacheWatchKindExistsInEvents(t *testing.T) {
t.Parallel()
clock := clockwork.NewFakeClockAt(time.Now())
cases := map[string]Config{
"ForAuth": ForAuth(Config{}),
"ForProxy": ForProxy(Config{}),
"ForRelay": ForRelay(Config{}),
"ForRemoteProxy": ForRemoteProxy(Config{}),
"ForNode": ForNode(Config{}),
"ForKubernetes": ForKubernetes(Config{}),
"ForApps": ForApps(Config{}),
"ForDatabases": ForDatabases(Config{}),
"ForOkta": ForOkta(Config{}),
}
events := map[string]types.Resource{
types.KindCertAuthority: &types.CertAuthorityV2{},
types.KindClusterName: &types.ClusterNameV2{},
types.KindClusterAuditConfig: types.DefaultClusterAuditConfig(),
types.KindClusterNetworkingConfig: types.DefaultClusterNetworkingConfig(),
types.KindClusterAuthPreference: types.DefaultAuthPreference(),
types.KindSessionRecordingConfig: types.DefaultSessionRecordingConfig(),
types.KindUIConfig: &types.UIConfigV1{},
types.KindStaticTokens: &types.StaticTokensV2{},
types.KindStaticScopedTokens: &types.StaticTokensV2{},
types.KindToken: &types.ProvisionTokenV2{},
types.KindUser: &types.UserV2{},
types.KindRole: &types.RoleV6{Version: types.V4},
types.KindNamespace: &types.Namespace{},
types.KindNode: &types.ServerV2{},
types.KindProxy: &types.ServerV2{},
types.KindAuthServer: &types.ServerV2{},
types.KindReverseTunnel: &types.ReverseTunnelV2{},
types.KindTunnelConnection: &types.TunnelConnectionV2{},
types.KindAccessRequest: &types.AccessRequestV3{},
types.KindAppServer: &types.AppServerV3{},
types.KindApp: &types.AppV3{},
types.KindWebSession: &types.WebSessionV2{SubKind: types.KindWebSession},
types.KindAppSession: &types.WebSessionV2{SubKind: types.KindAppSession},
types.KindSnowflakeSession: &types.WebSessionV2{SubKind: types.KindSnowflakeSession},
types.KindWebToken: &types.WebTokenV3{},
types.KindRemoteCluster: &types.RemoteClusterV3{},
types.KindKubeServer: &types.KubernetesServerV3{},
types.KindDatabaseService: &types.DatabaseServiceV1{},
types.KindDatabaseServer: &types.DatabaseServerV3{},
types.KindDatabase: &types.DatabaseV3{},
types.KindNetworkRestrictions: &types.NetworkRestrictionsV4{},
types.KindLock: &types.LockV2{},
types.KindWindowsDesktopService: &types.WindowsDesktopServiceV3{},
types.KindWindowsDesktop: &types.WindowsDesktopV3{},
types.KindDynamicWindowsDesktop: &types.DynamicWindowsDesktopV1{},
types.KindLinuxDesktop: types.ProtoResource153ToLegacy(newLinuxDesktop("linux-desktop")),
types.KindInstaller: &types.InstallerV1{},
types.KindKubernetesCluster: &types.KubernetesClusterV3{},
types.KindSAMLIdPServiceProvider: &types.SAMLIdPServiceProviderV1{},
types.KindUserGroup: &types.UserGroupV1{},
types.KindOktaImportRule: &types.OktaImportRuleV1{},
types.KindOktaAssignment: &types.OktaAssignmentV1{},
types.KindIntegration: &types.IntegrationV1{},
types.KindDiscoveryConfig: newDiscoveryConfig(t, "discovery-config"),
types.KindHeadlessAuthentication: &types.HeadlessAuthentication{},
types.KindUserLoginState: newUserLoginState(t, "user-login-state"),
types.KindAuditQuery: newAuditQuery(t, "audit-query"),
types.KindSecurityReport: newSecurityReport(t, "security-report"),
types.KindSecurityReportState: newSecurityReport(t, "security-report-state"),
types.KindAccessList: newAccessList(t, "access-list", clock),
types.KindAccessListMember: newAccessListMember(t, "access-list", "member"),
types.KindAccessListReview: newAccessListReview(t, "access-list", "review"),
types.KindKubeWaitingContainer: newKubeWaitingContainer(t),
types.KindNotification: types.Resource153ToLegacy(newUserNotification(t, "test")),
types.KindGlobalNotification: types.Resource153ToLegacy(newGlobalNotification(t, "test")),
types.KindAccessMonitoringRule: types.Resource153ToLegacy(newAccessMonitoringRule(t, "test")),
types.KindCrownJewel: types.Resource153ToLegacy(newCrownJewel(t, "test")),
types.KindDatabaseObject: types.Resource153ToLegacy(newDatabaseObject(t, "test")),
types.KindBeam: types.Resource153ToLegacy(newBeamResource("some-beam", "curious-harbor", clock.Now().Add(time.Hour))),
types.KindAccessGraphSettings: types.Resource153ToLegacy(newAccessGraphSettings(t)),
types.KindSPIFFEFederation: types.Resource153ToLegacy(newSPIFFEFederation("test")),
types.KindStaticHostUser: types.Resource153ToLegacy(newStaticHostUser(t, "test")),
types.KindAutoUpdateConfig: types.Resource153ToLegacy(newAutoUpdateConfig(t)),
types.KindAutoUpdateVersion: types.Resource153ToLegacy(newAutoUpdateVersion(t)),
types.KindAutoUpdateAgentRollout: types.Resource153ToLegacy(newAutoUpdateAgentRollout(t)),
types.KindAutoUpdateAgentReport: types.Resource153ToLegacy(newAutoUpdateAgentReport(t, "test")),
types.KindAutoUpdateBotInstanceReport: types.Resource153ToLegacy(newAutoUpdateBotInstanceReport(t)),
types.KindUserTask: types.Resource153ToLegacy(newUserTasks(t)),
types.KindProvisioningPrincipalState: types.Resource153ToLegacy(newProvisioningPrincipalState("u-alice@example.com")),
types.KindIdentityCenterAccount: types.Resource153ToLegacy(newIdentityCenterAccount("some_account")),
types.KindIdentityCenterAccountAssignment: types.Resource153ToLegacy(newIdentityCenterAccountAssignment("some_account_assignment")),
types.KindIdentityCenterPrincipalAssignment: types.Resource153ToLegacy(newIdentityCenterPrincipalAssignment("some_principal_assignment")),
types.KindPlugin: &types.PluginV1{},
types.KindPluginStaticCredentials: &types.PluginStaticCredentialsV1{},
types.KindGitServer: &types.ServerV2{},
types.KindWorkloadIdentity: types.Resource153ToLegacy(newWorkloadIdentity("some_identifier")),
types.KindRecordingEncryption: types.Resource153ToLegacy(newRecordingEncryption()),
types.KindHealthCheckConfig: types.Resource153ToLegacy(newHealthCheckConfig(t, "some-name")),
scopedaccess.KindScopedRole: types.Resource153ToLegacy(&scopedaccessv1.ScopedRole{}),
scopedaccess.KindScopedRoleAssignment: types.Resource153ToLegacy(&scopedaccessv1.ScopedRoleAssignment{}),
types.KindRelayServer: types.ProtoResource153ToLegacy(new(presencev1.RelayServer)),
types.KindBotInstance: types.ProtoResource153ToLegacy(new(machineidv1.BotInstance)),
types.KindAppAuthConfig: types.Resource153ToLegacy(new(appauthconfigv1.AppAuthConfig)),
types.KindWorkloadCluster: types.Resource153ToLegacy(newWorkloadCluster(t, "test")),
types.KindInferenceModel: types.Resource153ToLegacy(new(summaryv1.InferenceModel)),
types.KindInferenceSecret: types.Resource153ToLegacy(new(summaryv1.InferenceSecret)),
types.KindInferencePolicy: types.Resource153ToLegacy(new(summaryv1.InferencePolicy)),
types.KindRetrievalModel: types.Resource153ToLegacy(new(summaryv1.RetrievalModel)),
types.KindCertAuthorityOverride: types.Resource153ToLegacy(&subcav1.CertAuthorityOverride{}),
types.KindValidatedMFAChallenge: &types.ResourceHeader{Kind: types.KindValidatedMFAChallenge},
}
for name, cfg := range cases {
t.Run(name, func(t *testing.T) {
for _, watch := range cfg.Watches {
resource, ok := events[watch.Kind]
require.True(t, ok, "missing event for kind %q", watch.Kind)
protoEvent, err := client.EventToGRPC(types.Event{
Type: types.OpPut,
Resource: resource,
})
require.NoError(t, err)
event, err := client.EventFromGRPC(protoEvent)
require.NoError(t, err)
// unwrap the RFD 153 resource if necessary
switch uw := event.Resource.(type) {
case types.Resource153UnwrapperT[*workloadidentityv1.WorkloadIdentity]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*workloadidentityv1.WorkloadIdentity]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*identitycenterv1.PrincipalAssignment]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*identitycenterv1.PrincipalAssignment]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*identitycenterv1.AccountAssignment]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*identitycenterv1.AccountAssignment]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*identitycenterv1.Account]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*identitycenterv1.Account]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*provisioningv1.PrincipalState]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*provisioningv1.PrincipalState]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*usertasksv1.UserTask]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*usertasksv1.UserTask]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*autoupdate.AutoUpdateAgentReport]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*autoupdate.AutoUpdateAgentReport]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*autoupdate.AutoUpdateAgentRollout]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*autoupdate.AutoUpdateAgentRollout]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*autoupdate.AutoUpdateVersion]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*autoupdate.AutoUpdateVersion]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*autoupdate.AutoUpdateConfig]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*autoupdate.AutoUpdateConfig]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*autoupdate.AutoUpdateBotInstanceReport]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*autoupdate.AutoUpdateBotInstanceReport]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*recordingencryptionv1.RecordingEncryption]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*recordingencryptionv1.RecordingEncryption]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*userprovisioningpb.StaticHostUser]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*userprovisioningpb.StaticHostUser]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*machineidv1.SPIFFEFederation]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*machineidv1.SPIFFEFederation]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*clusterconfigpb.AccessGraphSettings]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*clusterconfigpb.AccessGraphSettings]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*dbobjectv1.DatabaseObject]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*dbobjectv1.DatabaseObject]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*crownjewelv1.CrownJewel]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*crownjewelv1.CrownJewel]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*accessmonitoringrulesv1.AccessMonitoringRule]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*accessmonitoringrulesv1.AccessMonitoringRule]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*notificationsv1.GlobalNotification]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*notificationsv1.GlobalNotification]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*notificationsv1.Notification]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*notificationsv1.Notification]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*kubewaitingcontainerpb.KubernetesWaitingContainer]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*kubewaitingcontainerpb.KubernetesWaitingContainer]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*healthcheckconfigv1.HealthCheckConfig]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*healthcheckconfigv1.HealthCheckConfig]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*scopedaccessv1.ScopedRole]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*scopedaccessv1.ScopedRole]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*scopedaccessv1.ScopedRoleAssignment]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*scopedaccessv1.ScopedRoleAssignment]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*presencev1.RelayServer]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*presencev1.RelayServer]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*machineidv1.BotInstance]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*machineidv1.BotInstance]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*appauthconfigv1.AppAuthConfig]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*appauthconfigv1.AppAuthConfig]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*workloadclusterv1.WorkloadCluster]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*workloadclusterv1.WorkloadCluster]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*summaryv1.InferenceModel]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*summaryv1.InferenceModel]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*summaryv1.InferenceSecret]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*summaryv1.InferenceSecret]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*summaryv1.InferencePolicy]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*summaryv1.InferencePolicy]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*summaryv1.RetrievalModel]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*summaryv1.RetrievalModel]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*subcav1.CertAuthorityOverride]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*subcav1.CertAuthorityOverride]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
case types.Resource153UnwrapperT[*beamsv1.Beam]:
require.Empty(t, cmp.Diff(resource.(types.Resource153UnwrapperT[*beamsv1.Beam]).UnwrapT(), uw.UnwrapT(), protocmp.Transform()))
default:
require.Empty(t, cmp.Diff(resource, event.Resource))
}
}
})
}
}
// TestPartialHealth ensures that when an event source confirms only some resource kinds specified on the watch request,
// Cache operates in partially healthy mode in which it serves reads of the confirmed kinds from the cache and
// lets everything else pass through.
func TestPartialHealth(t *testing.T) {
t.Parallel()
ctx := context.Background()
// setup cache such that role resources wouldn't be recognized by the event source and wouldn't be cached.
p, err := newPack(t, ForApps, ignoreKinds([]types.WatchKind{{Kind: types.KindRole}}))
require.NoError(t, err)
t.Cleanup(p.Close)
role, err := types.NewRole("editor", types.RoleSpecV6{})
require.NoError(t, err)
_, err = p.accessS.UpsertRole(ctx, role)
require.NoError(t, err)
user, err := types.NewUser("bob")
require.NoError(t, err)
user, err = p.usersS.UpsertUser(ctx, user)
require.NoError(t, err)
select {
case event := <-p.eventsC:
require.Equal(t, EventProcessed, event.Type)
require.Equal(t, types.KindUser, event.Event.Resource.GetKind())
case <-time.After(time.Second):
t.Fatal("timeout waiting for event")
}
// make sure that the user resource works as normal and gets replicated to cache
replicatedUsers, err := p.cache.GetUsers(ctx, false)
require.NoError(t, err)
require.Len(t, replicatedUsers, 1)
// now add a label to the user directly in the cache
meta := user.GetMetadata()
meta.Labels = map[string]string{"origin": "cache"}
user.SetMetadata(meta)
err = p.cache.collections.users.onPut(user)
require.NoError(t, err)
// the label on the returned user proves that it came from the cache
resultUser, err := p.cache.GetUser(ctx, "bob", false)
require.NoError(t, err)
require.Equal(t, "cache", resultUser.GetMetadata().Labels["origin"])
// query cache storage directly to ensure roles haven't been replicated
require.Empty(t, p.cache.collections.roles.store.len())
// non-empty result here proves that it was not served from cache
resultRoles, err := p.cache.GetRoles(ctx)
require.NoError(t, err)
require.Len(t, resultRoles, 1)
// ensure that cache cannot be watched for resources that weren't confirmed in regular mode
testWatch := types.Watch{
Kinds: []types.WatchKind{
{Kind: types.KindUser},
{Kind: types.KindRole},
},
}
_, err = p.cache.NewWatcher(ctx, testWatch)
require.Error(t, err)
// same request should work in partial success mode, but WatchStatus on the OpInit event should indicate
// that only user resources will be watched.
testWatch.AllowPartialSuccess = true
w, err := p.cache.NewWatcher(ctx, testWatch)
require.NoError(t, err)
select {
case e := <-w.Events():
require.Equal(t, types.OpInit, e.Type)
watchStatus, ok := e.Resource.(types.WatchStatus)
require.True(t, ok)
require.Equal(t, []types.WatchKind{{Kind: types.KindUser}}, watchStatus.GetKinds())
case <-time.After(time.Second):
t.Fatal("Timeout waiting for event.")
}
}
// TestInvalidDatbases given a database that returns an error on validation, the
// cache should not be impacted, and the database must be on it. This scenario
// is most common on Teleport upgrades/downgrades where the database validation
// can have new rules, causing the existing database to fail on validation.
func TestInvalidDatabases(t *testing.T) {
t.Parallel()
ctx := context.Background()
generateInvalidDB := func(t *testing.T, name string) types.Database {
db := &types.DatabaseV3{Metadata: types.Metadata{
Name: name,
}, Spec: types.DatabaseSpecV3{
Protocol: "invalid-protocol",
URI: "non-empty-uri",
}}
// Just ensures the database we're using on this test will trigger a
// validation failure.
require.Error(t, services.ValidateDatabase(db))
return db
}
for name, tc := range map[string]struct {
storeFunc func(*testing.T, *backend.Wrapper, *Cache)
}{
"CreateDatabase": {
storeFunc: func(t *testing.T, b *backend.Wrapper, _ *Cache) {
db := generateInvalidDB(t, "invalid-db")
value, err := services.MarshalDatabase(db)
require.NoError(t, err)
_, err = b.Create(ctx, backend.Item{
Key: backend.NewKey("db", db.GetName()),
Value: value,
Expires: db.Expiry(),
})
require.NoError(t, err)
},
},
"UpdateDatabase": {
storeFunc: func(t *testing.T, b *backend.Wrapper, c *Cache) {
dbName := "updated-db"
validDB, err := types.NewDatabaseV3(types.Metadata{
Name: dbName,
}, types.DatabaseSpecV3{
Protocol: defaults.ProtocolPostgres,
URI: "postgres://localhost",
})
require.NoError(t, err)
require.NoError(t, services.ValidateDatabase(validDB))
marshalledDB, err := services.MarshalDatabase(validDB)
require.NoError(t, err)
_, err = b.Create(ctx, backend.Item{
Key: backend.NewKey("db", validDB.GetName()),
Value: marshalledDB,
Expires: validDB.Expiry(),
})
require.NoError(t, err)
// Wait until the database appear on cache.
require.EventuallyWithT(t, func(t *assert.CollectT) {
dbs, err := c.GetDatabases(ctx)
require.NoError(t, err)
require.Len(t, dbs, 1)
}, time.Second, 100*time.Millisecond, "expected database to be on cache, but nothing found")
cacheDB, err := c.GetDatabase(ctx, dbName)
require.NoError(t, err)
invalidDB := generateInvalidDB(t, cacheDB.GetName())
value, err := services.MarshalDatabase(invalidDB)
require.NoError(t, err)
_, err = b.Update(ctx, backend.Item{
Key: backend.NewKey("db", cacheDB.GetName()),
Value: value,
Expires: invalidDB.Expiry(),
})
require.NoError(t, err)
},
},
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
p := newTestPack(t, ForAuth)
t.Cleanup(p.Close)
tc.storeFunc(t, p.backend, p.cache)
// Events processing should not restart due to an invalid database error.
unexpectedEvent(t, p.eventsC, Reloading)
})
}
}
func newDiscoveryConfig(t *testing.T, name string) *discoveryconfig.DiscoveryConfig {
t.Helper()
discoveryConfig, err := discoveryconfig.NewDiscoveryConfig(
header.Metadata{
Name: name,
},
discoveryconfig.Spec{
DiscoveryGroup: "mygroup",
AWS: []types.AWSMatcher{},
Azure: []types.AzureMatcher{},
GCP: []types.GCPMatcher{},
Kube: []types.KubernetesMatcher{},
},
)
require.NoError(t, err)
discoveryConfig.Status.State = "DISCOVERY_CONFIG_STATE_UNSPECIFIED"
return discoveryConfig
}
func newAuditQuery(t *testing.T, name string) *secreports.AuditQuery {
t.Helper()
item, err := secreports.NewAuditQuery(
header.Metadata{
Name: name,
},
secreports.AuditQuerySpec{
Name: name,
Title: "title",
Description: "desc",
Query: "query",
},
)
require.NoError(t, err)
return item
}
func newSecurityReport(t *testing.T, name string) *secreports.Report {
t.Helper()
item, err := secreports.NewReport(
header.Metadata{
Name: name,
},
secreports.ReportSpec{
Name: name,
Title: "title",
AuditQueries: []*secreports.AuditQuerySpec{
{
Name: "name",
Title: "title",
Description: "desc",
Query: "query",
},
},
Version: "0.0.0",
},
)
require.NoError(t, err)
return item
}
func newSecurityReportState(t *testing.T, name string) *secreports.ReportState {
t.Helper()
item, err := secreports.NewReportState(
header.Metadata{
Name: name,
},
secreports.ReportStateSpec{
Status: "RUNNING",
UpdatedAt: time.Now().UTC(),
},
)
require.NoError(t, err)
return item
}
func newUserLoginState(t *testing.T, name string) *userloginstate.UserLoginState {
t.Helper()
uls, err := userloginstate.New(
header.Metadata{
Name: name,
},
userloginstate.Spec{
Roles: []string{"role1", "role2"},
OriginalTraits: trait.Traits{},
Traits: trait.Traits{
"key1": []string{"value1"},
"key2": []string{"value2"},
},
},
)
require.NoError(t, err)
return uls
}
func newAccessList(t *testing.T, name string, clock clockwork.Clock) *accesslist.AccessList {
t.Helper()
accessList, err := accesslist.NewAccessList(
header.Metadata{
Name: name,
},
accesslist.Spec{
Title: "Title" + name,
Description: "test access list",
Owners: []accesslist.Owner{
{
Name: "test-user1",
Description: "test user 1",
},
{
Name: "test-user2",
Description: "test user 2",
},
},
Audit: accesslist.Audit{
NextAuditDate: clock.Now(),
},
MembershipRequires: accesslist.Requires{
Roles: []string{"mrole1", "mrole2"},
Traits: map[string][]string{
"mtrait1": {"mvalue1", "mvalue2"},
"mtrait2": {"mvalue3", "mvalue4"},
},
},
OwnershipRequires: accesslist.Requires{
Roles: []string{"orole1", "orole2"},
Traits: map[string][]string{
"otrait1": {"ovalue1", "ovalue2"},
"otrait2": {"ovalue3", "ovalue4"},
},
},
Grants: accesslist.Grants{
Roles: []string{"grole1", "grole2"},
Traits: map[string][]string{
"gtrait1": {"gvalue1", "gvalue2"},
"gtrait2": {"gvalue3", "gvalue4"},
},
},
},
)
require.NoError(t, err)
return accessList
}
func newAccessListMember(t *testing.T, accessList, name string) *accesslist.AccessListMember {
t.Helper()
member, err := accesslist.NewAccessListMember(
header.Metadata{
Name: name,
},
accesslist.AccessListMemberSpec{
AccessList: accessList,
Name: name,
Joined: time.Now(),
Expires: time.Now().Add(time.Hour * 24),
Reason: "a reason",
AddedBy: "dummy",
},
)
require.NoError(t, err)
return member
}
func newAccessListReview(t *testing.T, accessList, name string) *accesslist.Review {
t.Helper()
review, err := accesslist.NewReview(
header.Metadata{
Name: name,
},
accesslist.ReviewSpec{
AccessList: accessList,
Reviewers: []string{
"user1",
"user2",
},
ReviewDate: time.Date(2023, 1, 1, 0, 0, 0, 0, time.UTC),
Notes: "Some notes",
Changes: accesslist.ReviewChanges{
MembershipRequirementsChanged: &accesslist.Requires{
Roles: []string{
"role1",
"role2",
},
Traits: trait.Traits{
"trait1": []string{
"value1",
"value2",
},
"trait2": []string{
"value1",
"value2",
},
},
},
RemovedMembers: []string{
"member1",
"member2",
},
ReviewFrequencyChanged: accesslist.ThreeMonths,
ReviewDayOfMonthChanged: accesslist.FifteenthDayOfMonth,
},
},
)
require.NoError(t, err)
return review
}
func newKubeWaitingContainer(t *testing.T) types.Resource {
t.Helper()
waitingCont, err := kubewaitingcontainer.NewKubeWaitingContainer("container", &kubewaitingcontainerpb.KubernetesWaitingContainerSpec{
Username: "user",
Cluster: "cluster",
Namespace: "namespace",
PodName: "pod",
ContainerName: "container",
Patch: []byte("patch"),
PatchType: "application/json-patch+json",
})
require.NoError(t, err)
return types.Resource153ToLegacy(waitingCont)
}
func newCrownJewel(t *testing.T, name string) *crownjewelv1.CrownJewel {
t.Helper()
crownJewel := &crownjewelv1.CrownJewel{
Metadata: &headerv1.Metadata{
Name: name,
},
}
return crownJewel
}
func newDatabaseObject(t *testing.T, name string) *dbobjectv1.DatabaseObject {
t.Helper()
r, err := databaseobject.NewDatabaseObject(name, &dbobjectv1.DatabaseObjectSpec{
Name: name,
Protocol: "postgres",
DatabaseServiceName: "pg",
ObjectKind: "table",
})
require.NoError(t, err)
return r
}
func newAccessGraphSettings(t *testing.T) *clusterconfigpb.AccessGraphSettings {
t.Helper()
r, err := clusterconfig.NewAccessGraphSettings(&clusterconfigpb.AccessGraphSettingsSpec{
SecretsScanConfig: clusterconfigpb.AccessGraphSecretsScanConfig_ACCESS_GRAPH_SECRETS_SCAN_CONFIG_ENABLED,
})
require.NoError(t, err)
return r
}
func newLinuxDesktop(name string) *linuxdesktopv1.LinuxDesktop {
return &linuxdesktopv1.LinuxDesktop{
Kind: types.KindLinuxDesktop,
Version: types.V1,
Metadata: &headerv1.Metadata{
Name: name,
},
Spec: &linuxdesktopv1.LinuxDesktopSpec{
Addr: "127.0.0.1:22",
Hostname: "host",
},
}
}
func newUserNotification(t *testing.T, name string) *notificationsv1.Notification {
t.Helper()
notification := &notificationsv1.Notification{
SubKind: "test-subkind",
Spec: &notificationsv1.NotificationSpec{
Username: name,
},
Metadata: &headerv1.Metadata{
Labels: map[string]string{types.NotificationTitleLabel: "test-title"},
},
}
return notification
}
func newGlobalNotification(t *testing.T, title string) *notificationsv1.GlobalNotification {
t.Helper()
notification := &notificationsv1.GlobalNotification{
Spec: &notificationsv1.GlobalNotificationSpec{
Matcher: &notificationsv1.GlobalNotificationSpec_All{
All: true,
},
Notification: &notificationsv1.Notification{
SubKind: "test-subkind",
Spec: &notificationsv1.NotificationSpec{},
Metadata: &headerv1.Metadata{
Labels: map[string]string{types.NotificationTitleLabel: title},
},
},
},
}
return notification
}
func newAccessMonitoringRule(t *testing.T, name string) *accessmonitoringrulesv1.AccessMonitoringRule {
t.Helper()
notification := &accessmonitoringrulesv1.AccessMonitoringRule{
Kind: types.KindAccessMonitoringRule,
Version: types.V1,
Metadata: &headerv1.Metadata{
Name: name,
},
Spec: &accessmonitoringrulesv1.AccessMonitoringRuleSpec{
Notification: &accessmonitoringrulesv1.Notification{
Name: "test",
},
Subjects: []string{"llama", "shark"},
Condition: "test",
},
}
return notification
}
func newStaticHostUser(t *testing.T, name string) *userprovisioningpb.StaticHostUser {
t.Helper()
return userprovisioning.NewStaticHostUser(name, &userprovisioningpb.StaticHostUserSpec{
Matchers: []*userprovisioningpb.Matcher{
{
NodeLabels: []*labelv1.Label{
{
Name: "foo",
Values: []string{"bar"},
},
},
Groups: []string{"foo", "bar"},
},
},
})
}
func newAutoUpdateConfig(t *testing.T) *autoupdate.AutoUpdateConfig {
t.Helper()
r, err := update.NewAutoUpdateConfig(&autoupdate.AutoUpdateConfigSpec{
Tools: &autoupdate.AutoUpdateConfigSpecTools{
Mode: update.ToolsUpdateModeEnabled,
},
})
require.NoError(t, err)
return r
}
func newAutoUpdateVersion(t *testing.T) *autoupdate.AutoUpdateVersion {
t.Helper()
r, err := update.NewAutoUpdateVersion(&autoupdate.AutoUpdateVersionSpec{
Tools: &autoupdate.AutoUpdateVersionSpecTools{
TargetVersion: "1.2.3",
},
})
require.NoError(t, err)
return r
}
func newAutoUpdateAgentRollout(t *testing.T) *autoupdate.AutoUpdateAgentRollout {
t.Helper()
r, err := update.NewAutoUpdateAgentRollout(&autoupdate.AutoUpdateAgentRolloutSpec{
StartVersion: "1.2.3",
TargetVersion: "2.3.4",
Schedule: update.AgentsScheduleImmediate,
AutoupdateMode: update.AgentsUpdateModeEnabled,
Strategy: update.AgentsStrategyTimeBased,
})
require.NoError(t, err)
return r
}
func newAutoUpdateAgentReport(t *testing.T, name string) *autoupdate.AutoUpdateAgentReport {
t.Helper()
r, err := update.NewAutoUpdateAgentReport(&autoupdate.AutoUpdateAgentReportSpec{
Timestamp: timestamppb.Now(),
Groups: map[string]*autoupdate.AutoUpdateAgentReportSpecGroup{
"foo": {
Versions: map[string]*autoupdate.AutoUpdateAgentReportSpecGroupVersion{
"1.2.3": {Count: 1},
"1.2.4": {Count: 2},
},
},
"bar": {
Versions: map[string]*autoupdate.AutoUpdateAgentReportSpecGroupVersion{
"2.3.4": {Count: 3},
"2.3.5": {Count: 4},
},
},
},
}, name)
require.NoError(t, err)
return r
}
func newAutoUpdateBotInstanceReport(t *testing.T) *autoupdate.AutoUpdateBotInstanceReport {
t.Helper()
return &autoupdate.AutoUpdateBotInstanceReport{
Kind: types.KindAutoUpdateBotInstanceReport,
Version: types.V1,
Metadata: &headerv1.Metadata{
Name: types.MetaNameAutoUpdateBotInstanceReport,
},
Spec: &autoupdate.AutoUpdateBotInstanceReportSpec{
Timestamp: timestamppb.Now(),
Groups: map[string]*autoupdate.AutoUpdateBotInstanceReportSpecGroup{
"foo": {
Versions: map[string]*autoupdate.AutoUpdateBotInstanceReportSpecGroupVersion{
"1.2.3": {Count: 1},
"1.2.4": {Count: 2},
},
},
"bar": {
Versions: map[string]*autoupdate.AutoUpdateBotInstanceReportSpecGroupVersion{
"2.3.4": {Count: 3},
"2.3.5": {Count: 4},
},
},
},
},
}
}
func newWorkloadCluster(t *testing.T, name string) *workloadclusterv1.WorkloadCluster {
t.Helper()
workloadCluster := &workloadclusterv1.WorkloadCluster{
Metadata: &headerv1.Metadata{
Name: name,
},
}
return workloadCluster
}
func withKeepalive[T any](fn func(context.Context, T) (*types.KeepAlive, error)) func(context.Context, T) error {
return func(ctx context.Context, resource T) error {
_, err := fn(ctx, resource)
return err
}
}
func modifyNoContext[T any](fn func(T) error) func(context.Context, T) error {
return func(_ context.Context, resource T) error {
return fn(resource)
}
}
// A test entity descriptor from https://sptest.iamshowcase.com/testsp_metadata.xml.
const testEntityDescriptorFmt = `<?xml version="1.0" encoding="UTF-8"?>
<md:EntityDescriptor xmlns:md="urn:oasis:names:tc:SAML:2.0:metadata" xmlns:ds="http://www.w3.org/2000/09/xmldsig#" entityID="%s" validUntil="2025-12-09T09:13:31.006Z">
<md:SPSSODescriptor AuthnRequestsSigned="false" WantAssertionsSigned="true" protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<md:NameIDFormat>urn:oasis:names:tc:SAML:1.1:nameid-format:unspecified</md:NameIDFormat>
<md:NameIDFormat>urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress</md:NameIDFormat>
<md:AssertionConsumerService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="https://sptest.iamshowcase.com/acs" index="0" isDefault="true"/>
</md:SPSSODescriptor>
</md:EntityDescriptor>
`
func fetchEvent(t *testing.T, w types.Watcher, timeout time.Duration) types.Event {
t.Helper()
timeoutC := time.After(timeout)
var ev types.Event
select {
case <-timeoutC:
require.Fail(t, "Timeout waiting for event", w.Error())
case <-w.Done():
require.Fail(t, "Watcher exited with error", w.Error())
case ev = <-w.Events():
}
return ev
}
func testResourcePagination[T any](t *testing.T, p *testPack, funcs testFuncs[T]) {
t.Helper()
const defaultTestPageSize = 2
const numberOfFullPages = 2
const totalItemCount = (numberOfFullPages * defaultTestPageSize) + 1
ctx, cancel := context.WithCancel(t.Context())
defer cancel()
go func() {
for {
select {
case <-ctx.Done():
return
case <-p.eventsC:
// Discard events to avoid blocking the test.
}
}
}()
// Generate resources
for i := range totalItemCount {
name := fmt.Sprintf("resources-%d", i)
r, err := funcs.newResource(name)
require.NoError(t, err)
require.NoError(t, funcs.create(ctx, r))
}
// Fetch all of the created items from the upstream:
expected, err := funcs.listAll(ctx)
require.NoError(t, err)
require.Len(t, expected, totalItemCount)
cmpOpts := funcs.resource.cmpOpts
// Wait for all the resources to be replicated to the cache.
require.EventuallyWithT(t, func(t *assert.CollectT) {
items, _ := funcs.cacheListAll(ctx)
assert.Len(t, items, len(expected))
}, 15*time.Second, 100*time.Millisecond)
page1, page2Start, err := funcs.cacheList(ctx, defaultTestPageSize, "")
require.NoError(t, err)
assert.Len(t, page1, defaultTestPageSize)
assert.NotEmpty(t, page2Start)
page2, page3Start, err := funcs.cacheList(ctx, defaultTestPageSize, page2Start)
require.NoError(t, err)
assert.Len(t, page2, defaultTestPageSize)
assert.NotEmpty(t, page3Start)
page3, end, err := funcs.cacheList(ctx, defaultTestPageSize, page3Start)
require.NoError(t, err)
assert.Len(t, page3, 1)
assert.Empty(t, end)
var listed []T
listed = append(listed, page1...)
listed = append(listed, page2...)
listed = append(listed, page3...)
// All items have been returned as expected
assert.Empty(t, cmp.Diff(expected, listed, cmpOpts...))
// Small pages
pageSmall, pageSmallNext, err := funcs.cacheList(ctx, 1, "")
require.NoError(t, err)
assert.Len(t, pageSmall, 1)
assert.NotEmpty(t, pageSmallNext)
if funcs.Range != nil && funcs.cacheRange != nil {
out, err := stream.Collect(funcs.cacheRange(ctx, "", page2Start))
require.NoError(t, err)
assert.Len(t, out, len(page1))
assert.Empty(t, cmp.Diff(page1, out, cmpOpts...))
out, err = stream.Collect(funcs.cacheRange(ctx, "", ""))
require.NoError(t, err)
assert.Len(t, out, len(expected))
assert.Empty(t, cmp.Diff(expected, out, cmpOpts...))
out, err = stream.Collect(funcs.cacheRange(ctx, page2Start, ""))
require.NoError(t, err)
assert.Len(t, out, len(expected)-defaultTestPageSize)
assert.Empty(t, cmp.Diff(expected, append(page1, out...), cmpOpts...))
// invalidate the cache, cover upstream fallback
p.cache.ok = false
out, err = stream.Collect(funcs.cacheRange(ctx, "", ""))
require.NoError(t, err)
assert.Len(t, out, len(expected))
assert.Empty(t, cmp.Diff(expected, out, cmpOpts...))
}
// invalidate the cache, cover upstream fallback
p.cache.ok = false
out, err := funcs.cacheListAll(ctx)
require.NoError(t, err)
assert.Len(t, out, len(expected))
assert.Empty(t, cmp.Diff(expected, out, cmpOpts...))
require.NoError(t, funcs.deleteAll(ctx))
// Wait for the cache to be empty.
require.EventuallyWithT(t, func(t *assert.CollectT) {
items, err := funcs.cacheListAll(ctx)
assert.NoError(t, err)
assert.Empty(t, items)
}, 3*time.Second, 100*time.Millisecond)
}
type resourcesLister interface {
ListResources(ctx context.Context, req proto.ListResourcesRequest) (*types.ListResourcesResponse, error)
}
func listResource(ctx context.Context, lister resourcesLister, kind string, pageSize int, pageToken string) ([]types.ResourceWithLabels, string, error) {
resp, err := lister.ListResources(ctx, proto.ListResourcesRequest{
ResourceType: kind,
Limit: int32(pageSize),
StartKey: pageToken,
})
if err != nil {
return nil, "", trace.Wrap(err)
}
return resp.Resources, resp.NextKey, nil
}
// NewServer creates a new server resource
func NewServer(kind, name, addr, namespace string) *types.ServerV2 {
return &types.ServerV2{
Kind: kind,
Version: types.V2,
Metadata: types.Metadata{
Name: name,
Namespace: namespace,
},
Spec: types.ServerSpecV2{
Addr: addr,
},
}
}