Zanzana: Use shared auth interceptor for zanzana and pass tracer (#100968)

* Use shared auth interceptor for zanzana and pass tracer
This commit is contained in:
Karl Persson
2025-02-20 16:07:06 +01:00
committed by GitHub
parent 74e621f377
commit 14886410d6
11 changed files with 48 additions and 55 deletions
@@ -12,6 +12,7 @@ import (
"github.com/grafana/authlib/types"
"github.com/prometheus/client_golang/prometheus"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
@@ -22,7 +23,7 @@ import (
)
func NewInProcGrpcAuthenticator() interceptors.Authenticator {
return newAuthenticator(
return NewAuthenticatorInterceptor(
authn.NewDefaultAuthenticator(
authn.NewUnsafeAccessTokenVerifier(authn.VerifierConfig{}),
authn.NewUnsafeIDTokenVerifier(authn.VerifierConfig{}),
@@ -46,7 +47,7 @@ func NewAuthenticator(cfg *GrpcServerConfig, tracer tracing.Tracer) interceptors
authn.NewIDTokenVerifier(authn.VerifierConfig{}, kr),
)
return newAuthenticator(auth, tracer)
return NewAuthenticatorInterceptor(auth, tracer)
}
func NewAuthenticatorWithFallback(cfg *setting.Cfg, reg prometheus.Registerer, tracer tracing.Tracer, fallback interceptors.Authenticator) interceptors.Authenticator {
@@ -64,7 +65,7 @@ func NewAuthenticatorWithFallback(cfg *setting.Cfg, reg prometheus.Registerer, t
}
}
func newAuthenticator(auth authn.Authenticator, tracer tracing.Tracer) interceptors.Authenticator {
func NewAuthenticatorInterceptor(auth authn.Authenticator, tracer trace.Tracer) interceptors.Authenticator {
return interceptors.AuthenticatorFunc(func(ctx context.Context) (context.Context, error) {
ctx, span := tracer.Start(ctx, "grpcutils.Authenticate")
defer span.End()
+26 -34
View File
@@ -12,27 +12,26 @@ import (
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
healthv1pb "google.golang.org/grpc/health/grpc_health_v1"
"google.golang.org/grpc/metadata"
authnlib "github.com/grafana/authlib/authn"
authzv1 "github.com/grafana/authlib/authz/proto/v1"
claims "github.com/grafana/authlib/types"
"github.com/grafana/authlib/types"
"github.com/grafana/dskit/services"
"github.com/grafana/grafana/pkg/infra/db"
"github.com/grafana/grafana/pkg/infra/log"
"github.com/grafana/grafana/pkg/infra/tracing"
"github.com/grafana/grafana/pkg/services/authn/grpcutils"
authzextv1 "github.com/grafana/grafana/pkg/services/authz/proto/v1"
"github.com/grafana/grafana/pkg/services/authz/zanzana"
"github.com/grafana/grafana/pkg/services/featuremgmt"
"github.com/grafana/grafana/pkg/services/grpcserver"
"github.com/grafana/grafana/pkg/services/grpcserver/interceptors"
"github.com/grafana/grafana/pkg/setting"
)
// ProvideZanzana used to register ZanzanaClient.
// It will also start an embedded ZanzanaSever if mode is set to "embedded".
func ProvideZanzana(cfg *setting.Cfg, db db.DB, features featuremgmt.FeatureToggles) (zanzana.Client, error) {
func ProvideZanzana(cfg *setting.Cfg, db db.DB, tracer tracing.Tracer, features featuremgmt.FeatureToggles) (zanzana.Client, error) {
if !features.IsEnabledGlobally(featuremgmt.FlagZanzana) {
return zanzana.NewNoopClient(), nil
}
@@ -82,7 +81,7 @@ func ProvideZanzana(cfg *setting.Cfg, db db.DB, features featuremgmt.FeatureTogg
return nil, fmt.Errorf("failed to start zanzana: %w", err)
}
srv, err := zanzana.NewServer(cfg.ZanzanaServer, openfga, logger)
srv, err := zanzana.NewServer(cfg.ZanzanaServer, openfga, logger, tracer)
if err != nil {
return nil, fmt.Errorf("failed to start zanzana: %w", err)
}
@@ -90,7 +89,7 @@ func ProvideZanzana(cfg *setting.Cfg, db db.DB, features featuremgmt.FeatureTogg
channel := &inprocgrpc.Channel{}
// Put * as a namespace so we can properly authorize request with in-proc mode
channel.WithServerUnaryInterceptor(grpcAuth.UnaryServerInterceptor(func(ctx context.Context) (context.Context, error) {
ctx = claims.WithAuthInfo(ctx, authnlib.NewAccessTokenAuthInfo(authnlib.Claims[authnlib.AccessTokenClaims]{
ctx = types.WithAuthInfo(ctx, authnlib.NewAccessTokenAuthInfo(authnlib.Claims[authnlib.AccessTokenClaims]{
Rest: authnlib.AccessTokenClaims{
Namespace: "*",
},
@@ -144,6 +143,18 @@ type Zanzana struct {
}
func (z *Zanzana) start(ctx context.Context) error {
tracingCfg, err := tracing.ProvideTracingConfig(z.cfg)
if err != nil {
return err
}
tracingCfg.ServiceName = "zanzana"
tracer, err := tracing.ProvideService(tracingCfg)
if err != nil {
return err
}
store, err := zanzana.NewStore(z.cfg, z.logger)
if err != nil {
return fmt.Errorf("failed to initilize zanana store: %w", err)
@@ -154,46 +165,27 @@ func (z *Zanzana) start(ctx context.Context) error {
return fmt.Errorf("failed to start zanzana: %w", err)
}
zanzanaServer, err := zanzana.NewServer(z.cfg.ZanzanaServer, openfgaServer, z.logger)
zanzanaServer, err := zanzana.NewServer(z.cfg.ZanzanaServer, openfgaServer, z.logger, tracer)
if err != nil {
return fmt.Errorf("failed to start zanzana: %w", err)
}
tracingCfg, err := tracing.ProvideTracingConfig(z.cfg)
if err != nil {
return err
}
tracingCfg.ServiceName = "zanzana"
tracer, err := tracing.ProvideService(tracingCfg)
if err != nil {
return err
}
authenticator := authnlib.NewAccessTokenAuthenticator(
authnlib.NewAccessTokenVerifier(
authnlib.VerifierConfig{
AllowedAudiences: []string{authzServiceAudience},
},
authnlib.VerifierConfig{AllowedAudiences: []string{authzServiceAudience}},
authnlib.NewKeyRetriever(authnlib.KeyRetrieverConfig{
SigningKeysURL: z.cfg.ZanzanaServer.SigningKeysURL,
}),
),
)
authfn := interceptors.AuthenticatorFunc(func(ctx context.Context) (context.Context, error) {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return nil, fmt.Errorf("missing metadata")
}
c, err := authenticator.Authenticate(ctx, authnlib.NewGRPCTokenProvider(md))
if err != nil {
return nil, err
}
return claims.WithAuthInfo(ctx, c), nil
})
z.handle, err = grpcserver.ProvideService(z.cfg, z.features, authfn, tracer, prometheus.DefaultRegisterer)
z.handle, err = grpcserver.ProvideService(
z.cfg,
z.features,
grpcutils.NewAuthenticatorInterceptor(authenticator, tracer),
tracer,
prometheus.DefaultRegisterer,
)
if err != nil {
return fmt.Errorf("failed to create zanzana grpc server: %w", err)
}
+3 -2
View File
@@ -7,13 +7,14 @@ import (
openfgastorage "github.com/openfga/openfga/pkg/storage"
"github.com/grafana/grafana/pkg/infra/log"
"github.com/grafana/grafana/pkg/infra/tracing"
"github.com/grafana/grafana/pkg/services/authz/zanzana/server"
"github.com/grafana/grafana/pkg/services/grpcserver"
"github.com/grafana/grafana/pkg/setting"
)
func NewServer(cfg setting.ZanzanaServerSettings, openfga server.OpenFGAServer, logger log.Logger) (*server.Server, error) {
return server.NewServer(cfg, openfga, logger)
func NewServer(cfg setting.ZanzanaServerSettings, openfga server.OpenFGAServer, logger log.Logger, tracer tracing.Tracer) (*server.Server, error) {
return server.NewServer(cfg, openfga, logger, tracer)
}
func NewHealthServer(target server.DiagnosticServer) *server.HealthServer {
@@ -36,9 +36,6 @@ func NewOpenFGAServer(cfg setting.ZanzanaServerSettings, store storage.OpenFGADa
server.WithListObjectsDeadline(cfg.ListObjectsDeadline),
}
// FIXME(kalleep): Interceptors
// We probably need to at least need to add store id interceptor also
// would be nice to inject our own requestid?
srv, err := server.NewServerWithOpts(opts...)
if err != nil {
return nil, err
+6 -5
View File
@@ -9,12 +9,12 @@ import (
"github.com/fullstorydev/grpchan/inprocgrpc"
authzv1 "github.com/grafana/authlib/authz/proto/v1"
openfgav1 "github.com/openfga/api/proto/openfga/v1"
"go.opentelemetry.io/otel"
"google.golang.org/protobuf/types/known/wrapperspb"
dashboardalpha1 "github.com/grafana/grafana/pkg/apis/dashboard/v2alpha1"
"github.com/grafana/grafana/pkg/infra/localcache"
"github.com/grafana/grafana/pkg/infra/log"
"github.com/grafana/grafana/pkg/infra/tracing"
authzextv1 "github.com/grafana/grafana/pkg/services/authz/proto/v1"
"github.com/grafana/grafana/pkg/services/authz/zanzana/common"
"github.com/grafana/grafana/pkg/setting"
@@ -25,8 +25,6 @@ const cacheCleanInterval = 2 * time.Minute
var _ authzv1.AuthzServiceServer = (*Server)(nil)
var _ authzextv1.AuthzExtentionServiceServer = (*Server)(nil)
var tracer = otel.Tracer("github.com/grafana/grafana/pkg/services/authz/zanzana/server")
type OpenFGAServer interface {
openfgav1.OpenFGAServiceServer
IsReady(ctx context.Context) (bool, error)
@@ -40,10 +38,12 @@ type Server struct {
openfgaClient openfgav1.OpenFGAServiceClient
cfg setting.ZanzanaServerSettings
logger log.Logger
stores map[string]storeInfo
storesMU *sync.Mutex
cache *localcache.CacheService
logger log.Logger
tracer tracing.Tracer
}
type storeInfo struct {
@@ -51,7 +51,7 @@ type storeInfo struct {
ModelID string
}
func NewServer(cfg setting.ZanzanaServerSettings, openfga OpenFGAServer, logger log.Logger) (*Server, error) {
func NewServer(cfg setting.ZanzanaServerSettings, openfga OpenFGAServer, logger log.Logger, tracer tracing.Tracer) (*Server, error) {
channel := &inprocgrpc.Channel{}
openfgav1.RegisterOpenFGAServiceServer(channel, openfga)
openFGAClient := openfgav1.NewOpenFGAServiceClient(channel)
@@ -64,6 +64,7 @@ func NewServer(cfg setting.ZanzanaServerSettings, openfga OpenFGAServer, logger
cfg: cfg,
cache: localcache.New(cfg.CheckQueryCacheTTL, cacheCleanInterval),
logger: logger,
tracer: tracer,
}
return s, nil
@@ -11,7 +11,7 @@ import (
)
func (s *Server) BatchCheck(ctx context.Context, r *authzextv1.BatchCheckRequest) (*authzextv1.BatchCheckResponse, error) {
ctx, span := tracer.Start(ctx, "server.BatchCheck")
ctx, span := s.tracer.Start(ctx, "server.BatchCheck")
defer span.End()
if err := authorize(ctx, r.GetNamespace()); err != nil {
@@ -10,7 +10,7 @@ import (
)
func (s *Server) Check(ctx context.Context, r *authzv1.CheckRequest) (*authzv1.CheckResponse, error) {
ctx, span := tracer.Start(ctx, "server.Check")
ctx, span := s.tracer.Start(ctx, "server.Check")
defer span.End()
if err := authorize(ctx, r.GetNamespace()); err != nil {
@@ -15,7 +15,7 @@ import (
)
func (s *Server) List(ctx context.Context, r *authzv1.ListRequest) (*authzv1.ListResponse, error) {
ctx, span := tracer.Start(ctx, "server.List")
ctx, span := s.tracer.Start(ctx, "server.List")
defer span.End()
if err := authorize(ctx, r.GetNamespace()); err != nil {
@@ -141,7 +141,7 @@ func (s *Server) listObjects(ctx context.Context, req *openfgav1.ListObjectsRequ
type listFn func(ctx context.Context, req *openfgav1.ListObjectsRequest) (*openfgav1.ListObjectsResponse, error)
func (s *Server) listObjectCached(ctx context.Context, req *openfgav1.ListObjectsRequest, fn listFn) (*openfgav1.ListObjectsResponse, error) {
ctx, span := tracer.Start(ctx, "server.listObjectCached")
ctx, span := s.tracer.Start(ctx, "server.listObjectCached")
defer span.End()
key, err := getRequestHash(req)
@@ -163,7 +163,7 @@ func (s *Server) listObjectCached(ctx context.Context, req *openfgav1.ListObject
}
func (s *Server) streamedListObjects(ctx context.Context, req *openfgav1.ListObjectsRequest) (*openfgav1.ListObjectsResponse, error) {
ctx, span := tracer.Start(ctx, "server.streamedListObjects")
ctx, span := s.tracer.Start(ctx, "server.streamedListObjects")
defer span.End()
r := &openfgav1.StreamedListObjectsRequest{
@@ -10,7 +10,7 @@ import (
)
func (s *Server) Read(ctx context.Context, req *authzextv1.ReadRequest) (*authzextv1.ReadResponse, error) {
ctx, span := tracer.Start(ctx, "server.Read")
ctx, span := s.tracer.Start(ctx, "server.Read")
defer span.End()
if err := authorize(ctx, req.GetNamespace()); err != nil {
@@ -12,6 +12,7 @@ import (
"github.com/grafana/grafana/pkg/infra/db"
"github.com/grafana/grafana/pkg/infra/log"
"github.com/grafana/grafana/pkg/infra/tracing"
"github.com/grafana/grafana/pkg/services/authz/zanzana/common"
"github.com/grafana/grafana/pkg/services/authz/zanzana/store"
"github.com/grafana/grafana/pkg/services/sqlstore/migrator"
@@ -75,7 +76,7 @@ func setup(t *testing.T, testDB db.DB, cfg *setting.Cfg) *Server {
openfga, err := NewOpenFGAServer(cfg.ZanzanaServer, store, log.NewNopLogger())
require.NoError(t, err)
srv, err := NewServer(cfg.ZanzanaServer, openfga, log.NewNopLogger())
srv, err := NewServer(cfg.ZanzanaServer, openfga, log.NewNopLogger(), tracing.NewNoopTracerService())
require.NoError(t, err)
storeInf, err := srv.getStoreInfo(context.Background(), namespace)
@@ -10,7 +10,7 @@ import (
)
func (s *Server) Write(ctx context.Context, req *authzextv1.WriteRequest) (*authzextv1.WriteResponse, error) {
ctx, span := tracer.Start(ctx, "server.Write")
ctx, span := s.tracer.Start(ctx, "server.Write")
defer span.End()
if err := authorize(ctx, req.GetNamespace()); err != nil {