From 437dcc875cd220fd727c76f10b87b06e4f22e2bf Mon Sep 17 00:00:00 2001 From: Bruno Date: Tue, 28 Oct 2025 11:41:46 -0300 Subject: [PATCH] QueryCaching: Use CachingServiceClient for query caching (#112128) * Integrate mt querier with query caching * typo * let the caller set cache status response header * fix TestQueryAPI * make gen-go * handle CachingServiceClient being nil and make gen-go * include namespace in cache key * set signed in user namespace in query_test.go * fix test * remove commented out code * undo services/query/query.go changes * make gen-go * remove namespace requirement * fix tests * fix test * remove namespace from SignedInUser in tests * make gen-go --- pkg/api/plugin_resource_test.go | 2 +- pkg/registry/apis/query/query_test.go | 4 + pkg/server/wire.go | 2 + pkg/server/wire_gen.go | 8 +- pkg/services/caching/fake_caching_service.go | 12 +- .../caching_metrics.go => caching/metrics.go} | 2 +- pkg/services/caching/service.go | 169 ++++++++++++++-- pkg/services/caching/service_test.go | 130 ++++++++++++ .../clientmiddleware/caching_middleware.go | 186 +----------------- .../caching_middleware_test.go | 116 ++--------- .../pluginsintegration/pluginsintegration.go | 12 +- pkg/services/query/query_test.go | 7 +- 12 files changed, 345 insertions(+), 305 deletions(-) rename pkg/services/{pluginsintegration/clientmiddleware/caching_metrics.go => caching/metrics.go} (98%) create mode 100644 pkg/services/caching/service_test.go diff --git a/pkg/api/plugin_resource_test.go b/pkg/api/plugin_resource_test.go index 5f0412b5c4c..0ee413447da 100644 --- a/pkg/api/plugin_resource_test.go +++ b/pkg/api/plugin_resource_test.go @@ -173,7 +173,7 @@ func TestIntegrationCallResource(t *testing.T) { Backend: true, }, })) - middlewares := pluginsintegration.CreateMiddlewares(cfg, &oauthtokentest.Service{}, tracing.InitializeTracerForTest(), &caching.OSSCachingService{}, featuremgmt.WithFeatures(), prometheus.DefaultRegisterer, pluginRegistry) + middlewares := pluginsintegration.CreateMiddlewares(cfg, &oauthtokentest.Service{}, tracing.InitializeTracerForTest(), caching.ProvideCachingServiceClient(&caching.OSSCachingService{}, nil), featuremgmt.WithFeatures(), prometheus.DefaultRegisterer, pluginRegistry) pc, err := backend.HandlerFromMiddlewares(&pluginfakes.FakePluginClient{ CallResourceHandlerFunc: backend.CallResourceHandlerFunc(func(ctx context.Context, req *backend.CallResourceRequest, sender backend.CallResourceResponseSender) error { diff --git a/pkg/registry/apis/query/query_test.go b/pkg/registry/apis/query/query_test.go index f94cc92f317..53ee003882b 100644 --- a/pkg/registry/apis/query/query_test.go +++ b/pkg/registry/apis/query/query_test.go @@ -54,6 +54,10 @@ func (mu mockUser) GetOrgID() int64 { return -1 } +func (mu mockUser) GetNamespace() string { + return "ns" +} + func TestQueryAPI(t *testing.T) { testCases := []struct { name string diff --git a/pkg/server/wire.go b/pkg/server/wire.go index d4a98e8a3c9..b06e7e995ad 100644 --- a/pkg/server/wire.go +++ b/pkg/server/wire.go @@ -70,6 +70,7 @@ import ( "github.com/grafana/grafana/pkg/services/auth/jwt" "github.com/grafana/grafana/pkg/services/authn/authnimpl" "github.com/grafana/grafana/pkg/services/authz" + "github.com/grafana/grafana/pkg/services/caching" "github.com/grafana/grafana/pkg/services/cleanup" "github.com/grafana/grafana/pkg/services/cloudmigration/cloudmigrationimpl" "github.com/grafana/grafana/pkg/services/contexthandler" @@ -431,6 +432,7 @@ var wireBasicSet = wire.NewSet( idimpl.ProvideService, wire.Bind(new(auth.IDService), new(*idimpl.Service)), cloudmigrationimpl.ProvideService, + caching.ProvideCachingServiceClient, userimpl.ProvideVerifier, connectors.ProvideOrgRoleMapper, wire.Bind(new(user.Verifier), new(*userimpl.Verifier)), diff --git a/pkg/server/wire_gen.go b/pkg/server/wire_gen.go index ac23bdda879..bb260585b69 100644 --- a/pkg/server/wire_gen.go +++ b/pkg/server/wire_gen.go @@ -583,7 +583,8 @@ func Initialize(ctx context.Context, cfg *setting.Cfg, opts Options, apiOpts api } oauthtokenService := oauthtoken.ProvideService(socialService, authinfoimplService, cfg, registerer, serverLockService, tracingService, userAuthTokenService, featureToggles) ossCachingService := caching.ProvideCachingService() - middlewareHandler, err := pluginsintegration.ProvideClientWithMiddlewares(cfg, inMemory, oauthtokenService, tracingService, ossCachingService, featureToggles, registerer) + cachingServiceClient := caching.ProvideCachingServiceClient(ossCachingService, featureToggles) + middlewareHandler, err := pluginsintegration.ProvideClientWithMiddlewares(cfg, inMemory, oauthtokenService, tracingService, cachingServiceClient, featureToggles, registerer) if err != nil { return nil, err } @@ -1194,7 +1195,8 @@ func InitializeForTest(ctx context.Context, t sqlutil.ITestDB, testingT interfac service14 := service8.ProvideService(fileStoreManager, pluginService) oauthtokentestService := oauthtokentest.ProvideService() ossCachingService := caching.ProvideCachingService() - middlewareHandler, err := pluginsintegration.ProvideClientWithMiddlewares(cfg, inMemory, oauthtokentestService, tracingService, ossCachingService, featureToggles, registerer) + cachingServiceClient := caching.ProvideCachingServiceClient(ossCachingService, featureToggles) + middlewareHandler, err := pluginsintegration.ProvideClientWithMiddlewares(cfg, inMemory, oauthtokentestService, tracingService, cachingServiceClient, featureToggles, registerer) if err != nil { return nil, err } @@ -1712,7 +1714,7 @@ var withOTelSet = wire.NewSet( otelTracer, grpcserver.ProvideService, interceptors.ProvideAuthenticator, ) -var wireBasicSet = wire.NewSet(annotationsimpl.ProvideService, wire.Bind(new(annotations.Repository), new(*annotationsimpl.RepositoryImpl)), New, api.ProvideHTTPServer, query.ProvideService, wire.Bind(new(query.Service), new(*query.ServiceImpl)), bus.ProvideBus, wire.Bind(new(bus.Bus), new(*bus.InProcBus)), rendering.ProvideService, wire.Bind(new(rendering.Service), new(*rendering.RenderingService)), routing.ProvideRegister, wire.Bind(new(routing.RouteRegister), new(*routing.RouteRegisterImpl)), hooks.ProvideService, kvstore.ProvideService, localcache.ProvideService, bundleregistry.ProvideService, wire.Bind(new(supportbundles.Service), new(*bundleregistry.Service)), updatemanager.ProvideGrafanaService, updatemanager.ProvidePluginsService, service.ProvideService, wire.Bind(new(usagestats.Service), new(*service.UsageStats)), validator3.ProvideService, legacy.ProvideLegacyMigrator, pluginsintegration.WireSet, dashboards.ProvideFileStoreManager, wire.Bind(new(dashboards.FileStore), new(*dashboards.FileStoreManager)), cloudwatch.ProvideService, cloudmonitoring.ProvideService, azuremonitor.ProvideService, postgres.ProvideService, mysql.ProvideService, mssql.ProvideService, store.ProvideEntityEventsService, dualwrite.ProvideService, httpclientprovider.New, wire.Bind(new(httpclient.Provider), new(*httpclient2.Provider)), serverlock.ProvideService, wire.Bind(new(installsync.ServerLock), new(*serverlock.ServerLockService)), annotationsimpl.ProvideCleanupService, wire.Bind(new(annotations.Cleaner), new(*annotationsimpl.CleanupServiceImpl)), cleanup.ProvideService, shorturlimpl.ProvideService, wire.Bind(new(shorturls.Service), new(*shorturlimpl.ShortURLService)), queryhistory.ProvideService, wire.Bind(new(queryhistory.Service), new(*queryhistory.QueryHistoryService)), correlations.ProvideService, wire.Bind(new(correlations.Service), new(*correlations.CorrelationsService)), quotaimpl.ProvideService, remotecache.ProvideService, wire.Bind(new(remotecache.CacheStorage), new(*remotecache.RemoteCache)), authinfoimpl.ProvideService, wire.Bind(new(login.AuthInfoService), new(*authinfoimpl.Service)), authinfoimpl.ProvideStore, datasourceproxy.ProvideService, sort.ProvideService, search2.ProvideService, searchV2.ProvideService, searchV2.ProvideSearchHTTPService, store.ProvideService, store.ProvideSystemUsersService, live.ProvideService, pushhttp.ProvideService, contexthandler.ProvideService, service12.ProvideService, wire.Bind(new(service12.LDAP), new(*service12.LDAPImpl)), jwt.ProvideService, wire.Bind(new(jwt.JWTService), new(*jwt.AuthService)), store2.ProvideDBStore, image.ProvideDeleteExpiredService, ngalert.ProvideService, librarypanels.ProvideService, wire.Bind(new(librarypanels.Service), new(*librarypanels.LibraryPanelService)), libraryelements.ProvideService, wire.Bind(new(libraryelements.Service), new(*libraryelements.LibraryElementService)), notifications.ProvideService, notifications.ProvideSmtpService, github.ProvideFactory, tracing.ProvideService, tracing.ProvideTracingConfig, wire.Bind(new(tracing.Tracer), new(*tracing.TracingService)), withOTelSet, testdatasource.ProvideService, api4.ProvideService, opentsdb.ProvideService, socialimpl.ProvideService, influxdb.ProvideService, wire.Bind(new(social.Service), new(*socialimpl.SocialService)), tempo.ProvideService, loki.ProvideService, graphite.ProvideService, prometheus.ProvideService, elasticsearch.ProvideService, pyroscope.ProvideService, parca.ProvideService, zipkin.ProvideService, jaeger.ProvideService, service9.ProvideCacheService, wire.Bind(new(datasources.CacheService), new(*service9.CacheServiceImpl)), service2.ProvideEncryptionService, wire.Bind(new(encryption2.Internal), new(*service2.Service)), manager.ProvideSecretsService, wire.Bind(new(secrets.Service), new(*manager.SecretsService)), database.ProvideSecretsStore, wire.Bind(new(secrets.Store), new(*database.SecretsStoreImpl)), garbagecollectionworker.ProvideWorker, grafanads.ProvideService, wire.Bind(new(dashboardsnapshots.Store), new(*database5.DashboardSnapshotStore)), database5.ProvideStore, wire.Bind(new(dashboardsnapshots.Service), new(*service10.ServiceImpl)), service10.ProvideService, service9.ProvideService, wire.Bind(new(datasources.DataSourceService), new(*service9.Service)), service9.ProvideLegacyDataSourceLookup, retriever.ProvideService, wire.Bind(new(serviceaccounts.ServiceAccountRetriever), new(*retriever.Service)), ossaccesscontrol.ProvideServiceAccountPermissions, wire.Bind(new(accesscontrol.ServiceAccountPermissionsService), new(*ossaccesscontrol.ServiceAccountPermissionsService)), manager3.ProvideServiceAccountsService, proxy.ProvideServiceAccountsProxy, wire.Bind(new(serviceaccounts.Service), new(*proxy.ServiceAccountsProxy)), dsquerierclient.NewNullQSDatasourceClientBuilder, expr.ProvideService, featuremgmt.ProvideManagerService, featuremgmt.ProvideToggles, service7.ProvideDashboardServiceImpl, wire.Bind(new(dashboards2.PermissionsRegistrationService), new(*service7.DashboardServiceImpl)), service7.ProvideDashboardService, service7.ProvideDashboardProvisioningService, service7.ProvideDashboardPluginService, database2.ProvideDashboardStore, folderimpl.ProvideService, wire.Bind(new(folder.Service), new(*folderimpl.Service)), wire.Bind(new(folder.LegacyService), new(*folderimpl.Service)), folderimpl.ProvideStore, wire.Bind(new(folder.Store), new(*folderimpl.FolderStoreImpl)), service11.ProvideService, wire.Bind(new(dashboardimport.Service), new(*service11.ImportDashboardService)), service8.ProvideService, wire.Bind(new(plugindashboards.Service), new(*service8.Service)), service8.ProvideDashboardUpdater, kvstore2.ProvideService, avatar.ProvideAvatarCacheServer, statscollector.ProvideService, csrf.ProvideCSRFFilter, wire.Bind(new(csrf.Service), new(*csrf.CSRF)), ossaccesscontrol.ProvideTeamPermissions, wire.Bind(new(accesscontrol.TeamPermissionsService), new(*ossaccesscontrol.TeamPermissionsService)), ossaccesscontrol.ProvideFolderPermissions, wire.Bind(new(accesscontrol.FolderPermissionsService), new(*ossaccesscontrol.FolderPermissionsService)), ossaccesscontrol.ProvideDashboardPermissions, wire.Bind(new(accesscontrol.DashboardPermissionsService), new(*ossaccesscontrol.DashboardPermissionsService)), ossaccesscontrol.ProvideReceiverPermissionsService, wire.Bind(new(accesscontrol.ReceiverPermissionsService), new(*ossaccesscontrol.ReceiverPermissionsService)), starimpl.ProvideService, playlistimpl.ProvideService, apikeyimpl.ProvideService, dashverimpl.ProvideService, service3.ProvideService, wire.Bind(new(publicdashboards.Service), new(*service3.PublicDashboardServiceImpl)), database3.ProvideStore, wire.Bind(new(publicdashboards.Store), new(*database3.PublicDashboardStoreImpl)), metric.ProvideService, api2.ProvideApi, api3.ProvideApi, userimpl.ProvideService, orgimpl.ProvideService, orgimpl.ProvideDeletionService, statsimpl.ProvideService, grpccontext.ProvideContextHandler, grpcserver.ProvideHealthService, grpcserver.ProvideReflectionService, resolver.ProvideEntityReferenceResolver, teamimpl.ProvideService, teamapi.ProvideTeamAPI, tempuserimpl.ProvideService, loginattemptimpl.ProvideService, wire.Bind(new(loginattempt.Service), new(*loginattemptimpl.Service)), migrations2.ProvideDataSourceMigrationService, migrations2.ProvideSecretMigrationProvider, wire.Bind(new(migrations2.SecretMigrationProvider), new(*migrations2.SecretMigrationProviderImpl)), promtypemigration.ProvideAzurePromMigrationService, promtypemigration.ProvideAmazonPromMigrationService, promtypemigration.ProvidePromTypeMigrationProvider, wire.Bind(new(promtypemigration.PromTypeMigrationProvider), new(*promtypemigration.PromTypeMigrationProviderImpl)), resourcepermissions.NewActionSetService, wire.Bind(new(accesscontrol.ActionResolver), new(resourcepermissions.ActionSetService)), wire.Bind(new(pluginaccesscontrol.ActionSetRegistry), new(resourcepermissions.ActionSetService)), permreg.ProvidePermissionRegistry, acimpl.ProvideAccessControl, accesscontrol.ProvideFixedRolesLoader, dualwrite2.ProvideZanzanaReconciler, navtreeimpl.ProvideService, wire.Bind(new(accesscontrol.AccessControl), new(*acimpl.AccessControl)), wire.Bind(new(notifications.TempUserStore), new(tempuser.Service)), tagimpl.ProvideService, wire.Bind(new(tag.Service), new(*tagimpl.Service)), authnimpl.ProvideService, authnimpl.ProvideIdentitySynchronizer, authnimpl.ProvideAuthnService, authnimpl.ProvideAuthnServiceAuthenticateOnly, authnimpl.ProvideRegistration, supportbundlesimpl.ProvideService, extsvcaccounts.ProvideExtSvcAccountsService, wire.Bind(new(serviceaccounts.ExtSvcAccountsService), new(*extsvcaccounts.ExtSvcAccountsService)), registry2.ProvideExtSvcRegistry, wire.Bind(new(extsvcauth.ExternalServiceRegistry), new(*registry2.Registry)), anonstore.ProvideAnonDBStore, wire.Bind(new(anonstore.AnonStore), new(*anonstore.AnonDBStore)), loggermw.Provide, slogadapter.Provide, signingkeysimpl.ProvideEmbeddedSigningKeysService, wire.Bind(new(signingkeys.Service), new(*signingkeysimpl.Service)), ssosettingsimpl.ProvideService, wire.Bind(new(ssosettings.Service), new(*ssosettingsimpl.Service)), idimpl.ProvideService, wire.Bind(new(auth.IDService), new(*idimpl.Service)), cloudmigrationimpl.ProvideService, userimpl.ProvideVerifier, connectors.ProvideOrgRoleMapper, wire.Bind(new(user.Verifier), new(*userimpl.Verifier)), authz.WireSet, metadata.ProvideSecureValueMetadataStorage, metadata.ProvideKeeperMetadataStorage, metadata.ProvideDecryptStorage, decrypt.ProvideDecryptAuthorizer, wire.Value([]decrypt.ExtraOwnerDecrypter(nil)), decrypt.ProvideDecryptService, inline.ProvideInlineSecureValueService, encryption.ProvideDataKeyStorage, encryption.ProvideGlobalDataKeyStorage, encryption.ProvideEncryptedValueStorage, encryption.ProvideGlobalEncryptedValueStorage, service5.ProvideSecureValueService, validator.ProvideKeeperValidator, validator.ProvideSecureValueValidator, mutator.ProvideKeeperMutator, mutator.ProvideSecureValueMutator, migrator2.NewWithEngine, database4.ProvideDatabase, clock.ProvideClock, wire.Bind(new(contracts.Database), new(*database4.Database)), wire.Bind(new(contracts.Clock), new(*clock.Clock)), manager2.ProvideEncryptionManager, service4.ProvideAESGCMCipherService, resource.ProvideStorageMetrics, resource.ProvideIndexMetrics, apiserver.WireSet, apiregistry.WireSet, appregistry.WireSet, client.ProvideK8sClientWithFallback) +var wireBasicSet = wire.NewSet(annotationsimpl.ProvideService, wire.Bind(new(annotations.Repository), new(*annotationsimpl.RepositoryImpl)), New, api.ProvideHTTPServer, query.ProvideService, wire.Bind(new(query.Service), new(*query.ServiceImpl)), bus.ProvideBus, wire.Bind(new(bus.Bus), new(*bus.InProcBus)), rendering.ProvideService, wire.Bind(new(rendering.Service), new(*rendering.RenderingService)), routing.ProvideRegister, wire.Bind(new(routing.RouteRegister), new(*routing.RouteRegisterImpl)), hooks.ProvideService, kvstore.ProvideService, localcache.ProvideService, bundleregistry.ProvideService, wire.Bind(new(supportbundles.Service), new(*bundleregistry.Service)), updatemanager.ProvideGrafanaService, updatemanager.ProvidePluginsService, service.ProvideService, wire.Bind(new(usagestats.Service), new(*service.UsageStats)), validator3.ProvideService, legacy.ProvideLegacyMigrator, pluginsintegration.WireSet, dashboards.ProvideFileStoreManager, wire.Bind(new(dashboards.FileStore), new(*dashboards.FileStoreManager)), cloudwatch.ProvideService, cloudmonitoring.ProvideService, azuremonitor.ProvideService, postgres.ProvideService, mysql.ProvideService, mssql.ProvideService, store.ProvideEntityEventsService, dualwrite.ProvideService, httpclientprovider.New, wire.Bind(new(httpclient.Provider), new(*httpclient2.Provider)), serverlock.ProvideService, wire.Bind(new(installsync.ServerLock), new(*serverlock.ServerLockService)), annotationsimpl.ProvideCleanupService, wire.Bind(new(annotations.Cleaner), new(*annotationsimpl.CleanupServiceImpl)), cleanup.ProvideService, shorturlimpl.ProvideService, wire.Bind(new(shorturls.Service), new(*shorturlimpl.ShortURLService)), queryhistory.ProvideService, wire.Bind(new(queryhistory.Service), new(*queryhistory.QueryHistoryService)), correlations.ProvideService, wire.Bind(new(correlations.Service), new(*correlations.CorrelationsService)), quotaimpl.ProvideService, remotecache.ProvideService, wire.Bind(new(remotecache.CacheStorage), new(*remotecache.RemoteCache)), authinfoimpl.ProvideService, wire.Bind(new(login.AuthInfoService), new(*authinfoimpl.Service)), authinfoimpl.ProvideStore, datasourceproxy.ProvideService, sort.ProvideService, search2.ProvideService, searchV2.ProvideService, searchV2.ProvideSearchHTTPService, store.ProvideService, store.ProvideSystemUsersService, live.ProvideService, pushhttp.ProvideService, contexthandler.ProvideService, service12.ProvideService, wire.Bind(new(service12.LDAP), new(*service12.LDAPImpl)), jwt.ProvideService, wire.Bind(new(jwt.JWTService), new(*jwt.AuthService)), store2.ProvideDBStore, image.ProvideDeleteExpiredService, ngalert.ProvideService, librarypanels.ProvideService, wire.Bind(new(librarypanels.Service), new(*librarypanels.LibraryPanelService)), libraryelements.ProvideService, wire.Bind(new(libraryelements.Service), new(*libraryelements.LibraryElementService)), notifications.ProvideService, notifications.ProvideSmtpService, github.ProvideFactory, tracing.ProvideService, tracing.ProvideTracingConfig, wire.Bind(new(tracing.Tracer), new(*tracing.TracingService)), withOTelSet, testdatasource.ProvideService, api4.ProvideService, opentsdb.ProvideService, socialimpl.ProvideService, influxdb.ProvideService, wire.Bind(new(social.Service), new(*socialimpl.SocialService)), tempo.ProvideService, loki.ProvideService, graphite.ProvideService, prometheus.ProvideService, elasticsearch.ProvideService, pyroscope.ProvideService, parca.ProvideService, zipkin.ProvideService, jaeger.ProvideService, service9.ProvideCacheService, wire.Bind(new(datasources.CacheService), new(*service9.CacheServiceImpl)), service2.ProvideEncryptionService, wire.Bind(new(encryption2.Internal), new(*service2.Service)), manager.ProvideSecretsService, wire.Bind(new(secrets.Service), new(*manager.SecretsService)), database.ProvideSecretsStore, wire.Bind(new(secrets.Store), new(*database.SecretsStoreImpl)), garbagecollectionworker.ProvideWorker, grafanads.ProvideService, wire.Bind(new(dashboardsnapshots.Store), new(*database5.DashboardSnapshotStore)), database5.ProvideStore, wire.Bind(new(dashboardsnapshots.Service), new(*service10.ServiceImpl)), service10.ProvideService, service9.ProvideService, wire.Bind(new(datasources.DataSourceService), new(*service9.Service)), service9.ProvideLegacyDataSourceLookup, retriever.ProvideService, wire.Bind(new(serviceaccounts.ServiceAccountRetriever), new(*retriever.Service)), ossaccesscontrol.ProvideServiceAccountPermissions, wire.Bind(new(accesscontrol.ServiceAccountPermissionsService), new(*ossaccesscontrol.ServiceAccountPermissionsService)), manager3.ProvideServiceAccountsService, proxy.ProvideServiceAccountsProxy, wire.Bind(new(serviceaccounts.Service), new(*proxy.ServiceAccountsProxy)), dsquerierclient.NewNullQSDatasourceClientBuilder, expr.ProvideService, featuremgmt.ProvideManagerService, featuremgmt.ProvideToggles, service7.ProvideDashboardServiceImpl, wire.Bind(new(dashboards2.PermissionsRegistrationService), new(*service7.DashboardServiceImpl)), service7.ProvideDashboardService, service7.ProvideDashboardProvisioningService, service7.ProvideDashboardPluginService, database2.ProvideDashboardStore, folderimpl.ProvideService, wire.Bind(new(folder.Service), new(*folderimpl.Service)), wire.Bind(new(folder.LegacyService), new(*folderimpl.Service)), folderimpl.ProvideStore, wire.Bind(new(folder.Store), new(*folderimpl.FolderStoreImpl)), service11.ProvideService, wire.Bind(new(dashboardimport.Service), new(*service11.ImportDashboardService)), service8.ProvideService, wire.Bind(new(plugindashboards.Service), new(*service8.Service)), service8.ProvideDashboardUpdater, kvstore2.ProvideService, avatar.ProvideAvatarCacheServer, statscollector.ProvideService, csrf.ProvideCSRFFilter, wire.Bind(new(csrf.Service), new(*csrf.CSRF)), ossaccesscontrol.ProvideTeamPermissions, wire.Bind(new(accesscontrol.TeamPermissionsService), new(*ossaccesscontrol.TeamPermissionsService)), ossaccesscontrol.ProvideFolderPermissions, wire.Bind(new(accesscontrol.FolderPermissionsService), new(*ossaccesscontrol.FolderPermissionsService)), ossaccesscontrol.ProvideDashboardPermissions, wire.Bind(new(accesscontrol.DashboardPermissionsService), new(*ossaccesscontrol.DashboardPermissionsService)), ossaccesscontrol.ProvideReceiverPermissionsService, wire.Bind(new(accesscontrol.ReceiverPermissionsService), new(*ossaccesscontrol.ReceiverPermissionsService)), starimpl.ProvideService, playlistimpl.ProvideService, apikeyimpl.ProvideService, dashverimpl.ProvideService, service3.ProvideService, wire.Bind(new(publicdashboards.Service), new(*service3.PublicDashboardServiceImpl)), database3.ProvideStore, wire.Bind(new(publicdashboards.Store), new(*database3.PublicDashboardStoreImpl)), metric.ProvideService, api2.ProvideApi, api3.ProvideApi, userimpl.ProvideService, orgimpl.ProvideService, orgimpl.ProvideDeletionService, statsimpl.ProvideService, grpccontext.ProvideContextHandler, grpcserver.ProvideHealthService, grpcserver.ProvideReflectionService, resolver.ProvideEntityReferenceResolver, teamimpl.ProvideService, teamapi.ProvideTeamAPI, tempuserimpl.ProvideService, loginattemptimpl.ProvideService, wire.Bind(new(loginattempt.Service), new(*loginattemptimpl.Service)), migrations2.ProvideDataSourceMigrationService, migrations2.ProvideSecretMigrationProvider, wire.Bind(new(migrations2.SecretMigrationProvider), new(*migrations2.SecretMigrationProviderImpl)), promtypemigration.ProvideAzurePromMigrationService, promtypemigration.ProvideAmazonPromMigrationService, promtypemigration.ProvidePromTypeMigrationProvider, wire.Bind(new(promtypemigration.PromTypeMigrationProvider), new(*promtypemigration.PromTypeMigrationProviderImpl)), resourcepermissions.NewActionSetService, wire.Bind(new(accesscontrol.ActionResolver), new(resourcepermissions.ActionSetService)), wire.Bind(new(pluginaccesscontrol.ActionSetRegistry), new(resourcepermissions.ActionSetService)), permreg.ProvidePermissionRegistry, acimpl.ProvideAccessControl, accesscontrol.ProvideFixedRolesLoader, dualwrite2.ProvideZanzanaReconciler, navtreeimpl.ProvideService, wire.Bind(new(accesscontrol.AccessControl), new(*acimpl.AccessControl)), wire.Bind(new(notifications.TempUserStore), new(tempuser.Service)), tagimpl.ProvideService, wire.Bind(new(tag.Service), new(*tagimpl.Service)), authnimpl.ProvideService, authnimpl.ProvideIdentitySynchronizer, authnimpl.ProvideAuthnService, authnimpl.ProvideAuthnServiceAuthenticateOnly, authnimpl.ProvideRegistration, supportbundlesimpl.ProvideService, extsvcaccounts.ProvideExtSvcAccountsService, wire.Bind(new(serviceaccounts.ExtSvcAccountsService), new(*extsvcaccounts.ExtSvcAccountsService)), registry2.ProvideExtSvcRegistry, wire.Bind(new(extsvcauth.ExternalServiceRegistry), new(*registry2.Registry)), anonstore.ProvideAnonDBStore, wire.Bind(new(anonstore.AnonStore), new(*anonstore.AnonDBStore)), loggermw.Provide, slogadapter.Provide, signingkeysimpl.ProvideEmbeddedSigningKeysService, wire.Bind(new(signingkeys.Service), new(*signingkeysimpl.Service)), ssosettingsimpl.ProvideService, wire.Bind(new(ssosettings.Service), new(*ssosettingsimpl.Service)), idimpl.ProvideService, wire.Bind(new(auth.IDService), new(*idimpl.Service)), cloudmigrationimpl.ProvideService, caching.ProvideCachingServiceClient, userimpl.ProvideVerifier, connectors.ProvideOrgRoleMapper, wire.Bind(new(user.Verifier), new(*userimpl.Verifier)), authz.WireSet, metadata.ProvideSecureValueMetadataStorage, metadata.ProvideKeeperMetadataStorage, metadata.ProvideDecryptStorage, decrypt.ProvideDecryptAuthorizer, wire.Value([]decrypt.ExtraOwnerDecrypter(nil)), decrypt.ProvideDecryptService, inline.ProvideInlineSecureValueService, encryption.ProvideDataKeyStorage, encryption.ProvideGlobalDataKeyStorage, encryption.ProvideEncryptedValueStorage, encryption.ProvideGlobalEncryptedValueStorage, service5.ProvideSecureValueService, validator.ProvideKeeperValidator, validator.ProvideSecureValueValidator, mutator.ProvideKeeperMutator, mutator.ProvideSecureValueMutator, migrator2.NewWithEngine, database4.ProvideDatabase, clock.ProvideClock, wire.Bind(new(contracts.Database), new(*database4.Database)), wire.Bind(new(contracts.Clock), new(*clock.Clock)), manager2.ProvideEncryptionManager, service4.ProvideAESGCMCipherService, resource.ProvideStorageMetrics, resource.ProvideIndexMetrics, apiserver.WireSet, apiregistry.WireSet, appregistry.WireSet, client.ProvideK8sClientWithFallback) var wireSet = wire.NewSet( wireBasicSet, metrics.WireSet, sqlstore.ProvideService, metrics2.ProvideService, wire.Bind(new(notifications.Service), new(*notifications.NotificationService)), wire.Bind(new(notifications.WebhookSender), new(*notifications.NotificationService)), wire.Bind(new(notifications.EmailSender), new(*notifications.NotificationService)), wire.Bind(new(db.DB), new(*sqlstore.SQLStore)), prefimpl.ProvideService, oauthtoken.ProvideService, wire.Bind(new(oauthtoken.OAuthTokenService), new(*oauthtoken.Service)), wire.Bind(new(cleanup.AlertRuleService), new(*store2.DBstore)), diff --git a/pkg/services/caching/fake_caching_service.go b/pkg/services/caching/fake_caching_service.go index 522be45f02f..d4afc545bc0 100644 --- a/pkg/services/caching/fake_caching_service.go +++ b/pkg/services/caching/fake_caching_service.go @@ -10,19 +10,20 @@ import ( type FakeOSSCachingService struct { calls map[string]int + ReturnStatus CacheStatus ReturnHit bool ReturnResourceResponse CachedResourceDataResponse ReturnQueryResponse CachedQueryDataResponse } -func (f *FakeOSSCachingService) HandleQueryRequest(ctx context.Context, req *backend.QueryDataRequest) (bool, CachedQueryDataResponse) { +func (f *FakeOSSCachingService) HandleQueryRequest(ctx context.Context, req *backend.QueryDataRequest) (bool, CachedQueryDataResponse, CacheStatus) { f.calls["HandleQueryRequest"]++ - return f.ReturnHit, f.ReturnQueryResponse + return f.ReturnHit, f.ReturnQueryResponse, f.ReturnStatus } -func (f *FakeOSSCachingService) HandleResourceRequest(ctx context.Context, req *backend.CallResourceRequest) (bool, CachedResourceDataResponse) { +func (f *FakeOSSCachingService) HandleResourceRequest(ctx context.Context, req *backend.CallResourceRequest) (bool, CachedResourceDataResponse, CacheStatus) { f.calls["HandleResourceRequest"]++ - return f.ReturnHit, f.ReturnResourceResponse + return f.ReturnHit, f.ReturnResourceResponse, f.ReturnStatus } func (f *FakeOSSCachingService) AssertCalls(t *testing.T, fn string, times int) { @@ -35,7 +36,8 @@ func (f *FakeOSSCachingService) Reset() { func NewFakeOSSCachingService() *FakeOSSCachingService { fake := &FakeOSSCachingService{ - calls: map[string]int{}, + calls: map[string]int{}, + ReturnStatus: "unset", } return fake diff --git a/pkg/services/pluginsintegration/clientmiddleware/caching_metrics.go b/pkg/services/caching/metrics.go similarity index 98% rename from pkg/services/pluginsintegration/clientmiddleware/caching_metrics.go rename to pkg/services/caching/metrics.go index 2f452d550c8..ff9c229afbf 100644 --- a/pkg/services/pluginsintegration/clientmiddleware/caching_metrics.go +++ b/pkg/services/caching/metrics.go @@ -1,4 +1,4 @@ -package clientmiddleware +package caching import ( "github.com/grafana/grafana/pkg/infra/metrics" diff --git a/pkg/services/caching/service.go b/pkg/services/caching/service.go index c3da3d22222..ad494bcaee4 100644 --- a/pkg/services/caching/service.go +++ b/pkg/services/caching/service.go @@ -7,20 +7,34 @@ import ( "encoding/hex" "encoding/json" "io" + + "strconv" "strings" + "time" + + "github.com/grafana/grafana-aws-sdk/pkg/awsds" + "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana-plugin-sdk-go/backend" + "github.com/grafana/grafana/pkg/services/contexthandler" + "github.com/grafana/grafana/pkg/services/featuremgmt" + "github.com/prometheus/client_golang/prometheus" ) +type CacheStatus string + const ( - XCacheHeader = "X-Cache" - StatusHit = "HIT" - StatusMiss = "MISS" - StatusBypass = "BYPASS" - StatusError = "ERROR" - StatusDisabled = "DISABLED" + XCacheHeader = "X-Cache" + StatusHit CacheStatus = "HIT" + StatusMiss CacheStatus = "MISS" + StatusBypass CacheStatus = "BYPASS" + StatusError CacheStatus = "ERROR" + StatusDisabled CacheStatus = "DISABLED" ) +// needed to mock the function for testing +var ShouldCacheQuery = awsds.ShouldCacheQuery + type CacheQueryResponseFn func(context.Context, *backend.QueryDataResponse) type CacheResourceResponseFn func(context.Context, *backend.CallResourceResponse) @@ -49,22 +63,22 @@ type CachingService interface { // HandleQueryRequest uses a QueryDataRequest to check the cache for any existing results for that query. // If none are found, it should return false and a CachedQueryDataResponse with an UpdateCacheFn which can be used to update the results cache after the fact. // This function may populate any response headers (accessible through the context) with the cache status using the X-Cache header. - HandleQueryRequest(context.Context, *backend.QueryDataRequest) (bool, CachedQueryDataResponse) + HandleQueryRequest(ctx context.Context, req *backend.QueryDataRequest) (bool, CachedQueryDataResponse, CacheStatus) // HandleResourceRequest uses a CallResourceRequest to check the cache for any existing results for that request. If none are found, it should return false. // This function may populate any response headers (accessible through the context) with the cache status using the X-Cache header. - HandleResourceRequest(context.Context, *backend.CallResourceRequest) (bool, CachedResourceDataResponse) + HandleResourceRequest(ctx context.Context, req *backend.CallResourceRequest) (bool, CachedResourceDataResponse, CacheStatus) } // Implementation of interface - does nothing type OSSCachingService struct { } -func (s *OSSCachingService) HandleQueryRequest(ctx context.Context, req *backend.QueryDataRequest) (bool, CachedQueryDataResponse) { - return false, CachedQueryDataResponse{} +func (s *OSSCachingService) HandleQueryRequest(ctx context.Context, req *backend.QueryDataRequest) (bool, CachedQueryDataResponse, CacheStatus) { + return false, CachedQueryDataResponse{}, "" } -func (s *OSSCachingService) HandleResourceRequest(ctx context.Context, req *backend.CallResourceRequest) (bool, CachedResourceDataResponse) { - return false, CachedResourceDataResponse{} +func (s *OSSCachingService) HandleResourceRequest(ctx context.Context, req *backend.CallResourceRequest) (bool, CachedResourceDataResponse, CacheStatus) { + return false, CachedResourceDataResponse{}, "" } var _ CachingService = &OSSCachingService{} @@ -133,3 +147,134 @@ func (e *JSONEncoder) Encode(w io.Writer, v interface{}) error { func (e *JSONEncoder) Decode(r io.Reader, v interface{}) error { return json.NewDecoder(r).Decode(v) } + +// A service that provides methods to cache requests. +// It can be used to cache requests using `caching.CachingService` without reimplementing +// the caching logic at every call site. +type CachingServiceClient struct { + cachingService CachingService + features featuremgmt.FeatureToggles +} + +func ProvideCachingServiceClient(cachingService CachingService, features featuremgmt.FeatureToggles) *CachingServiceClient { + log := log.New("caching_service_client") + if err := prometheus.Register(QueryCachingRequestHistogram); err != nil { + log.Error("Error registering prometheus collector 'QueryRequestHistogram'", "error", err) + } + if err := prometheus.Register(ResourceCachingRequestHistogram); err != nil { + log.Error("Error registering prometheus collector 'ResourceRequestHistogram'", "error", err) + } + return &CachingServiceClient{cachingService: cachingService, features: features} +} + +// WithQueryDataCaching calls `f` and caches the returned value if `req` has not been cached already. +// Returns the cached value otherwise. +func (c *CachingServiceClient) WithQueryDataCaching(ctx context.Context, req *backend.QueryDataRequest, f func() (*backend.QueryDataResponse, error)) (*backend.QueryDataResponse, error) { + if c == nil || req == nil { + return f() + } + + reqCtx := contexthandler.FromContext(ctx) + + // time how long this request takes + start := time.Now() + + // First look in the query cache if enabled + hit, cr, status := c.cachingService.HandleQueryRequest(ctx, req) + + // record request duration if caching was used + if reqCtx != nil { + reqCtx.Resp.Header().Set(XCacheHeader, string(status)) + defer func() { + QueryCachingRequestHistogram.With(prometheus.Labels{ + "datasource_type": getDatasourceType(req.PluginContext), + "cache": string(status), + "query_type": getQueryType(reqCtx), + }).Observe(time.Since(start).Seconds()) + }() + } + + // Cache hit; return the response + if hit { + return cr.Response, nil + } + + // Cache miss; do the actual queries + resp, err := f() + // Update the query cache with the result for this metrics request + if err == nil && cr.UpdateCacheFn != nil { + // If AWS async caching is not enabled, use the old code path + if c.features == nil || !c.features.IsEnabled(ctx, featuremgmt.FlagAwsAsyncQueryCaching) { + cr.UpdateCacheFn(ctx, resp) + } else if reqCtx != nil { + // time how long shouldCacheQuery takes + startShouldCacheQuery := time.Now() + shouldCache := ShouldCacheQuery(resp) + ShouldCacheQueryHistogram.With(prometheus.Labels{ + "datasource_type": req.PluginContext.DataSourceInstanceSettings.Type, + "cache": string(status), + "shouldCache": strconv.FormatBool(shouldCache), + "query_type": getQueryType(reqCtx), + }).Observe(time.Since(startShouldCacheQuery).Seconds()) + + // If AWS async caching is enabled and resp is for a running async query, don't cache it + if shouldCache { + cr.UpdateCacheFn(ctx, resp) + } + } + } + + return resp, err +} + +// WithCallResourceCaching calls `f` and caches the returned value if `req` has not been cached already. +// Returns the cached value otherwise. +func (c *CachingServiceClient) WithCallResourceCaching(ctx context.Context, req *backend.CallResourceRequest, sender backend.CallResourceResponseSender, f func(backend.CallResourceResponseSender) error) error { + if c == nil || req == nil { + return f(sender) + } + + reqCtx := contexthandler.FromContext(ctx) + + // time how long this request takes + start := time.Now() + + // First look in the resource cache if enabled + hit, cr, status := c.cachingService.HandleResourceRequest(ctx, req) + + if reqCtx != nil { + reqCtx.Resp.Header().Set(XCacheHeader, string(status)) + } + // record request duration if caching was used + defer func() { + ResourceCachingRequestHistogram.With(prometheus.Labels{ + "plugin_id": req.PluginContext.PluginID, + "cache": string(status), + }).Observe(time.Since(start).Seconds()) + }() + + // Cache hit; send the response and return + if hit { + return sender.Send(cr.Response) + } + + // Cache miss; do the actual request + // If there is no update cache func, just pass in the original sender + if cr.UpdateCacheFn == nil { + return f(sender) + } + // Otherwise, intercept the responses in a wrapped sender so we can cache them first + cacheSender := backend.CallResourceResponseSenderFunc(func(res *backend.CallResourceResponse) error { + cr.UpdateCacheFn(ctx, res) + return sender.Send(res) + }) + + return f(cacheSender) +} + +func getDatasourceType(pluginCtx backend.PluginContext) string { + if pluginCtx.DataSourceInstanceSettings == nil { + return "unknown" + } + return pluginCtx.DataSourceInstanceSettings.Name +} diff --git a/pkg/services/caching/service_test.go b/pkg/services/caching/service_test.go new file mode 100644 index 00000000000..90864171469 --- /dev/null +++ b/pkg/services/caching/service_test.go @@ -0,0 +1,130 @@ +package caching + +import ( + "context" + "errors" + "net/http/httptest" + "testing" + + "github.com/grafana/grafana-plugin-sdk-go/backend" + "github.com/grafana/grafana/pkg/services/contexthandler/ctxkey" + contextmodel "github.com/grafana/grafana/pkg/services/contexthandler/model" + "github.com/grafana/grafana/pkg/web" + "github.com/stretchr/testify/require" +) + +func TestWithQueryDataCaching(t *testing.T) { + t.Run("caching is a no-op when service is nil", func(t *testing.T) { + var s *CachingServiceClient + req := backend.QueryDataRequest{} + fakeResponse := &backend.QueryDataResponse{} + response, err := s.WithQueryDataCaching(t.Context(), &req, func() (*backend.QueryDataResponse, error) { + return fakeResponse, nil + }) + require.NoError(t, err) + require.Equal(t, fakeResponse, response) + }) + + t.Run("cache status is included in the response if a request context is available", func(t *testing.T) { + fakeCachingService := NewFakeOSSCachingService() + fakeCachingService.ReturnStatus = StatusMiss + client := ProvideCachingServiceClient(fakeCachingService, nil) + + req := backend.QueryDataRequest{} + + reqCtx := &contextmodel.ReqContext{ + Context: &web.Context{ + Resp: web.NewResponseWriter("", httptest.NewRecorder()), + }, + } + ctx := context.WithValue(t.Context(), ctxkey.Key{}, reqCtx) + fakeResponse := &backend.QueryDataResponse{} + response, err := client.WithQueryDataCaching(ctx, &req, func() (*backend.QueryDataResponse, error) { + return fakeResponse, nil + }) + require.NoError(t, err) + require.Equal(t, fakeResponse, response) + require.EqualValues(t, StatusMiss, reqCtx.Resp.Header().Get(XCacheHeader)) + }) + + t.Run("caching can be used without a request context", func(t *testing.T) { + fakeCachingService := NewFakeOSSCachingService() + fakeCachingService.ReturnStatus = StatusMiss + client := ProvideCachingServiceClient(fakeCachingService, nil) + + req := backend.QueryDataRequest{} + + fakeResponse := &backend.QueryDataResponse{} + // Using the default test context, no request context. + response, err := client.WithQueryDataCaching(t.Context(), &req, func() (*backend.QueryDataResponse, error) { + return fakeResponse, nil + }) + require.NoError(t, err) + require.Equal(t, fakeResponse, response) + }) +} + +func TestWithCallResourceCaching(t *testing.T) { + t.Run("caching is a no-op when service is nil", func(t *testing.T) { + var s *CachingServiceClient + req := backend.CallResourceRequest{} + fakeErr := errors.New("oops") + err := s.WithCallResourceCaching(t.Context(), &req, nil, func(backend.CallResourceResponseSender) error { + return fakeErr + }) + require.ErrorIs(t, err, fakeErr) + }) + + t.Run("cache status is included in the response if a request context is available", func(t *testing.T) { + fakeCachingService := NewFakeOSSCachingService() + fakeCachingService.ReturnStatus = StatusMiss + client := ProvideCachingServiceClient(fakeCachingService, nil) + + req := backend.CallResourceRequest{} + + reqCtx := &contextmodel.ReqContext{ + Context: &web.Context{ + Resp: web.NewResponseWriter("", httptest.NewRecorder()), + }, + } + ctx := context.WithValue(t.Context(), ctxkey.Key{}, reqCtx) + sender := func(*backend.CallResourceResponse) error { + return nil + } + var fakeErr = errors.New("oops") + err := client.WithCallResourceCaching(ctx, &req, backend.CallResourceResponseSenderFunc(sender), func(backend.CallResourceResponseSender) error { + return fakeErr + }) + require.ErrorIs(t, err, fakeErr) + require.EqualValues(t, StatusMiss, reqCtx.Resp.Header().Get(XCacheHeader)) + }) + + t.Run("caching can be used without a request context", func(t *testing.T) { + fakeCachingService := NewFakeOSSCachingService() + fakeCachingService.ReturnStatus = StatusMiss + client := ProvideCachingServiceClient(fakeCachingService, nil) + + req := backend.CallResourceRequest{} + + sender := func(*backend.CallResourceResponse) error { + return nil + } + var fakeErr = errors.New("oops") + // Using the default test context, no request context. + err := client.WithCallResourceCaching(t.Context(), &req, backend.CallResourceResponseSenderFunc(sender), func(_ backend.CallResourceResponseSender) error { + return fakeErr + }) + require.ErrorIs(t, err, fakeErr) + }) +} + +func TestGetDatasourceType(t *testing.T) { + t.Parallel() + + require.Equal(t, "unknown", getDatasourceType(backend.PluginContext{})) + require.Equal(t, "name", getDatasourceType(backend.PluginContext{ + DataSourceInstanceSettings: &backend.DataSourceInstanceSettings{ + Name: "name", + }, + })) +} diff --git a/pkg/services/pluginsintegration/clientmiddleware/caching_middleware.go b/pkg/services/pluginsintegration/clientmiddleware/caching_middleware.go index 2ef0b24c162..14471178bd9 100644 --- a/pkg/services/pluginsintegration/clientmiddleware/caching_middleware.go +++ b/pkg/services/pluginsintegration/clientmiddleware/caching_middleware.go @@ -2,222 +2,56 @@ package clientmiddleware import ( "context" - "fmt" - "strconv" - "time" - "github.com/grafana/grafana-aws-sdk/pkg/awsds" "github.com/grafana/grafana-plugin-sdk-go/backend" - "github.com/prometheus/client_golang/prometheus" - "golang.org/x/sync/singleflight" - "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/services/caching" "github.com/grafana/grafana/pkg/services/contexthandler" - "github.com/grafana/grafana/pkg/services/featuremgmt" ) -// needed to mock the function for testing -var shouldCacheQuery = awsds.ShouldCacheQuery - // NewCachingMiddleware creates a new backend.HandlerMiddleware that will // attempt to read and write query results to the cache -func NewCachingMiddleware(cachingService caching.CachingService) backend.HandlerMiddleware { - return NewCachingMiddlewareWithFeatureManager(cachingService, nil) -} - -// NewCachingMiddlewareWithFeatureManager creates a new backend.HandlerMiddleware that will -// attempt to read and write query results to the cache with a feature manager -func NewCachingMiddlewareWithFeatureManager(cachingService caching.CachingService, features featuremgmt.FeatureToggles) backend.HandlerMiddleware { - log := log.New("caching_middleware") - if err := prometheus.Register(QueryCachingRequestHistogram); err != nil { - log.Error("Error registering prometheus collector 'QueryRequestHistogram'", "error", err) - } - if err := prometheus.Register(ResourceCachingRequestHistogram); err != nil { - log.Error("Error registering prometheus collector 'ResourceRequestHistogram'", "error", err) - } +func NewCachingMiddleware(cachingServiceClient *caching.CachingServiceClient) backend.HandlerMiddleware { cachingMiddlewareHandler := func(next backend.Handler) backend.Handler { - cachingMiddleware := &CachingMiddleware{ - BaseHandler: backend.NewBaseHandler(next), - caching: cachingService, - log: log, - features: features, + return &CachingMiddleware{ + BaseHandler: backend.NewBaseHandler(next), + cachingServiceClient: cachingServiceClient, } - if features != nil && features.IsEnabled(context.Background(), featuremgmt.FlagQueryCacheRequestDeduplication) { - return newRequestDeduplicationMiddleware(log, cachingMiddleware) - } - return cachingMiddleware } return backend.HandlerMiddlewareFunc(cachingMiddlewareHandler) } +// An adapter to use CachingServiceClient as a middleware. If possible prefer to use `CachingServiceClient` directly. type CachingMiddleware struct { backend.BaseHandler - caching caching.CachingService - log log.Logger - features featuremgmt.FeatureToggles + cachingServiceClient *caching.CachingServiceClient } // QueryData receives a data request and attempts to access results already stored in the cache for that request. // If data is found, it will return it immediately. Otherwise, it will perform the queries as usual, then write the response to the cache. // If the cache service is implemented, we capture the request duration as a metric. The service is expected to write any response headers. func (m *CachingMiddleware) QueryData(ctx context.Context, req *backend.QueryDataRequest) (*backend.QueryDataResponse, error) { - if req == nil { - return m.BaseHandler.QueryData(ctx, req) - } - reqCtx := contexthandler.FromContext(ctx) if reqCtx == nil { return m.BaseHandler.QueryData(ctx, req) } - - // time how long this request takes - start := time.Now() - - // First look in the query cache if enabled - hit, cr := m.caching.HandleQueryRequest(ctx, req) - - // record request duration if caching was used - ch := reqCtx.Resp.Header().Get(caching.XCacheHeader) - if ch != "" { - defer func() { - QueryCachingRequestHistogram.With(prometheus.Labels{ - "datasource_type": req.PluginContext.DataSourceInstanceSettings.Type, - "cache": ch, - "query_type": getQueryType(reqCtx), - }).Observe(time.Since(start).Seconds()) - }() - } - - // Cache hit; return the response - if hit { - return cr.Response, nil - } - - // Cache miss; do the actual queries - resp, err := m.BaseHandler.QueryData(ctx, req) - - // Update the query cache with the result for this metrics request - if err == nil && cr.UpdateCacheFn != nil { - // If AWS async caching is not enabled, use the old code path - if m.features == nil || !m.features.IsEnabled(ctx, featuremgmt.FlagAwsAsyncQueryCaching) { - cr.UpdateCacheFn(ctx, resp) - } else { - // time how long shouldCacheQuery takes - startShouldCacheQuery := time.Now() - shouldCache := shouldCacheQuery(resp) - ShouldCacheQueryHistogram.With(prometheus.Labels{ - "datasource_type": req.PluginContext.DataSourceInstanceSettings.Type, - "cache": ch, - "shouldCache": strconv.FormatBool(shouldCache), - "query_type": getQueryType(reqCtx), - }).Observe(time.Since(startShouldCacheQuery).Seconds()) - - // If AWS async caching is enabled and resp is for a running async query, don't cache it - if shouldCache { - cr.UpdateCacheFn(ctx, resp) - } - } - } - - return resp, err + return m.cachingServiceClient.WithQueryDataCaching(ctx, req, func() (*backend.QueryDataResponse, error) { + return m.BaseHandler.QueryData(ctx, req) + }) } // CallResource receives a resource request and attempts to access results already stored in the cache for that request. // If data is found, it will return it immediately. Otherwise, it will perform the request as usual. The caller of CallResource is expected to explicitly update the cache with any responses. // If the cache service is implemented, we capture the request duration as a metric. The service is expected to write any response headers. func (m *CachingMiddleware) CallResource(ctx context.Context, req *backend.CallResourceRequest, sender backend.CallResourceResponseSender) error { - if req == nil { - return m.BaseHandler.CallResource(ctx, req, sender) - } - reqCtx := contexthandler.FromContext(ctx) if reqCtx == nil { return m.BaseHandler.CallResource(ctx, req, sender) } - // time how long this request takes - start := time.Now() - - // First look in the resource cache if enabled - hit, cr := m.caching.HandleResourceRequest(ctx, req) - - // record request duration if caching was used - if ch := reqCtx.Resp.Header().Get(caching.XCacheHeader); ch != "" { - defer func() { - ResourceCachingRequestHistogram.With(prometheus.Labels{ - "plugin_id": req.PluginContext.PluginID, - "cache": ch, - }).Observe(time.Since(start).Seconds()) - }() - } - - // Cache hit; send the response and return - if hit { - return sender.Send(cr.Response) - } - - // Cache miss; do the actual request - // If there is no update cache func, just pass in the original sender - if cr.UpdateCacheFn == nil { + return m.cachingServiceClient.WithCallResourceCaching(ctx, req, sender, func(sender backend.CallResourceResponseSender) error { return m.BaseHandler.CallResource(ctx, req, sender) - } - // Otherwise, intercept the responses in a wrapped sender so we can cache them first - cacheSender := backend.CallResourceResponseSenderFunc(func(res *backend.CallResourceResponse) error { - cr.UpdateCacheFn(ctx, res) - return sender.Send(res) }) - - return m.BaseHandler.CallResource(ctx, req, cacheSender) -} - -// Given N requests happening at the same time and issuing the same query, only one request will execute -// and the other ones will wait for the response received by the request being executed. -type requestDeduplicationMiddleware struct { - backend.BaseHandler - log *log.ConcreteLogger - singleflight *singleflight.Group -} - -func newRequestDeduplicationMiddleware(log *log.ConcreteLogger, next backend.Handler) *requestDeduplicationMiddleware { - return &requestDeduplicationMiddleware{log: log, BaseHandler: backend.NewBaseHandler(next), singleflight: &singleflight.Group{}} -} - -func (m *requestDeduplicationMiddleware) QueryData(ctx context.Context, req *backend.QueryDataRequest) (*backend.QueryDataResponse, error) { - if req.PluginContext.DataSourceInstanceSettings == nil || req.PluginContext.DataSourceInstanceSettings.UID == "" { - return m.BaseHandler.QueryData(ctx, req) - } - key, err := caching.GetKey(req.PluginContext.DataSourceInstanceSettings.UID, req) - if err != nil { - m.log.Error("error building cache key for request deduplication, skipping request deduplication", "error", err) - return m.BaseHandler.QueryData(ctx, req) - } - v, err, _ := m.singleflight.Do(key, func() (interface{}, error) { - return m.BaseHandler.QueryData(ctx, req) - }) - if err != nil { - return nil, fmt.Errorf("request deduplication middleware: calling BaseHandler.QueryData: %w", err) - } - return v.(*backend.QueryDataResponse), nil -} - -func (m *requestDeduplicationMiddleware) CallResource(ctx context.Context, req *backend.CallResourceRequest, sender backend.CallResourceResponseSender) error { - if req.PluginContext.DataSourceInstanceSettings == nil || req.PluginContext.DataSourceInstanceSettings.UID == "" { - return m.BaseHandler.CallResource(ctx, req, sender) - } - - key, err := caching.GetKey(req.PluginContext.DataSourceInstanceSettings.UID, req) - if err != nil { - m.log.Error("error building cache key for request deduplication, skipping request deduplication", "error", err) - return m.BaseHandler.CallResource(ctx, req, sender) - } - _, err, _ = m.singleflight.Do(key, func() (interface{}, error) { - return nil, m.BaseHandler.CallResource(ctx, req, sender) - }) - if err != nil { - return fmt.Errorf("request deduplication middleware: calling BaseHandler.CallResource: %w", err) - } - return nil } diff --git a/pkg/services/pluginsintegration/clientmiddleware/caching_middleware_test.go b/pkg/services/pluginsintegration/clientmiddleware/caching_middleware_test.go index 8185a9c66c3..f0e985b3047 100644 --- a/pkg/services/pluginsintegration/clientmiddleware/caching_middleware_test.go +++ b/pkg/services/pluginsintegration/clientmiddleware/caching_middleware_test.go @@ -4,10 +4,7 @@ import ( "context" "encoding/json" "net/http" - "sync" - "sync/atomic" "testing" - "time" "github.com/grafana/grafana-plugin-sdk-go/backend" "github.com/grafana/grafana-plugin-sdk-go/backend/handlertest" @@ -25,9 +22,10 @@ func TestCachingMiddleware(t *testing.T) { require.NoError(t, err) cs := caching.NewFakeOSSCachingService() + cachingServiceClient := caching.ProvideCachingServiceClient(cs, nil) cdt := handlertest.NewHandlerMiddlewareTest(t, WithReqContext(req, &user.SignedInUser{}), - handlertest.WithMiddlewares(NewCachingMiddleware(cs)), + handlertest.WithMiddlewares(NewCachingMiddleware(cachingServiceClient)), ) jsonDataMap := map[string]any{} @@ -78,9 +76,9 @@ func TestCachingMiddleware(t *testing.T) { }) t.Run("If cache returns a miss, queries are issued and the update cache function is called", func(t *testing.T) { - origShouldCacheQuery := shouldCacheQuery + origShouldCacheQuery := caching.ShouldCacheQuery var shouldCacheQueryCalled bool - shouldCacheQuery = func(resp *backend.QueryDataResponse) bool { + caching.ShouldCacheQuery = func(resp *backend.QueryDataResponse) bool { shouldCacheQueryCalled = true return true } @@ -88,7 +86,7 @@ func TestCachingMiddleware(t *testing.T) { t.Cleanup(func() { updateCacheCalled = false shouldCacheQueryCalled = false - shouldCacheQuery = origShouldCacheQuery + caching.ShouldCacheQuery = origShouldCacheQuery cs.Reset() }) @@ -108,15 +106,16 @@ func TestCachingMiddleware(t *testing.T) { }) t.Run("with async queries", func(t *testing.T) { + cachingServiceClient := caching.ProvideCachingServiceClient(cs, featuremgmt.WithFeatures(featuremgmt.FlagAwsAsyncQueryCaching)) asyncCdt := handlertest.NewHandlerMiddlewareTest(t, WithReqContext(req, &user.SignedInUser{}), handlertest.WithMiddlewares( - NewCachingMiddlewareWithFeatureManager(cs, featuremgmt.WithFeatures(featuremgmt.FlagAwsAsyncQueryCaching))), + NewCachingMiddleware(cachingServiceClient)), ) t.Run("If shoudCacheQuery returns true update cache function is called", func(t *testing.T) { - origShouldCacheQuery := shouldCacheQuery + origShouldCacheQuery := caching.ShouldCacheQuery var shouldCacheQueryCalled bool - shouldCacheQuery = func(resp *backend.QueryDataResponse) bool { + caching.ShouldCacheQuery = func(resp *backend.QueryDataResponse) bool { shouldCacheQueryCalled = true return true } @@ -124,7 +123,7 @@ func TestCachingMiddleware(t *testing.T) { t.Cleanup(func() { updateCacheCalled = false shouldCacheQueryCalled = false - shouldCacheQuery = origShouldCacheQuery + caching.ShouldCacheQuery = origShouldCacheQuery cs.Reset() }) @@ -144,9 +143,9 @@ func TestCachingMiddleware(t *testing.T) { }) t.Run("If shoudCacheQuery returns false update cache function is not called", func(t *testing.T) { - origShouldCacheQuery := shouldCacheQuery + origShouldCacheQuery := caching.ShouldCacheQuery var shouldCacheQueryCalled bool - shouldCacheQuery = func(resp *backend.QueryDataResponse) bool { + caching.ShouldCacheQuery = func(resp *backend.QueryDataResponse) bool { shouldCacheQueryCalled = true return false } @@ -154,7 +153,7 @@ func TestCachingMiddleware(t *testing.T) { t.Cleanup(func() { updateCacheCalled = false shouldCacheQueryCalled = false - shouldCacheQuery = origShouldCacheQuery + caching.ShouldCacheQuery = origShouldCacheQuery cs.Reset() }) @@ -199,9 +198,10 @@ func TestCachingMiddleware(t *testing.T) { } cs := caching.NewFakeOSSCachingService() + cachingServiceClient := caching.ProvideCachingServiceClient(cs, nil) cdt := handlertest.NewHandlerMiddlewareTest(t, WithReqContext(req, &user.SignedInUser{}), - handlertest.WithMiddlewares(NewCachingMiddleware(cs)), + handlertest.WithMiddlewares(NewCachingMiddleware(cachingServiceClient)), handlertest.WithResourceResponses([]*backend.CallResourceResponse{simulatedPluginResponse}), ) @@ -275,9 +275,10 @@ func TestCachingMiddleware(t *testing.T) { require.NoError(t, err) cs := caching.NewFakeOSSCachingService() + cachingServiceClient := caching.ProvideCachingServiceClient(cs, nil) cdt := handlertest.NewHandlerMiddlewareTest(t, // Skip the request context in this case - handlertest.WithMiddlewares(NewCachingMiddleware(cs)), + handlertest.WithMiddlewares(NewCachingMiddleware(cachingServiceClient)), ) reqCtx := contexthandler.FromContext(req.Context()) require.Nil(t, reqCtx) @@ -325,86 +326,3 @@ func TestCachingMiddleware(t *testing.T) { }) }) } - -func TestRequestDeduplicationMiddleware(t *testing.T) { - t.Parallel() - - t.Run("deduplicates requests issuing the same query", func(t *testing.T) { - t.Parallel() - - handler := newMockMiddlewareHandler() - middleware := newRequestDeduplicationMiddleware(nil, handler) - - req := backend.QueryDataRequest{ - PluginContext: backend.PluginContext{ - DataSourceInstanceSettings: &backend.DataSourceInstanceSettings{ - UID: "uid", - }, - }, - } - - wg := &sync.WaitGroup{} - wg.Add(2) - - for range 2 { - go func() { - defer wg.Done() - resp, err := middleware.QueryData(t.Context(), &req) - require.NoError(t, err) - require.Equal(t, &backend.QueryDataResponse{}, resp) - }() - } - - wg.Wait() - - require.EqualValues(t, 1, handler.QueryDataCalls) - }) - - t.Run("requests where DataSourceInstanceSettings is nil bypass request deduplication", func(t *testing.T) { - t.Parallel() - - handler := newMockMiddlewareHandler() - middleware := newRequestDeduplicationMiddleware(nil, handler) - - { - req := backend.QueryDataRequest{ - PluginContext: backend.PluginContext{ - DataSourceInstanceSettings: nil, - }, - } - - resp, err := middleware.QueryData(t.Context(), &req) - require.NoError(t, err) - require.Empty(t, resp) - } - - { - req := backend.CallResourceRequest{ - PluginContext: backend.PluginContext{ - DataSourceInstanceSettings: nil, - }, - } - - require.NoError(t, middleware.CallResource(t.Context(), &req, nil)) - } - }) -} - -type mockMiddlewareHandler struct { - backend.BaseHandler - QueryDataCalls int32 -} - -func newMockMiddlewareHandler() *mockMiddlewareHandler { - return &mockMiddlewareHandler{} -} - -func (m *mockMiddlewareHandler) QueryData(ctx context.Context, req *backend.QueryDataRequest) (*backend.QueryDataResponse, error) { - atomic.AddInt32(&m.QueryDataCalls, 1) - time.Sleep(10 * time.Millisecond) - return &backend.QueryDataResponse{}, nil -} - -func (m *mockMiddlewareHandler) CallResource(ctx context.Context, req *backend.CallResourceRequest, sender backend.CallResourceResponseSender) error { - return nil -} diff --git a/pkg/services/pluginsintegration/pluginsintegration.go b/pkg/services/pluginsintegration/pluginsintegration.go index 46bc10cb76b..a510ead1910 100644 --- a/pkg/services/pluginsintegration/pluginsintegration.go +++ b/pkg/services/pluginsintegration/pluginsintegration.go @@ -167,25 +167,25 @@ func ProvideClientWithMiddlewares( pluginRegistry registry.Service, oAuthTokenService oauthtoken.OAuthTokenService, tracer tracing.Tracer, - cachingService caching.CachingService, + cachingServiceClient *caching.CachingServiceClient, features featuremgmt.FeatureToggles, promRegisterer prometheus.Registerer, ) (*backend.MiddlewareHandler, error) { - return NewMiddlewareHandler(cfg, pluginRegistry, oAuthTokenService, tracer, cachingService, features, promRegisterer, pluginRegistry) + return NewMiddlewareHandler(cfg, pluginRegistry, oAuthTokenService, tracer, cachingServiceClient, features, promRegisterer, pluginRegistry) } func NewMiddlewareHandler( cfg *setting.Cfg, pluginRegistry registry.Service, oAuthTokenService oauthtoken.OAuthTokenService, - tracer tracing.Tracer, cachingService caching.CachingService, features featuremgmt.FeatureToggles, + tracer tracing.Tracer, cachingServiceClient *caching.CachingServiceClient, features featuremgmt.FeatureToggles, promRegisterer prometheus.Registerer, registry registry.Service, ) (*backend.MiddlewareHandler, error) { c := client.ProvideService(pluginRegistry) - middlewares := CreateMiddlewares(cfg, oAuthTokenService, tracer, cachingService, features, promRegisterer, registry) + middlewares := CreateMiddlewares(cfg, oAuthTokenService, tracer, cachingServiceClient, features, promRegisterer, registry) return backend.HandlerFromMiddlewares(c, middlewares...) } -func CreateMiddlewares(cfg *setting.Cfg, oAuthTokenService oauthtoken.OAuthTokenService, tracer tracing.Tracer, cachingService caching.CachingService, features featuremgmt.FeatureToggles, promRegisterer prometheus.Registerer, registry registry.Service) []backend.HandlerMiddleware { +func CreateMiddlewares(cfg *setting.Cfg, oAuthTokenService oauthtoken.OAuthTokenService, tracer tracing.Tracer, cachingServiceClient *caching.CachingServiceClient, features featuremgmt.FeatureToggles, promRegisterer prometheus.Registerer, registry registry.Service) []backend.HandlerMiddleware { middlewares := []backend.HandlerMiddleware{ clientmiddleware.NewTracingMiddleware(tracer), clientmiddleware.NewMetricsMiddleware(promRegisterer, registry), @@ -203,7 +203,7 @@ func CreateMiddlewares(cfg *setting.Cfg, oAuthTokenService oauthtoken.OAuthToken clientmiddleware.NewClearAuthHeadersMiddleware(), clientmiddleware.NewOAuthTokenMiddleware(oAuthTokenService), clientmiddleware.NewCookiesMiddleware(skipCookiesNames), - clientmiddleware.NewCachingMiddlewareWithFeatureManager(cachingService, features), + clientmiddleware.NewCachingMiddleware(cachingServiceClient), clientmiddleware.NewForwardIDMiddleware(), clientmiddleware.NewUseAlertHeadersMiddleware(), ) diff --git a/pkg/services/query/query_test.go b/pkg/services/query/query_test.go index aead5e3cc00..02a4dea4824 100644 --- a/pkg/services/query/query_test.go +++ b/pkg/services/query/query_test.go @@ -844,7 +844,7 @@ func setup(t *testing.T, isMultiTenant bool, mockClient clientapi.QueryDataClien secretStore: ss, pluginRequestValidator: rv, queryService: queryService, - signedInUser: &user.SignedInUser{OrgID: 1, Login: "login", Name: "name", Email: "email", OrgRole: identity.RoleAdmin}, + signedInUser: &user.SignedInUser{OrgID: 1, Login: "login", Name: "name", Email: "email", OrgRole: identity.RoleAdmin, Namespace: "ns1"}, } } @@ -928,12 +928,15 @@ func (c *fakePluginClient) QueryData(ctx context.Context, req *backend.QueryData } type testClient struct { - queryDataLastCalledWith data.QueryDataRequest + queryDataLastCalledWith data.QueryDataRequest + // The number of times the QueryData method has been called + queryDataCalls int queryDataStubbedResponse *backend.QueryDataResponse queryDataStubbedError error } func (c *testClient) QueryData(ctx context.Context, req data.QueryDataRequest) (*backend.QueryDataResponse, error) { + c.queryDataCalls++ c.queryDataLastCalledWith = req if c.queryDataStubbedError != nil { return nil, c.queryDataStubbedError