Merge remote-tracking branch 'origin/main' into resource-store
This commit is contained in:
@@ -108,6 +108,10 @@ func (hs *HTTPServer) registerRoutes() {
|
||||
r.Get("/admin/storage/*", reqSignedIn, hs.Index)
|
||||
}
|
||||
|
||||
if hs.Features.IsEnabledGlobally(featuremgmt.FlagOnPremToCloudMigrations) {
|
||||
r.Get("/admin/migrate-to-cloud", reqOrgAdmin, hs.Index)
|
||||
}
|
||||
|
||||
// feature toggle admin page
|
||||
if hs.Features.IsEnabledGlobally(featuremgmt.FlagFeatureToggleAdminPage) {
|
||||
r.Get("/admin/featuretoggles", authorize(ac.EvalPermission(ac.ActionFeatureManagementRead)), hs.Index)
|
||||
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
acdb "github.com/grafana/grafana/pkg/services/accesscontrol/database"
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol/ossaccesscontrol"
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol/resourcepermissions"
|
||||
"github.com/grafana/grafana/pkg/services/authz/zanzana"
|
||||
"github.com/grafana/grafana/pkg/services/contexthandler/ctxkey"
|
||||
contextmodel "github.com/grafana/grafana/pkg/services/contexthandler/model"
|
||||
"github.com/grafana/grafana/pkg/services/dashboards"
|
||||
@@ -460,7 +461,10 @@ func setupServer(b testing.TB, sc benchScenario, features featuremgmt.FeatureTog
|
||||
|
||||
cfg := setting.NewCfg()
|
||||
actionSets := resourcepermissions.NewActionSetService()
|
||||
acSvc := acimpl.ProvideOSSService(sc.cfg, acdb.ProvideService(sc.db), actionSets, localcache.ProvideService(), features, tracing.InitializeTracerForTest())
|
||||
acSvc := acimpl.ProvideOSSService(
|
||||
sc.cfg, acdb.ProvideService(sc.db), actionSets, localcache.ProvideService(),
|
||||
features, tracing.InitializeTracerForTest(), zanzana.NewNoopClient(), sc.db,
|
||||
)
|
||||
|
||||
folderPermissions, err := ossaccesscontrol.ProvideFolderPermissions(
|
||||
cfg, features, routing.NewRouteRegister(), sc.db, ac, license, &dashboards.FakeDashboardStore{}, folderServiceWithFlagOn, acSvc, sc.teamSvc, sc.userSvc, actionSets)
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
"github.com/grafana/grafana/pkg/infra/tracing"
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol"
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol/acimpl"
|
||||
"github.com/grafana/grafana/pkg/services/authz/zanzana"
|
||||
"github.com/grafana/grafana/pkg/services/featuremgmt"
|
||||
"github.com/grafana/grafana/pkg/services/quota/quotaimpl"
|
||||
"github.com/grafana/grafana/pkg/services/sqlstore"
|
||||
@@ -89,7 +90,7 @@ func initializeConflictResolver(cmd *utils.ContextCommandLine, f Formatter, ctx
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%v: %w", "failed to initialize tracer service", err)
|
||||
}
|
||||
acService, err := acimpl.ProvideService(cfg, s, routing, nil, nil, nil, features, tracer)
|
||||
acService, err := acimpl.ProvideService(cfg, s, routing, nil, nil, nil, features, tracer, zanzana.NewNoopClient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%v: %w", "failed to get access control", err)
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"os"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
"go.opentelemetry.io/otel"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
genericapiserver "k8s.io/apiserver/pkg/server"
|
||||
"k8s.io/component-base/cli"
|
||||
|
||||
@@ -108,14 +110,31 @@ func RunCLI(opts commands.ServerOptions) int {
|
||||
return cli.Run(cmd)
|
||||
}
|
||||
|
||||
type lateInitializedTracingProvider struct {
|
||||
trace.TracerProvider
|
||||
tracer *lateInitializedTracingService
|
||||
}
|
||||
|
||||
func (tp lateInitializedTracingProvider) Tracer(name string, options ...trace.TracerOption) trace.Tracer {
|
||||
return tp.tracer
|
||||
}
|
||||
|
||||
type lateInitializedTracingService struct {
|
||||
tracing.Tracer
|
||||
}
|
||||
|
||||
func newLateInitializedTracingService() *lateInitializedTracingService {
|
||||
return &lateInitializedTracingService{
|
||||
ts := &lateInitializedTracingService{
|
||||
Tracer: tracing.InitializeTracerForTest(),
|
||||
}
|
||||
|
||||
tp := &lateInitializedTracingProvider{
|
||||
tracer: ts,
|
||||
}
|
||||
|
||||
otel.SetTracerProvider(tp)
|
||||
|
||||
return ts
|
||||
}
|
||||
|
||||
func (s *lateInitializedTracingService) InitTracer(tracer tracing.Tracer) {
|
||||
|
||||
@@ -150,7 +150,7 @@ func (r *queryREST) Connect(connectCtx context.Context, name string, _ runtime.O
|
||||
func (b *QueryAPIBuilder) execute(ctx context.Context, req parsedRequestInfo) (qdr *backend.QueryDataResponse, err error) {
|
||||
switch len(req.Requests) {
|
||||
case 0:
|
||||
break // nothing to do
|
||||
qdr = &backend.QueryDataResponse{}
|
||||
case 1:
|
||||
qdr, err = b.handleQuerySingleDatasource(ctx, req.Requests[0])
|
||||
default:
|
||||
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol/migrator"
|
||||
"github.com/grafana/grafana/pkg/services/accesscontrol/pluginutils"
|
||||
"github.com/grafana/grafana/pkg/services/authn"
|
||||
"github.com/grafana/grafana/pkg/services/authz/zanzana"
|
||||
"github.com/grafana/grafana/pkg/services/dashboards"
|
||||
"github.com/grafana/grafana/pkg/services/featuremgmt"
|
||||
"github.com/grafana/grafana/pkg/services/folder"
|
||||
@@ -46,8 +47,12 @@ var SharedWithMeFolderPermission = accesscontrol.Permission{
|
||||
|
||||
var OSSRolesPrefixes = []string{accesscontrol.ManagedRolePrefix, accesscontrol.ExternalServiceRolePrefix}
|
||||
|
||||
func ProvideService(cfg *setting.Cfg, db db.DB, routeRegister routing.RouteRegister, cache *localcache.CacheService, accessControl accesscontrol.AccessControl, actionResolver accesscontrol.ActionResolver, features featuremgmt.FeatureToggles, tracer tracing.Tracer) (*Service, error) {
|
||||
service := ProvideOSSService(cfg, database.ProvideService(db), actionResolver, cache, features, tracer)
|
||||
func ProvideService(
|
||||
cfg *setting.Cfg, db db.DB, routeRegister routing.RouteRegister, cache *localcache.CacheService,
|
||||
accessControl accesscontrol.AccessControl, actionResolver accesscontrol.ActionResolver,
|
||||
features featuremgmt.FeatureToggles, tracer tracing.Tracer, zclient zanzana.Client,
|
||||
) (*Service, error) {
|
||||
service := ProvideOSSService(cfg, database.ProvideService(db), actionResolver, cache, features, tracer, zclient, db)
|
||||
|
||||
api.NewAccessControlAPI(routeRegister, accessControl, service, features).RegisterAPIEndpoints()
|
||||
if err := accesscontrol.DeclareFixedRoles(service, cfg); err != nil {
|
||||
@@ -65,7 +70,11 @@ func ProvideService(cfg *setting.Cfg, db db.DB, routeRegister routing.RouteRegis
|
||||
return service, nil
|
||||
}
|
||||
|
||||
func ProvideOSSService(cfg *setting.Cfg, store accesscontrol.Store, actionResolver accesscontrol.ActionResolver, cache *localcache.CacheService, features featuremgmt.FeatureToggles, tracer tracing.Tracer) *Service {
|
||||
func ProvideOSSService(
|
||||
cfg *setting.Cfg, store accesscontrol.Store, actionResolver accesscontrol.ActionResolver,
|
||||
cache *localcache.CacheService, features featuremgmt.FeatureToggles, tracer tracing.Tracer,
|
||||
zclient zanzana.Client, db db.DB,
|
||||
) *Service {
|
||||
s := &Service{
|
||||
actionResolver: actionResolver,
|
||||
cache: cache,
|
||||
@@ -75,6 +84,7 @@ func ProvideOSSService(cfg *setting.Cfg, store accesscontrol.Store, actionResolv
|
||||
roles: accesscontrol.BuildBasicRoleDefinitions(),
|
||||
store: store,
|
||||
tracer: tracer,
|
||||
sync: migrator.NewZanzanaSynchroniser(zclient, db),
|
||||
}
|
||||
|
||||
return s
|
||||
@@ -91,6 +101,7 @@ type Service struct {
|
||||
roles map[string]*accesscontrol.RoleDTO
|
||||
store accesscontrol.Store
|
||||
tracer tracing.Tracer
|
||||
sync *migrator.ZanzanaSynchroniser
|
||||
}
|
||||
|
||||
func (s *Service) GetUsageStats(_ context.Context) map[string]any {
|
||||
@@ -397,6 +408,13 @@ func (s *Service) RegisterFixedRoles(ctx context.Context) error {
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
if s.features.IsEnabledGlobally(featuremgmt.FlagZanzana) {
|
||||
if err := s.sync.Sync(context.Background()); err != nil {
|
||||
s.log.Error("Failed to synchronise permissions to zanzana ", "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -69,6 +69,8 @@ func TestUsageMetrics(t *testing.T) {
|
||||
localcache.ProvideService(),
|
||||
featuremgmt.WithFeatures(),
|
||||
tracing.InitializeTracerForTest(),
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
assert.Equal(t, tt.expectedValue, s.GetUsageStats(context.Background())["stats.oss.accesscontrol.enabled.count"])
|
||||
})
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
package migrator
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
openfgav1 "github.com/openfga/api/proto/openfga/v1"
|
||||
|
||||
"github.com/grafana/grafana/pkg/infra/db"
|
||||
"github.com/grafana/grafana/pkg/infra/log"
|
||||
"github.com/grafana/grafana/pkg/services/authz/zanzana"
|
||||
)
|
||||
|
||||
// A TupleCollector is responsible to build and store [openfgav1.TupleKey] into provided tuple map.
|
||||
// They key used should be a unique group key for the collector so we can skip over an already synced group.
|
||||
type TupleCollector func(ctx context.Context, tuples map[string][]*openfgav1.TupleKey) error
|
||||
|
||||
// ZanzanaSynchroniser is a component to sync RBAC permissions to zanzana.
|
||||
// We should rewrite the migration after we have "migrated" all possible actions
|
||||
// into our schema. This will only do a one time migration for each action so its
|
||||
// is not really syncing the full rbac state. If a fresh sync is needed the tuple
|
||||
// needs to be cleared first.
|
||||
type ZanzanaSynchroniser struct {
|
||||
log log.Logger
|
||||
client zanzana.Client
|
||||
collectors []TupleCollector
|
||||
}
|
||||
|
||||
func NewZanzanaSynchroniser(client zanzana.Client, store db.DB, collectors ...TupleCollector) *ZanzanaSynchroniser {
|
||||
// Append shared collectors that is used by both enterprise and oss
|
||||
collectors = append(collectors, managedPermissionsCollector(store))
|
||||
|
||||
return &ZanzanaSynchroniser{
|
||||
log: log.New("zanzana.sync"),
|
||||
collectors: collectors,
|
||||
}
|
||||
}
|
||||
|
||||
// Sync runs all collectors and tries to write all collected tuples.
|
||||
// It will skip over any "sync group" that has already been written.
|
||||
func (z *ZanzanaSynchroniser) Sync(ctx context.Context) error {
|
||||
tuplesMap := make(map[string][]*openfgav1.TupleKey)
|
||||
|
||||
for _, c := range z.collectors {
|
||||
if err := c(ctx, tuplesMap); err != nil {
|
||||
return fmt.Errorf("failed to collect permissions: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
for key, tuples := range tuplesMap {
|
||||
if err := batch(len(tuples), 100, func(start, end int) error {
|
||||
return z.client.Write(ctx, &openfgav1.WriteRequest{
|
||||
Writes: &openfgav1.WriteRequestWrites{
|
||||
TupleKeys: tuples[start:end],
|
||||
},
|
||||
})
|
||||
}); err != nil {
|
||||
if strings.Contains(err.Error(), "cannot write a tuple which already exists") {
|
||||
z.log.Debug("Skipping already synced permissions", "sync_key", key)
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// managedPermissionsCollector collects managed permissions into provided tuple map.
|
||||
// It will only store actions that are supported by our schema. Managed permissions can
|
||||
// be directly mapped to user/team/role without having to write an intermediate role.
|
||||
func managedPermissionsCollector(store db.DB) TupleCollector {
|
||||
return func(ctx context.Context, tuples map[string][]*openfgav1.TupleKey) error {
|
||||
const collectorID = "managed"
|
||||
const query = `
|
||||
SELECT ur.user_id, p.action, p.kind, p.identifier, r.org_id FROM permission p
|
||||
INNER JOIN role r on p.role_id = r.id
|
||||
LEFT JOIN user_role ur on r.id = ur.role_id
|
||||
LEFT JOIN team_role tr on r.id = tr.role_id
|
||||
LEFT JOIN builtin_role br on r.id = br.role_id
|
||||
WHERE r.name LIKE 'managed:%'
|
||||
`
|
||||
type Permission struct {
|
||||
RoleName string `xorm:"role_name"`
|
||||
OrgID int64 `xorm:"org_id"`
|
||||
Action string `xorm:"action"`
|
||||
Kind string
|
||||
Identifier string
|
||||
UserID int64 `xorm:"user_id"`
|
||||
TeamID int64 `xorm:"user_id"`
|
||||
}
|
||||
|
||||
var permissions []Permission
|
||||
err := store.WithDbSession(ctx, func(sess *db.Session) error {
|
||||
return sess.SQL(query).Find(&permissions)
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, p := range permissions {
|
||||
var subject string
|
||||
if p.UserID > 0 {
|
||||
subject = zanzana.NewObject(zanzana.TypeUser, strconv.FormatInt(p.UserID, 10))
|
||||
} else if p.TeamID > 0 {
|
||||
subject = zanzana.NewObject(zanzana.TypeTeam, strconv.FormatInt(p.TeamID, 10))
|
||||
} else {
|
||||
// FIXME(kalleep): Unsuported role binding (org role). We need to have basic roles in place
|
||||
continue
|
||||
}
|
||||
|
||||
tuple, ok := zanzana.TranslateToTuple(subject, p.Action, p.Kind, p.Identifier, p.OrgID)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
// our "sync key" is a combination of collectorID and action so we can run this
|
||||
// sync new data when more actions are supported
|
||||
key := fmt.Sprintf("%s-%s", collectorID, p.Action)
|
||||
tuples[key] = append(tuples[key], tuple)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -136,6 +136,16 @@ func (z *Zanzana) start(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (z *Zanzana) running(ctx context.Context) error {
|
||||
if z.cfg.Env == setting.Dev && z.cfg.Zanzana.ListenHTTP {
|
||||
go func() {
|
||||
z.logger.Info("Starting OpenFGA HTTP server")
|
||||
err := zanzana.StartOpenFGAHttpSever(z.cfg, z.handle, z.logger)
|
||||
if err != nil {
|
||||
z.logger.Error("failed to start OpenFGA HTTP server", "error", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// Run is blocking so we can just run it here
|
||||
return z.handle.Run(ctx)
|
||||
}
|
||||
|
||||
@@ -15,8 +15,9 @@ import (
|
||||
|
||||
// Client is a wrapper around [openfgav1.OpenFGAServiceClient]
|
||||
type Client interface {
|
||||
Check(ctx context.Context, in *openfgav1.CheckRequest, opts ...grpc.CallOption) (*openfgav1.CheckResponse, error)
|
||||
ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest, opts ...grpc.CallOption) (*openfgav1.ListObjectsResponse, error)
|
||||
Check(ctx context.Context, in *openfgav1.CheckRequest) (*openfgav1.CheckResponse, error)
|
||||
ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest) (*openfgav1.ListObjectsResponse, error)
|
||||
Write(ctx context.Context, in *openfgav1.WriteRequest) error
|
||||
}
|
||||
|
||||
func NewClient(ctx context.Context, cc grpc.ClientConnInterface, cfg *setting.Cfg) (*client.Client, error) {
|
||||
@@ -27,3 +28,7 @@ func NewClient(ctx context.Context, cc grpc.ClientConnInterface, cfg *setting.Cf
|
||||
client.WithLogger(log.New("zanzana-client")),
|
||||
)
|
||||
}
|
||||
|
||||
func NewNoopClient() *client.NoopClient {
|
||||
return client.NewNoop()
|
||||
}
|
||||
|
||||
@@ -70,12 +70,23 @@ func New(ctx context.Context, cc grpc.ClientConnInterface, opts ...ClientOption)
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func (c *Client) Check(ctx context.Context, in *openfgav1.CheckRequest, opts ...grpc.CallOption) (*openfgav1.CheckResponse, error) {
|
||||
return c.client.Check(ctx, in, opts...)
|
||||
func (c *Client) Check(ctx context.Context, in *openfgav1.CheckRequest) (*openfgav1.CheckResponse, error) {
|
||||
in.StoreId = c.storeID
|
||||
in.AuthorizationModelId = c.modelID
|
||||
return c.client.Check(ctx, in)
|
||||
}
|
||||
|
||||
func (c *Client) ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest, opts ...grpc.CallOption) (*openfgav1.ListObjectsResponse, error) {
|
||||
return c.client.ListObjects(ctx, in, opts...)
|
||||
func (c *Client) ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest) (*openfgav1.ListObjectsResponse, error) {
|
||||
in.StoreId = c.storeID
|
||||
in.AuthorizationModelId = c.modelID
|
||||
return c.client.ListObjects(ctx, in)
|
||||
}
|
||||
|
||||
func (c *Client) Write(ctx context.Context, in *openfgav1.WriteRequest) error {
|
||||
in.StoreId = c.storeID
|
||||
in.AuthorizationModelId = c.modelID
|
||||
_, err := c.client.Write(ctx, in)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) getOrCreateStore(ctx context.Context, name string) (*openfgav1.Store, error) {
|
||||
|
||||
@@ -3,8 +3,6 @@ package client
|
||||
import (
|
||||
"context"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
|
||||
openfgav1 "github.com/openfga/api/proto/openfga/v1"
|
||||
)
|
||||
|
||||
@@ -14,10 +12,14 @@ func NewNoop() *NoopClient {
|
||||
|
||||
type NoopClient struct{}
|
||||
|
||||
func (nc NoopClient) Check(ctx context.Context, in *openfgav1.CheckRequest, opts ...grpc.CallOption) (*openfgav1.CheckResponse, error) {
|
||||
func (nc NoopClient) Check(ctx context.Context, in *openfgav1.CheckRequest) (*openfgav1.CheckResponse, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (nc NoopClient) ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest, opts ...grpc.CallOption) (*openfgav1.ListObjectsResponse, error) {
|
||||
func (nc NoopClient) ListObjects(ctx context.Context, in *openfgav1.ListObjectsRequest) (*openfgav1.ListObjectsResponse, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (nc NoopClient) Write(ctx context.Context, in *openfgav1.WriteRequest) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,10 +1,27 @@
|
||||
package zanzana
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
openfgav1 "github.com/openfga/api/proto/openfga/v1"
|
||||
httpmiddleware "github.com/openfga/openfga/pkg/middleware/http"
|
||||
"github.com/openfga/openfga/pkg/server"
|
||||
serverErrors "github.com/openfga/openfga/pkg/server/errors"
|
||||
"github.com/openfga/openfga/pkg/storage"
|
||||
"github.com/rs/cors"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
healthv1pb "google.golang.org/grpc/health/grpc_health_v1"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/grafana/grafana/pkg/infra/log"
|
||||
"github.com/grafana/grafana/pkg/services/grpcserver"
|
||||
"github.com/grafana/grafana/pkg/setting"
|
||||
)
|
||||
|
||||
func NewServer(store storage.OpenFGADatastore, logger log.Logger) (*server.Server, error) {
|
||||
@@ -24,3 +41,70 @@ func NewServer(store storage.OpenFGADatastore, logger log.Logger) (*server.Serve
|
||||
|
||||
return srv, nil
|
||||
}
|
||||
|
||||
// StartOpenFGAHttpSever starts HTTP server which allows to use fga cli.
|
||||
func StartOpenFGAHttpSever(cfg *setting.Cfg, srv grpcserver.Provider, logger log.Logger) error {
|
||||
dialOpts := []grpc.DialOption{
|
||||
grpc.WithTransportCredentials(insecure.NewCredentials()),
|
||||
}
|
||||
|
||||
addr := srv.GetAddress()
|
||||
// Wait until GRPC server is initialized
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
maxRetries := 100
|
||||
retries := 0
|
||||
for addr == "" && retries < maxRetries {
|
||||
<-ticker.C
|
||||
addr = srv.GetAddress()
|
||||
retries++
|
||||
}
|
||||
if addr == "" {
|
||||
return fmt.Errorf("failed to start HTTP server: GRPC server unavailable")
|
||||
}
|
||||
|
||||
conn, err := grpc.NewClient(addr, dialOpts...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("unable to dial GRPC: %w", err)
|
||||
}
|
||||
|
||||
muxOpts := []runtime.ServeMuxOption{
|
||||
runtime.WithForwardResponseOption(httpmiddleware.HTTPResponseModifier),
|
||||
runtime.WithErrorHandler(func(c context.Context,
|
||||
sr *runtime.ServeMux, mm runtime.Marshaler, w http.ResponseWriter, r *http.Request, e error) {
|
||||
intCode := serverErrors.ConvertToEncodedErrorCode(status.Convert(e))
|
||||
httpmiddleware.CustomHTTPErrorHandler(c, w, r, serverErrors.NewEncodedError(intCode, e.Error()))
|
||||
}),
|
||||
runtime.WithStreamErrorHandler(func(ctx context.Context, e error) *status.Status {
|
||||
intCode := serverErrors.ConvertToEncodedErrorCode(status.Convert(e))
|
||||
encodedErr := serverErrors.NewEncodedError(intCode, e.Error())
|
||||
return status.Convert(encodedErr)
|
||||
}),
|
||||
runtime.WithHealthzEndpoint(healthv1pb.NewHealthClient(conn)),
|
||||
runtime.WithOutgoingHeaderMatcher(func(s string) (string, bool) { return s, true }),
|
||||
}
|
||||
mux := runtime.NewServeMux(muxOpts...)
|
||||
if err := openfgav1.RegisterOpenFGAServiceHandler(context.TODO(), mux, conn); err != nil {
|
||||
return fmt.Errorf("failed to register gateway handler: %w", err)
|
||||
}
|
||||
|
||||
httpServer := &http.Server{
|
||||
Addr: cfg.Zanzana.HttpAddr,
|
||||
Handler: cors.New(cors.Options{
|
||||
AllowedOrigins: []string{"*"},
|
||||
AllowCredentials: true,
|
||||
AllowedHeaders: []string{"*"},
|
||||
AllowedMethods: []string{http.MethodGet, http.MethodPost,
|
||||
http.MethodHead, http.MethodPatch, http.MethodDelete, http.MethodPut},
|
||||
}).Handler(mux),
|
||||
ReadHeaderTimeout: 30 * time.Second,
|
||||
}
|
||||
go func() {
|
||||
err = httpServer.ListenAndServe()
|
||||
if err != nil {
|
||||
logger.Error("failed to start http server", zapcore.Field{Key: "err", Type: zapcore.ErrorType, Interface: err})
|
||||
}
|
||||
}()
|
||||
logger.Info(fmt.Sprintf("OpenFGA HTTP server listening on '%s'...", httpServer.Addr))
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
package zanzana
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
|
||||
openfgav1 "github.com/openfga/api/proto/openfga/v1"
|
||||
)
|
||||
|
||||
const (
|
||||
TypeUser string = "user"
|
||||
TypeTeam string = "team"
|
||||
)
|
||||
|
||||
func NewObject(typ, id string) string {
|
||||
return fmt.Sprintf("%s:%s", typ, id)
|
||||
}
|
||||
|
||||
func NewScopedObject(typ, id, scope string) string {
|
||||
return NewObject(typ, fmt.Sprintf("%s-%s", scope, id))
|
||||
}
|
||||
|
||||
// rbac action to relation translation
|
||||
var actionTranslations = map[string]string{}
|
||||
|
||||
type kindTranslation struct {
|
||||
typ string
|
||||
orgScoped bool
|
||||
}
|
||||
|
||||
// all kinds that we can translate into a openFGA object
|
||||
var kindTranslations = map[string]kindTranslation{}
|
||||
|
||||
func TranslateToTuple(user string, action, kind, identifier string, orgID int64) (*openfgav1.TupleKey, bool) {
|
||||
relation, ok := actionTranslations[action]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
t, ok := kindTranslations[kind]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
tuple := &openfgav1.TupleKey{
|
||||
Relation: relation,
|
||||
}
|
||||
|
||||
tuple.User = user
|
||||
tuple.Relation = relation
|
||||
|
||||
// UID in grafana are not guarantee to be unique across orgs so we need to scope them.
|
||||
if t.orgScoped {
|
||||
tuple.Object = NewScopedObject(t.typ, identifier, strconv.FormatInt(orgID, 10))
|
||||
} else {
|
||||
tuple.Object = NewObject(t.typ, identifier)
|
||||
}
|
||||
|
||||
return tuple, true
|
||||
}
|
||||
@@ -112,16 +112,16 @@ func (h *ContextHandler) Middleware(next http.Handler) http.Handler {
|
||||
reqContext.Logger = reqContext.Logger.New("traceID", traceID)
|
||||
}
|
||||
|
||||
identity, err := h.authnService.Authenticate(ctx, &authn.Request{HTTPRequest: reqContext.Req, Resp: reqContext.Resp})
|
||||
id, err := h.authnService.Authenticate(ctx, &authn.Request{HTTPRequest: reqContext.Req, Resp: reqContext.Resp})
|
||||
if err != nil {
|
||||
// Hack: set all errors on LookupTokenErr, so we can check it in auth middlewares
|
||||
reqContext.LookupTokenErr = err
|
||||
} else {
|
||||
reqContext.SignedInUser = identity.SignedInUser()
|
||||
reqContext.UserToken = identity.SessionToken
|
||||
reqContext.SignedInUser = id.SignedInUser()
|
||||
reqContext.UserToken = id.SessionToken
|
||||
reqContext.IsSignedIn = !reqContext.SignedInUser.IsAnonymous
|
||||
reqContext.AllowAnonymous = reqContext.SignedInUser.IsAnonymous
|
||||
reqContext.IsRenderCall = identity.IsAuthenticatedBy(login.RenderModule)
|
||||
reqContext.IsRenderCall = id.IsAuthenticatedBy(login.RenderModule)
|
||||
}
|
||||
|
||||
reqContext.Logger = reqContext.Logger.New("userId", reqContext.UserID, "orgId", reqContext.OrgID, "uname", reqContext.Login)
|
||||
@@ -138,7 +138,7 @@ func (h *ContextHandler) Middleware(next http.Handler) http.Handler {
|
||||
// End the span to make next handlers not wrapped within middleware span
|
||||
span.End()
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
next.ServeHTTP(w, r.WithContext(identity.WithRequester(ctx, id)))
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/grafana/grafana/pkg/api/routing"
|
||||
"github.com/grafana/grafana/pkg/apimachinery/identity"
|
||||
"github.com/grafana/grafana/pkg/infra/tracing"
|
||||
"github.com/grafana/grafana/pkg/services/authn"
|
||||
"github.com/grafana/grafana/pkg/services/authn/authntest"
|
||||
@@ -44,20 +45,24 @@ func TestContextHandler(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("should set identity on successful authentication", func(t *testing.T) {
|
||||
identity := &authn.Identity{ID: authn.NewNamespaceID(authn.NamespaceUser, 1), OrgID: 1}
|
||||
id := &authn.Identity{ID: authn.NewNamespaceID(authn.NamespaceUser, 1), OrgID: 1}
|
||||
handler := contexthandler.ProvideService(
|
||||
setting.NewCfg(),
|
||||
tracing.InitializeTracerForTest(),
|
||||
featuremgmt.WithFeatures(),
|
||||
&authntest.FakeService{ExpectedIdentity: identity},
|
||||
&authntest.FakeService{ExpectedIdentity: id},
|
||||
)
|
||||
|
||||
server := webtest.NewServer(t, routing.NewRouteRegister())
|
||||
server.Mux.Use(handler.Middleware)
|
||||
server.Mux.Get("/api/handler", func(c *contextmodel.ReqContext) {
|
||||
require.True(t, c.IsSignedIn)
|
||||
require.EqualValues(t, identity.SignedInUser(), c.SignedInUser)
|
||||
require.EqualValues(t, id.SignedInUser(), c.SignedInUser)
|
||||
require.NoError(t, c.LookupTokenErr)
|
||||
|
||||
requester, err := identity.GetRequester(c.Req.Context())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, id, requester)
|
||||
})
|
||||
|
||||
res, err := server.Send(server.NewGetRequest("/api/handler"))
|
||||
|
||||
@@ -1371,6 +1371,13 @@ var (
|
||||
HideFromDocs: true,
|
||||
HideFromAdminPage: true,
|
||||
},
|
||||
{
|
||||
Name: "cloudwatchMetricInsightsCrossAccount",
|
||||
Description: "Enables cross account observability for Cloudwatch Metric Insights",
|
||||
Stage: FeatureStageExperimental,
|
||||
Owner: awsDatasourcesSquad,
|
||||
FrontendOnly: true,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -181,3 +181,4 @@ alertingApiServer,experimental,@grafana/alerting-squad,false,true,false
|
||||
dashboardRestoreUI,experimental,@grafana/grafana-frontend-platform,false,false,false
|
||||
cloudWatchRoundUpEndTime,GA,@grafana/aws-datasources,false,false,false
|
||||
bodyScrolling,experimental,@grafana/grafana-frontend-platform,false,false,true
|
||||
cloudwatchMetricInsightsCrossAccount,experimental,@grafana/aws-datasources,false,false,true
|
||||
|
||||
|
@@ -734,4 +734,8 @@ const (
|
||||
// FlagBodyScrolling
|
||||
// Adjusts Page to make body the scrollable element
|
||||
FlagBodyScrolling = "bodyScrolling"
|
||||
|
||||
// FlagCloudwatchMetricInsightsCrossAccount
|
||||
// Enables cross account observability for Cloudwatch Metric Insights
|
||||
FlagCloudwatchMetricInsightsCrossAccount = "cloudwatchMetricInsightsCrossAccount"
|
||||
)
|
||||
|
||||
@@ -590,6 +590,19 @@
|
||||
"codeowner": "@grafana/aws-datasources"
|
||||
}
|
||||
},
|
||||
{
|
||||
"metadata": {
|
||||
"name": "cloudwatchMetricInsightsCrossAccount",
|
||||
"resourceVersion": "1719497905377",
|
||||
"creationTimestamp": "2024-06-27T14:18:25Z"
|
||||
},
|
||||
"spec": {
|
||||
"description": "Enables cross account observability for Cloudwatch Metric Insights",
|
||||
"stage": "experimental",
|
||||
"codeowner": "@grafana/aws-datasources",
|
||||
"frontend": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"metadata": {
|
||||
"name": "configurableSchedulerTick",
|
||||
|
||||
@@ -45,7 +45,10 @@ func setupTestEnv(t *testing.T) *TestEnv {
|
||||
}
|
||||
logger := log.New("extsvcaccounts.test")
|
||||
env.S = &ExtSvcAccountsService{
|
||||
acSvc: acimpl.ProvideOSSService(cfg, env.AcStore, &resourcepermissions.FakeActionSetSvc{}, localcache.New(0, 0), fmgt, tracing.InitializeTracerForTest()),
|
||||
acSvc: acimpl.ProvideOSSService(
|
||||
cfg, env.AcStore, &resourcepermissions.FakeActionSetSvc{},
|
||||
localcache.New(0, 0), fmgt, tracing.InitializeTracerForTest(), nil, nil,
|
||||
),
|
||||
features: fmgt,
|
||||
logger: logger,
|
||||
metrics: newMetrics(nil, env.SaSvc, logger),
|
||||
|
||||
@@ -62,6 +62,7 @@ func ProvideService(cfg *setting.Cfg, sqlStore db.DB, ac ac.AccessControl,
|
||||
|
||||
if features.IsEnabledGlobally(featuremgmt.FlagSsoSettingsLDAP) {
|
||||
providersList = append(providersList, social.LDAPProviderName)
|
||||
configurableProviders[social.LDAPProviderName] = true
|
||||
}
|
||||
|
||||
if licensing.FeatureEnabled(social.SAMLProviderName) {
|
||||
@@ -320,21 +321,23 @@ func (s *Service) getFallbackStrategyFor(provider string) (ssosettings.FallbackS
|
||||
}
|
||||
|
||||
func (s *Service) encryptSecrets(ctx context.Context, settings map[string]any) (map[string]any, error) {
|
||||
result := make(map[string]any)
|
||||
for k, v := range settings {
|
||||
if IsSecretField(k) && v != "" {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
return result, fmt.Errorf("failed to encrypt %s setting because it is not a string: %v", k, v)
|
||||
}
|
||||
result := deepCopyMap(settings)
|
||||
configs := getConfigMaps(result)
|
||||
|
||||
encryptedSecret, err := s.secrets.Encrypt(ctx, []byte(strValue), secrets.WithoutScope())
|
||||
if err != nil {
|
||||
return result, err
|
||||
for _, config := range configs {
|
||||
for k, v := range config {
|
||||
if IsSecretField(k) && v != "" {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
return result, fmt.Errorf("failed to encrypt %s setting because it is not a string: %v", k, v)
|
||||
}
|
||||
|
||||
encryptedSecret, err := s.secrets.Encrypt(ctx, []byte(strValue), secrets.WithoutScope())
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
config[k] = base64.RawStdEncoding.EncodeToString(encryptedSecret)
|
||||
}
|
||||
result[k] = base64.RawStdEncoding.EncodeToString(encryptedSecret)
|
||||
} else {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
@@ -411,29 +414,34 @@ func (s *Service) mergeSSOSettings(dbSettings, systemSettings *models.SSOSetting
|
||||
}
|
||||
|
||||
func (s *Service) decryptSecrets(ctx context.Context, settings map[string]any) (map[string]any, error) {
|
||||
for k, v := range settings {
|
||||
if IsSecretField(k) && v != "" {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
s.logger.Error("Failed to parse secret value, it is not a string", "key", k)
|
||||
return nil, fmt.Errorf("secret value is not a string")
|
||||
}
|
||||
configs := getConfigMaps(settings)
|
||||
|
||||
decoded, err := base64.RawStdEncoding.DecodeString(strValue)
|
||||
if err != nil {
|
||||
s.logger.Error("Failed to decode secret string", "err", err, "value")
|
||||
return nil, err
|
||||
}
|
||||
for _, config := range configs {
|
||||
for k, v := range config {
|
||||
if IsSecretField(k) && v != "" {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
s.logger.Error("Failed to parse secret value, it is not a string", "key", k)
|
||||
return nil, fmt.Errorf("secret value is not a string")
|
||||
}
|
||||
|
||||
decrypted, err := s.secrets.Decrypt(ctx, decoded)
|
||||
if err != nil {
|
||||
s.logger.Error("Failed to decrypt secret", "err", err)
|
||||
return nil, err
|
||||
}
|
||||
decoded, err := base64.RawStdEncoding.DecodeString(strValue)
|
||||
if err != nil {
|
||||
s.logger.Error("Failed to decode secret string", "err", err, "value")
|
||||
return nil, err
|
||||
}
|
||||
|
||||
settings[k] = string(decrypted)
|
||||
decrypted, err := s.secrets.Decrypt(ctx, decoded)
|
||||
if err != nil {
|
||||
s.logger.Error("Failed to decrypt secret", "err", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
config[k] = string(decrypted)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
@@ -445,18 +453,39 @@ func (s *Service) isProviderConfigurable(provider string) bool {
|
||||
// removeSecrets removes all the secrets from the map and replaces them with a redacted password
|
||||
// and returns a new map
|
||||
func removeSecrets(settings map[string]any) map[string]any {
|
||||
result := make(map[string]any)
|
||||
for k, v := range settings {
|
||||
val, ok := v.(string)
|
||||
if ok && val != "" && IsSecretField(k) {
|
||||
result[k] = setting.RedactedPassword
|
||||
continue
|
||||
result := deepCopyMap(settings)
|
||||
configs := getConfigMaps(result)
|
||||
|
||||
for _, config := range configs {
|
||||
for k, v := range config {
|
||||
val, ok := v.(string)
|
||||
if ok && val != "" && IsSecretField(k) {
|
||||
config[k] = setting.RedactedPassword
|
||||
}
|
||||
}
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// getConfigMaps returns a list of maps that may contain secrets
|
||||
func getConfigMaps(settings map[string]any) []map[string]any {
|
||||
// always include the main settings map
|
||||
result := []map[string]any{settings}
|
||||
|
||||
// for LDAP include settings for each server
|
||||
if config, ok := settings["config"].(map[string]any); ok {
|
||||
if servers, ok := config["servers"].([]any); ok {
|
||||
for _, server := range servers {
|
||||
if serverSettings, ok := server.(map[string]any); ok {
|
||||
result = append(result, serverSettings)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// mergeSettings merges two maps in a way that the values from the first map are preserved
|
||||
// and the values from the second map are added only if they don't exist in the first map
|
||||
// or if they contain empty URLs.
|
||||
@@ -500,23 +529,25 @@ func isMergingAllowed(fieldName string) bool {
|
||||
|
||||
// mergeSecrets returns a new map with the current value for secrets that have not been updated
|
||||
func mergeSecrets(settings map[string]any, storedSettings map[string]any) (map[string]any, error) {
|
||||
settingsWithSecrets := map[string]any{}
|
||||
for k, v := range settings {
|
||||
if IsSecretField(k) {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("secret value is not a string")
|
||||
}
|
||||
settingsWithSecrets := deepCopyMap(settings)
|
||||
newConfigs := getConfigMaps(settingsWithSecrets)
|
||||
storedConfigs := getConfigMaps(storedSettings)
|
||||
|
||||
if isNewSecretValue(strValue) {
|
||||
settingsWithSecrets[k] = strValue // use the new value
|
||||
continue
|
||||
for i, config := range newConfigs {
|
||||
for k, v := range config {
|
||||
if IsSecretField(k) {
|
||||
strValue, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("secret value is not a string")
|
||||
}
|
||||
|
||||
if !isNewSecretValue(strValue) && len(storedConfigs) > i {
|
||||
config[k] = storedConfigs[i][k] // use the currently stored value
|
||||
}
|
||||
}
|
||||
settingsWithSecrets[k] = storedSettings[k] // keep the currently stored value
|
||||
} else {
|
||||
settingsWithSecrets[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
return settingsWithSecrets, nil
|
||||
}
|
||||
|
||||
@@ -532,7 +563,7 @@ func overrideMaps(maps ...map[string]any) map[string]any {
|
||||
|
||||
// IsSecretField returns true if the SSO settings field provided is a secret
|
||||
func IsSecretField(fieldName string) bool {
|
||||
secretFieldPatterns := []string{"secret", "private", "certificate"}
|
||||
secretFieldPatterns := []string{"secret", "private", "certificate", "password", "client_key"}
|
||||
|
||||
for _, v := range secretFieldPatterns {
|
||||
if strings.Contains(strings.ToLower(fieldName), strings.ToLower(v)) {
|
||||
@@ -554,3 +585,37 @@ func isEmptyString(val any) bool {
|
||||
func isNewSecretValue(value string) bool {
|
||||
return value != setting.RedactedPassword
|
||||
}
|
||||
|
||||
func deepCopyMap(settings map[string]any) map[string]any {
|
||||
newSettings := make(map[string]any)
|
||||
|
||||
for key, value := range settings {
|
||||
switch v := value.(type) {
|
||||
case map[string]any:
|
||||
newSettings[key] = deepCopyMap(v)
|
||||
case []any:
|
||||
newSettings[key] = deepCopySlice(v)
|
||||
default:
|
||||
newSettings[key] = value
|
||||
}
|
||||
}
|
||||
|
||||
return newSettings
|
||||
}
|
||||
|
||||
func deepCopySlice(s []any) []any {
|
||||
newSlice := make([]any, len(s))
|
||||
|
||||
for i, value := range s {
|
||||
switch v := value.(type) {
|
||||
case map[string]any:
|
||||
newSlice[i] = deepCopyMap(v)
|
||||
case []any:
|
||||
newSlice[i] = deepCopySlice(v)
|
||||
default:
|
||||
newSlice[i] = value
|
||||
}
|
||||
}
|
||||
|
||||
return newSlice
|
||||
}
|
||||
|
||||
@@ -158,6 +158,62 @@ func TestService_GetForProvider(t *testing.T) {
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should decrypt secrets for LDAP if data is coming from store",
|
||||
provider: "ldap",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedSSOSetting = &models.SSOSettings{
|
||||
Provider: "ldap",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": base64.RawStdEncoding.EncodeToString([]byte("bind_password_1")),
|
||||
"client_key": base64.RawStdEncoding.EncodeToString([]byte("client_key_1")),
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": base64.RawStdEncoding.EncodeToString([]byte("bind_password_2")),
|
||||
"client_key": base64.RawStdEncoding.EncodeToString([]byte("client_key_2")),
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Source: models.DB,
|
||||
}
|
||||
env.fallbackStrategy.ExpectedIsMatch = true
|
||||
env.fallbackStrategy.ExpectedConfigs = map[string]map[string]any{}
|
||||
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("bind_password_1"), mock.Anything).Return([]byte("decrypted-bind-password-1"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_key_1"), mock.Anything).Return([]byte("decrypted-client-key-1"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("bind_password_2"), mock.Anything).Return([]byte("decrypted-bind-password-2"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_key_2"), mock.Anything).Return([]byte("decrypted-client-key-2"), nil).Once()
|
||||
},
|
||||
want: &models.SSOSettings{
|
||||
Provider: "ldap",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": "decrypted-bind-password-1",
|
||||
"client_key": "decrypted-client-key-1",
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": "decrypted-bind-password-2",
|
||||
"client_key": "decrypted-client-key-2",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Source: models.DB,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should not decrypt secrets if data is coming from the fallback strategy",
|
||||
provider: "github",
|
||||
@@ -290,7 +346,7 @@ func TestService_GetForProvider(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, true, false, true)
|
||||
env := setupTestEnv(t, true, false, true, true)
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
}
|
||||
@@ -314,13 +370,15 @@ func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
testCases := []struct {
|
||||
name string
|
||||
setup func(env testEnv)
|
||||
want *models.SSOSettings
|
||||
wantErr bool
|
||||
name string
|
||||
provider string
|
||||
setup func(env testEnv)
|
||||
want *models.SSOSettings
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "should return successfully and redact secrets",
|
||||
name: "should return successfully and redact secrets",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedSSOSetting = &models.SSOSettings{
|
||||
Provider: "github",
|
||||
@@ -347,13 +405,67 @@ func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should return error if store returns an error different than not found",
|
||||
setup: func(env testEnv) { env.store.ExpectedError = fmt.Errorf("error") },
|
||||
want: nil,
|
||||
wantErr: true,
|
||||
name: "should return successfully and redact secrets for LDAP",
|
||||
provider: "ldap",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedSSOSetting = &models.SSOSettings{
|
||||
Provider: "ldap",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": base64.RawStdEncoding.EncodeToString([]byte("bind_password_1")),
|
||||
"client_key": base64.RawStdEncoding.EncodeToString([]byte("client_key_1")),
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": base64.RawStdEncoding.EncodeToString([]byte("bind_password_2")),
|
||||
"client_key": base64.RawStdEncoding.EncodeToString([]byte("client_key_2")),
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Source: models.DB,
|
||||
}
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("bind_password_1"), mock.Anything).Return([]byte("decrypted-bind-password-1"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_key_1"), mock.Anything).Return([]byte("decrypted-client-key-1"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("bind_password_2"), mock.Anything).Return([]byte("decrypted-bind-password-2"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_key_2"), mock.Anything).Return([]byte("decrypted-client-key-2"), nil).Once()
|
||||
},
|
||||
want: &models.SSOSettings{
|
||||
Provider: "ldap",
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": "*********",
|
||||
"client_key": "*********",
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": "*********",
|
||||
"client_key": "*********",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should fallback to strategy if store returns not found",
|
||||
name: "should return error if store returns an error different than not found",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) { env.store.ExpectedError = fmt.Errorf("error") },
|
||||
want: nil,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "should fallback to strategy if store returns not found",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedError = ssosettings.ErrNotFound
|
||||
env.fallbackStrategy.ExpectedIsMatch = true
|
||||
@@ -371,7 +483,8 @@ func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "should return error if the fallback strategy was not found",
|
||||
name: "should return error if the fallback strategy was not found",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedError = ssosettings.ErrNotFound
|
||||
env.fallbackStrategy.ExpectedIsMatch = false
|
||||
@@ -380,7 +493,8 @@ func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "should return error if fallback strategy returns error",
|
||||
name: "should return error if fallback strategy returns error",
|
||||
provider: "github",
|
||||
setup: func(env testEnv) {
|
||||
env.store.ExpectedError = ssosettings.ErrNotFound
|
||||
env.fallbackStrategy.ExpectedIsMatch = true
|
||||
@@ -399,12 +513,12 @@ func TestService_GetForProviderWithRedactedSecrets(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, true)
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
}
|
||||
|
||||
actual, err := env.service.GetForProviderWithRedactedSecrets(context.Background(), "github")
|
||||
actual, err := env.service.GetForProviderWithRedactedSecrets(context.Background(), tc.provider)
|
||||
|
||||
if tc.wantErr {
|
||||
require.Error(t, err)
|
||||
@@ -550,7 +664,7 @@ func TestService_List(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
}
|
||||
@@ -852,7 +966,7 @@ func TestService_ListWithRedactedSecrets(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
}
|
||||
@@ -876,7 +990,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("successfully upsert SSO settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -936,10 +1050,80 @@ func TestService_Upsert(t *testing.T) {
|
||||
require.EqualValues(t, settings, env.store.ActualSSOSettings)
|
||||
})
|
||||
|
||||
t.Run("successfully upsert SSO settings for LDAP", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false, true)
|
||||
|
||||
provider := social.LDAPProviderName
|
||||
settings := models.SSOSettings{
|
||||
Provider: provider,
|
||||
Settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": "bind_password_1",
|
||||
"client_key": "client_key_1",
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": "bind_password_2",
|
||||
"client_key": "client_key_2",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
reloadable := ssosettingstests.NewMockReloadable(t)
|
||||
reloadable.On("Validate", mock.Anything, settings, mock.Anything, mock.Anything).Return(nil)
|
||||
reloadable.On("Reload", mock.Anything, mock.MatchedBy(func(settings models.SSOSettings) bool {
|
||||
defer wg.Done()
|
||||
return settings.Provider == provider &&
|
||||
settings.ID == "someid" &&
|
||||
maps.Equal(settings.Settings["config"].(map[string]any)["servers"].([]any)[0].(map[string]any), map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": "bind_password_1",
|
||||
"client_key": "client_key_1",
|
||||
}) && maps.Equal(settings.Settings["config"].(map[string]any)["servers"].([]any)[1].(map[string]any), map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": "bind_password_2",
|
||||
"client_key": "client_key_2",
|
||||
})
|
||||
})).Return(nil).Once()
|
||||
env.reloadables[provider] = reloadable
|
||||
env.secrets.On("Encrypt", mock.Anything, []byte("bind_password_1"), mock.Anything).Return([]byte("encrypted-bind-password-1"), nil).Once()
|
||||
env.secrets.On("Encrypt", mock.Anything, []byte("bind_password_2"), mock.Anything).Return([]byte("encrypted-bind-password-2"), nil).Once()
|
||||
env.secrets.On("Encrypt", mock.Anything, []byte("client_key_1"), mock.Anything).Return([]byte("encrypted-client-key-1"), nil).Once()
|
||||
env.secrets.On("Encrypt", mock.Anything, []byte("client_key_2"), mock.Anything).Return([]byte("encrypted-client-key-2"), nil).Once()
|
||||
|
||||
env.store.UpsertFn = func(ctx context.Context, settings *models.SSOSettings) error {
|
||||
currentTime := time.Now()
|
||||
settings.ID = "someid"
|
||||
settings.Created = currentTime
|
||||
settings.Updated = currentTime
|
||||
|
||||
env.store.ActualSSOSettings = *settings
|
||||
return nil
|
||||
}
|
||||
|
||||
err := env.service.Upsert(context.Background(), &settings, &user.SignedInUser{})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for the goroutine first to assert the Reload call
|
||||
wg.Wait()
|
||||
|
||||
require.EqualValues(t, settings, env.store.ActualSSOSettings)
|
||||
})
|
||||
|
||||
t.Run("returns error if provider is not configurable", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.GrafanaComProviderName
|
||||
settings := &models.SSOSettings{
|
||||
@@ -962,7 +1146,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if provider was not found in reloadables", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := &models.SSOSettings{
|
||||
@@ -986,7 +1170,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if validation fails", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1010,7 +1194,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if a fallback strategy is not available for the provider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
settings := &models.SSOSettings{
|
||||
Provider: social.AzureADProviderName,
|
||||
@@ -1031,7 +1215,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if a secret does not have the type string", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.OktaProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1054,7 +1238,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if secrets encryption failed", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.OktaProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1079,7 +1263,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("should not update the current secret if the secret has not been updated", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1123,7 +1307,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("run validation with all new and current secrets available in settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1176,7 +1360,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("returns error if store failed to upsert settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1208,7 +1392,7 @@ func TestService_Upsert(t *testing.T) {
|
||||
t.Run("successfully upsert SSO settings if reload fails", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
settings := models.SSOSettings{
|
||||
@@ -1241,7 +1425,7 @@ func TestService_Delete(t *testing.T) {
|
||||
t.Run("successfully delete SSO settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
@@ -1279,7 +1463,7 @@ func TestService_Delete(t *testing.T) {
|
||||
t.Run("return error if SSO setting was not found for the specified provider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
reloadable := ssosettingstests.NewMockReloadable(t)
|
||||
@@ -1295,7 +1479,7 @@ func TestService_Delete(t *testing.T) {
|
||||
t.Run("should not delete the SSO settings if the provider is not configurable", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
env.cfg.SSOSettingsConfigurableProviders = map[string]bool{social.AzureADProviderName: true}
|
||||
|
||||
provider := social.GrafanaComProviderName
|
||||
@@ -1308,7 +1492,7 @@ func TestService_Delete(t *testing.T) {
|
||||
t.Run("return error when store fails to delete the SSO settings for the specified provider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
env.store.ExpectedError = errors.New("delete sso settings failed")
|
||||
@@ -1321,7 +1505,7 @@ func TestService_Delete(t *testing.T) {
|
||||
t.Run("return successfully when the deletion was successful but reloading the settings fail", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := social.AzureADProviderName
|
||||
reloadable := ssosettingstests.NewMockReloadable(t)
|
||||
@@ -1337,13 +1521,51 @@ func TestService_Delete(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// we might not need this test because it is not testing the public interface
|
||||
// it was added for convenient testing of the internal deep copy and remove secrets
|
||||
func TestRemoveSecrets(t *testing.T) {
|
||||
settings := map[string]any{
|
||||
"enabled": true,
|
||||
"client_secret": "client_secret",
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"host": "192.168.0.1",
|
||||
"bind_password": "bind_password_1",
|
||||
"client_key": "client_key_1",
|
||||
},
|
||||
map[string]any{
|
||||
"host": "192.168.0.2",
|
||||
"bind_password": "bind_password_2",
|
||||
"client_key": "client_key_2",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
copiedSettings := deepCopyMap(settings)
|
||||
copiedSettings["client_secret"] = "client_secret_updated"
|
||||
copiedSettings["config"].(map[string]any)["servers"].([]any)[0].(map[string]any)["bind_password"] = "bind_password_1_updated"
|
||||
|
||||
require.Equal(t, "client_secret", settings["client_secret"])
|
||||
require.Equal(t, "client_secret_updated", copiedSettings["client_secret"])
|
||||
require.Equal(t, "bind_password_1", settings["config"].(map[string]any)["servers"].([]any)[0].(map[string]any)["bind_password"])
|
||||
require.Equal(t, "bind_password_1_updated", copiedSettings["config"].(map[string]any)["servers"].([]any)[0].(map[string]any)["bind_password"])
|
||||
|
||||
settingsWithRedactedSecrets := removeSecrets(settings)
|
||||
require.Equal(t, "client_secret", settings["client_secret"])
|
||||
require.Equal(t, "*********", settingsWithRedactedSecrets["client_secret"])
|
||||
require.Equal(t, "bind_password_1", settings["config"].(map[string]any)["servers"].([]any)[0].(map[string]any)["bind_password"])
|
||||
require.Equal(t, "*********", settingsWithRedactedSecrets["config"].(map[string]any)["servers"].([]any)[0].(map[string]any)["bind_password"])
|
||||
}
|
||||
|
||||
func TestService_DoReload(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("successfully reload settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
settingsList := []*models.SSOSettings{
|
||||
{
|
||||
@@ -1383,7 +1605,7 @@ func TestService_DoReload(t *testing.T) {
|
||||
t.Run("successfully reload settings when some providers have empty settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
settingsList := []*models.SSOSettings{
|
||||
{
|
||||
@@ -1413,7 +1635,7 @@ func TestService_DoReload(t *testing.T) {
|
||||
t.Run("failed fetching the SSO settings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
provider := "github"
|
||||
|
||||
@@ -1459,6 +1681,35 @@ func TestService_decryptSecrets(t *testing.T) {
|
||||
"certificate": "decrypted-certificate",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "should decrypt LDAP secrets successfully",
|
||||
setup: func(env testEnv) {
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("client_key"), mock.Anything).Return([]byte("decrypted-client-key"), nil).Once()
|
||||
env.secrets.On("Decrypt", mock.Anything, []byte("bind_password"), mock.Anything).Return([]byte("decrypted-bind-password"), nil).Once()
|
||||
},
|
||||
settings: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"client_key": base64.RawStdEncoding.EncodeToString([]byte("client_key")),
|
||||
"bind_password": base64.RawStdEncoding.EncodeToString([]byte("bind_password")),
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
want: map[string]any{
|
||||
"enabled": true,
|
||||
"config": map[string]any{
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"client_key": "decrypted-client-key",
|
||||
"bind_password": "decrypted-bind-password",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "should not decrypt when a secret is empty",
|
||||
setup: func(env testEnv) {
|
||||
@@ -1514,7 +1765,7 @@ func TestService_decryptSecrets(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, false, false, false)
|
||||
env := setupTestEnv(t, false, false, false, false)
|
||||
|
||||
if tc.setup != nil {
|
||||
tc.setup(env)
|
||||
@@ -1593,7 +1844,7 @@ func Test_ProviderService(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := setupTestEnv(t, tc.isLicenseEnabled, true, tc.samlEnabled)
|
||||
env := setupTestEnv(t, tc.isLicenseEnabled, true, tc.samlEnabled, false)
|
||||
|
||||
require.Equal(t, tc.expectedProvidersList, env.service.providersList)
|
||||
require.Equal(t, tc.strategiesLength, len(env.service.fbStrategies))
|
||||
@@ -1601,7 +1852,7 @@ func Test_ProviderService(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func setupTestEnv(t *testing.T, isLicensingEnabled, keepFallbackStratergies, samlEnabled bool) testEnv {
|
||||
func setupTestEnv(t *testing.T, isLicensingEnabled, keepFallbackStratergies, samlEnabled bool, ldapEnabled bool) testEnv {
|
||||
t.Helper()
|
||||
|
||||
store := ssosettingstests.NewFakeStore()
|
||||
@@ -1631,10 +1882,14 @@ func setupTestEnv(t *testing.T, isLicensingEnabled, keepFallbackStratergies, sam
|
||||
licensing := licensingtest.NewFakeLicensing()
|
||||
licensing.On("FeatureEnabled", "saml").Return(isLicensingEnabled)
|
||||
|
||||
featureManager := featuremgmt.WithManager()
|
||||
features := make([]any, 0)
|
||||
if samlEnabled {
|
||||
featureManager = featuremgmt.WithManager(featuremgmt.FlagSsoSettingsSAML)
|
||||
features = append(features, featuremgmt.FlagSsoSettingsSAML)
|
||||
}
|
||||
if ldapEnabled {
|
||||
features = append(features, featuremgmt.FlagSsoSettingsLDAP)
|
||||
}
|
||||
featureManager := featuremgmt.WithManager(features...)
|
||||
|
||||
svc := ProvideService(
|
||||
cfg,
|
||||
|
||||
@@ -16,6 +16,10 @@ type ZanzanaSettings struct {
|
||||
Addr string
|
||||
// Mode can either be embedded or client
|
||||
Mode ZanzanaMode
|
||||
// ListenHTTP enables OpenFGA http server which allows to use fga cli
|
||||
ListenHTTP bool
|
||||
// OpenFGA http server address which allows to connect with fga cli
|
||||
HttpAddr string
|
||||
}
|
||||
|
||||
func (cfg *Cfg) readZanzanaSettings() {
|
||||
@@ -32,6 +36,8 @@ func (cfg *Cfg) readZanzanaSettings() {
|
||||
}
|
||||
|
||||
s.Addr = sec.Key("address").MustString("")
|
||||
s.ListenHTTP = sec.Key("listen_http").MustBool(false)
|
||||
s.HttpAddr = sec.Key("http_addr").MustString("127.0.0.1:8080")
|
||||
|
||||
cfg.Zanzana = s
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user