diff --git a/pkg/clientauth/providers.go b/pkg/clientauth/providers.go new file mode 100644 index 00000000000..fcb2cd77240 --- /dev/null +++ b/pkg/clientauth/providers.go @@ -0,0 +1,43 @@ +package clientauth + +import ( + "context" +) + +// NamespaceProvider is a strategy for determining the namespace to use in token exchange requests. +type NamespaceProvider interface { + GetNamespace(ctx context.Context) string +} + +// AudienceProvider is a strategy for determining the audiences to use in token exchange requests. +type AudienceProvider interface { + GetAudiences(ctx context.Context) []string +} + +// StaticNamespaceProvider returns a fixed namespace for all requests. +type StaticNamespaceProvider struct { + namespace string +} + +// NewStaticNamespaceProvider creates a namespace provider that always returns the same namespace. +func NewStaticNamespaceProvider(namespace string) *StaticNamespaceProvider { + return &StaticNamespaceProvider{namespace: namespace} +} + +func (p *StaticNamespaceProvider) GetNamespace(ctx context.Context) string { + return p.namespace +} + +// StaticAudienceProvider returns a fixed set of audiences for all requests. +type StaticAudienceProvider struct { + audiences []string +} + +// NewStaticAudienceProvider creates an audience provider that always returns the same audiences. +func NewStaticAudienceProvider(audiences ...string) *StaticAudienceProvider { + return &StaticAudienceProvider{audiences: audiences} +} + +func (p *StaticAudienceProvider) GetAudiences(ctx context.Context) []string { + return p.audiences +} diff --git a/pkg/clientauth/providers_test.go b/pkg/clientauth/providers_test.go new file mode 100644 index 00000000000..fdacfbeafb6 --- /dev/null +++ b/pkg/clientauth/providers_test.go @@ -0,0 +1,78 @@ +package clientauth + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestStaticNamespaceProvider(t *testing.T) { + tests := []struct { + name string + namespace string + expectedNamespace string + }{ + { + name: "wildcard namespace", + namespace: "*", + expectedNamespace: "*", + }, + { + name: "specific namespace", + namespace: "my-namespace", + expectedNamespace: "my-namespace", + }, + { + name: "empty namespace", + namespace: "", + expectedNamespace: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider := NewStaticNamespaceProvider(tt.namespace) + result := provider.GetNamespace(context.Background()) + require.Equal(t, tt.expectedNamespace, result) + }) + } +} + +func TestStaticAudienceProvider(t *testing.T) { + tests := []struct { + name string + audiences []string + expectedAudiences []string + }{ + { + name: "single audience", + audiences: []string{"folder.grafana.app"}, + expectedAudiences: []string{"folder.grafana.app"}, + }, + { + name: "multiple audiences", + audiences: []string{"audience1", "audience2", "audience3"}, + expectedAudiences: []string{"audience1", "audience2", "audience3"}, + }, + { + name: "empty audiences", + audiences: []string{}, + expectedAudiences: []string{}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider := NewStaticAudienceProvider(tt.audiences...) + result := provider.GetAudiences(context.Background()) + require.Equal(t, tt.expectedAudiences, result) + }) + } +} + +func TestProviderInterfaces(t *testing.T) { + // Verify that all providers implement their interfaces + var _ NamespaceProvider = (*StaticNamespaceProvider)(nil) + var _ AudienceProvider = (*StaticAudienceProvider)(nil) +} diff --git a/pkg/clientauth/roundtripper.go b/pkg/clientauth/roundtripper.go new file mode 100644 index 00000000000..225f6afcad9 --- /dev/null +++ b/pkg/clientauth/roundtripper.go @@ -0,0 +1,90 @@ +package clientauth + +import ( + "fmt" + "net/http" + + authnlib "github.com/grafana/authlib/authn" + utilnet "k8s.io/apimachinery/pkg/util/net" + "k8s.io/client-go/transport" +) + +// tokenExchangeRoundTripper wraps an http.RoundTripper and injects an exchanged +// access token into outgoing requests via the X-Access-Token header. +type tokenExchangeRoundTripper struct { + exchanger authnlib.TokenExchanger + transport http.RoundTripper + namespaceProvider NamespaceProvider + audienceProvider AudienceProvider +} + +var _ http.RoundTripper = (*tokenExchangeRoundTripper)(nil) + +// newTokenExchangeRoundTripperWithStrategies creates a new RoundTripper with custom +// namespace and audience strategies, allowing for flexible configuration. +func newTokenExchangeRoundTripperWithStrategies( + exchanger authnlib.TokenExchanger, + transport http.RoundTripper, + namespaceProvider NamespaceProvider, + audienceProvider AudienceProvider, +) *tokenExchangeRoundTripper { + return &tokenExchangeRoundTripper{ + exchanger: exchanger, + transport: transport, + namespaceProvider: namespaceProvider, + audienceProvider: audienceProvider, + } +} + +// RoundTrip implements http.RoundTripper by exchanging a token and setting it +// in the X-Access-Token header before forwarding the request. +func (t *tokenExchangeRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + ctx := req.Context() + + tokenResponse, err := t.exchanger.Exchange(ctx, authnlib.TokenExchangeRequest{ + Audiences: t.audienceProvider.GetAudiences(ctx), + Namespace: t.namespaceProvider.GetNamespace(ctx), + }) + if err != nil { + return nil, fmt.Errorf("failed to exchange token: %w", err) + } + + // Clone the request as RoundTrippers are not expected to mutate the passed request + req = utilnet.CloneRequest(req) + + req.Header.Set("X-Access-Token", "Bearer "+tokenResponse.Token) + + return t.transport.RoundTrip(req) +} + +// NewStaticTokenExchangeTransportWrapper creates a transport.WrapperFunc that wraps +// an http.RoundTripper with token exchange authentication for use with k8s +// rest.Config.WrapTransport. +func NewStaticTokenExchangeTransportWrapper( + exchanger authnlib.TokenExchanger, + audience string, + namespace string, +) transport.WrapperFunc { + return func(rt http.RoundTripper) http.RoundTripper { + return newTokenExchangeRoundTripperWithStrategies(exchanger, rt, NewStaticNamespaceProvider(namespace), NewStaticAudienceProvider(audience)) + } +} + +// NewTokenExchangeTransportWrapperWithStrategies creates a transport.WrapperFunc with custom strategies. +func NewTokenExchangeTransportWrapper( + exchanger authnlib.TokenExchanger, + namespaceProvider NamespaceProvider, + audienceProvider AudienceProvider, +) transport.WrapperFunc { + return func(rt http.RoundTripper) http.RoundTripper { + return newTokenExchangeRoundTripperWithStrategies( + exchanger, + rt, + namespaceProvider, + audienceProvider, + ) + } +} + +// WildcardNamespace is a convenience constant for the wildcard namespace. +const WildcardNamespace = "*" diff --git a/pkg/clientauth/roundtripper_test.go b/pkg/clientauth/roundtripper_test.go new file mode 100644 index 00000000000..96ee99406b3 --- /dev/null +++ b/pkg/clientauth/roundtripper_test.go @@ -0,0 +1,259 @@ +package clientauth + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/grafana/authlib/authn" + "github.com/stretchr/testify/require" +) + +type fakeExchanger struct { + resp *authn.TokenExchangeResponse + err error + gotReq *authn.TokenExchangeRequest +} + +func (f *fakeExchanger) Exchange(_ context.Context, req authn.TokenExchangeRequest) (*authn.TokenExchangeResponse, error) { + f.gotReq = &req + return f.resp, f.err +} + +// roundTripperFunc allows building a stub transport inline +type roundTripperFunc func(*http.Request) (*http.Response, error) + +func (f roundTripperFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func TestTokenExchangeRoundTripper_SetsAccessTokenHeader(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "test-token-123"}} + + var capturedHeader string + transport := roundTripperFunc(func(r *http.Request) (*http.Response, error) { + capturedHeader = r.Header.Get("X-Access-Token") + rr := httptest.NewRecorder() + rr.WriteHeader(http.StatusOK) + return rr.Result(), nil + }) + + rt := newTokenExchangeRoundTripperWithStrategies(exchanger, transport, NewStaticNamespaceProvider("test-namespace"), NewStaticAudienceProvider("test-audience")) + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + + resp, err := rt.RoundTrip(req) + require.NoError(t, err) + if resp != nil { + _ = resp.Body.Close() + } + + // Clean up response + _ = resp.Body.Close() + + require.Equal(t, "Bearer test-token-123", capturedHeader) +} + +func TestTokenExchangeRoundTripper_PropagatesExchangeError(t *testing.T) { + expectedErr := errors.New("token exchange failed") + exchanger := &fakeExchanger{err: expectedErr} + + transport := roundTripperFunc(func(_ *http.Request) (*http.Response, error) { + t.Fatal("transport should not be called on exchange error") + return nil, nil + }) + + rt := newTokenExchangeRoundTripperWithStrategies(exchanger, transport, NewStaticNamespaceProvider("test-namespace"), NewStaticAudienceProvider("test-audience")) + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + + resp, err := rt.RoundTrip(req) + require.Error(t, err) + if resp != nil { + _ = resp.Body.Close() + } + require.ErrorContains(t, err, "failed to exchange token") + require.ErrorIs(t, err, expectedErr) +} + +func TestTokenExchangeRoundTripper_SendsCorrectAudienceAndNamespace(t *testing.T) { + tests := []struct { + name string + audience string + namespace string + expectedAudiences []string + expectedNamespace string + }{ + { + name: "single audience with wildcard namespace", + audience: "folder.grafana.app", + namespace: "*", + expectedAudiences: []string{"folder.grafana.app"}, + expectedNamespace: "*", + }, + { + name: "different audience with wildcard namespace", + audience: "dashboard.grafana.app", + namespace: "*", + expectedAudiences: []string{"dashboard.grafana.app"}, + expectedNamespace: "*", + }, + { + name: "audience with specific namespace", + audience: "test-audience", + namespace: "test-namespace", + expectedAudiences: []string{"test-audience"}, + expectedNamespace: "test-namespace", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "token"}} + transport := roundTripperFunc(func(_ *http.Request) (*http.Response, error) { + rr := httptest.NewRecorder() + rr.WriteHeader(http.StatusOK) + return rr.Result(), nil + }) + + rt := newTokenExchangeRoundTripperWithStrategies(exchanger, transport, NewStaticNamespaceProvider(tt.namespace), NewStaticAudienceProvider(tt.audience)) + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + + resp, err := rt.RoundTrip(req) + require.NoError(t, err) + if resp != nil { + _ = resp.Body.Close() + } + + require.NotNil(t, exchanger.gotReq) + require.Equal(t, tt.expectedAudiences, exchanger.gotReq.Audiences) + require.Equal(t, tt.expectedNamespace, exchanger.gotReq.Namespace) + }) + } +} + +func TestTokenExchangeRoundTripper_DoesNotMutateOriginalRequest(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "token"}} + transport := roundTripperFunc(func(r *http.Request) (*http.Response, error) { + rr := httptest.NewRecorder() + rr.WriteHeader(http.StatusOK) + return rr.Result(), nil + }) + + rt := newTokenExchangeRoundTripperWithStrategies(exchanger, transport, NewStaticNamespaceProvider("namespace"), NewStaticAudienceProvider("audience")) + originalReq, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + + // Ensure original request has no X-Access-Token header + originalReq.Header.Set("X-Custom-Header", "original-value") + require.Empty(t, originalReq.Header.Get("X-Access-Token")) + + resp, err := rt.RoundTrip(originalReq) + require.NoError(t, err) + _ = resp.Body.Close() + + // Original request should not have been mutated + require.Empty(t, originalReq.Header.Get("X-Access-Token")) + require.Equal(t, "original-value", originalReq.Header.Get("X-Custom-Header")) +} + +func TestTokenExchangeRoundTripper_PropagatesTransportError(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "token"}} + expectedErr := errors.New("transport error") + transport := roundTripperFunc(func(_ *http.Request) (*http.Response, error) { + return nil, expectedErr + }) + + rt := newTokenExchangeRoundTripperWithStrategies(exchanger, transport, NewStaticNamespaceProvider("namespace"), NewStaticAudienceProvider("audience")) + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + + resp, err := rt.RoundTrip(req) + require.Error(t, err) + if resp != nil { + _ = resp.Body.Close() + } + require.ErrorIs(t, err, expectedErr) +} + +func TestNewTokenExchangeTransportWrapper(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "wrapped-token"}} + + var capturedHeader string + baseTransport := roundTripperFunc(func(r *http.Request) (*http.Response, error) { + capturedHeader = r.Header.Get("X-Access-Token") + rr := httptest.NewRecorder() + rr.WriteHeader(http.StatusOK) + return rr.Result(), nil + }) + + wrapper := NewStaticTokenExchangeTransportWrapper(exchanger, "test-audience", "test-namespace") + wrappedTransport := wrapper(baseTransport) + + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + resp, err := wrappedTransport.RoundTrip(req) + require.NoError(t, err) + _ = resp.Body.Close() + + require.Equal(t, "Bearer wrapped-token", capturedHeader) + require.NotNil(t, exchanger.gotReq) + require.Equal(t, []string{"test-audience"}, exchanger.gotReq.Audiences) + require.Equal(t, "test-namespace", exchanger.gotReq.Namespace) +} + +func TestTokenExchangeRoundTripperWithStrategies(t *testing.T) { + tests := []struct { + name string + namespaceProvider NamespaceProvider + audienceProvider AudienceProvider + expectedNamespace string + expectedAudiences []string + expectedHeader string + }{ + { + name: "static providers with bearer prefix", + namespaceProvider: NewStaticNamespaceProvider("*"), + audienceProvider: NewStaticAudienceProvider("folder.grafana.app"), + expectedNamespace: "*", + expectedAudiences: []string{"folder.grafana.app"}, + expectedHeader: "Bearer test-token", + }, + { + name: "multiple audiences", + namespaceProvider: NewStaticNamespaceProvider("*"), + audienceProvider: NewStaticAudienceProvider("audience1", "audience2"), + expectedNamespace: "*", + expectedAudiences: []string{"audience1", "audience2"}, + expectedHeader: "Bearer test-token", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + exchanger := &fakeExchanger{resp: &authn.TokenExchangeResponse{Token: "test-token"}} + + var capturedHeader string + transport := roundTripperFunc(func(r *http.Request) (*http.Response, error) { + capturedHeader = r.Header.Get("X-Access-Token") + rr := httptest.NewRecorder() + rr.WriteHeader(http.StatusOK) + return rr.Result(), nil + }) + + rt := newTokenExchangeRoundTripperWithStrategies( + exchanger, + transport, + tt.namespaceProvider, + tt.audienceProvider, + ) + + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.org", nil) + resp, err := rt.RoundTrip(req) + require.NoError(t, err) + if resp != nil { + _ = resp.Body.Close() + } + + require.Equal(t, tt.expectedHeader, capturedHeader) + require.NotNil(t, exchanger.gotReq) + require.Equal(t, tt.expectedAudiences, exchanger.gotReq.Audiences) + require.Equal(t, tt.expectedNamespace, exchanger.gotReq.Namespace) + }) + } +} diff --git a/pkg/operators/iam/zanzana_folder_reconciler.go b/pkg/operators/iam/zanzana_folder_reconciler.go index ee2233d7877..2fabef43dbd 100644 --- a/pkg/operators/iam/zanzana_folder_reconciler.go +++ b/pkg/operators/iam/zanzana_folder_reconciler.go @@ -6,24 +6,22 @@ import ( "errors" "fmt" "log/slog" - "net/http" "os" "os/signal" "syscall" "github.com/prometheus/client_golang/prometheus" "k8s.io/client-go/rest" - "k8s.io/client-go/transport" "github.com/grafana/grafana-app-sdk/logging" "github.com/grafana/grafana-app-sdk/operator" folder "github.com/grafana/grafana/apps/folder/pkg/apis/folder/v1beta1" "github.com/grafana/grafana/apps/iam/pkg/app" + "github.com/grafana/grafana/pkg/clientauth" "github.com/grafana/grafana/pkg/server" "github.com/grafana/grafana/pkg/setting" "github.com/grafana/authlib/authn" - utilnet "k8s.io/apimachinery/pkg/util/net" ) func RunIAMFolderReconciler(deps server.OperatorDependencies) error { @@ -151,12 +149,11 @@ func buildKubeConfigFromFolderAppURL( return &rest.Config{ APIPath: "/apis", Host: folderAppURL, - WrapTransport: transport.WrapperFunc(func(rt http.RoundTripper) http.RoundTripper { - return &authRoundTripper{ - tokenExchangeClient: tokenExchangeClient, - transport: rt, - } - }), + WrapTransport: clientauth.NewStaticTokenExchangeTransportWrapper( + tokenExchangeClient, + folder.GROUP, + clientauth.WildcardNamespace, + ), TLSClientConfig: tlsConfig, }, nil } @@ -189,23 +186,3 @@ func buildTLSConfig(insecure bool, certFile, keyFile, caFile string) (rest.TLSCl return tlsConfig, nil } - -type authRoundTripper struct { - tokenExchangeClient *authn.TokenExchangeClient - transport http.RoundTripper -} - -func (t *authRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { - tokenResponse, err := t.tokenExchangeClient.Exchange(req.Context(), authn.TokenExchangeRequest{ - Audiences: []string{folder.GROUP}, - Namespace: "*", - }) - if err != nil { - return nil, fmt.Errorf("failed to exchange token: %w", err) - } - - // clone the request as RTs are not expected to mutate the passed request - req = utilnet.CloneRequest(req) - req.Header.Set("X-Access-Token", "Bearer "+tokenResponse.Token) - return t.transport.RoundTrip(req) -} diff --git a/pkg/registry/apis/iam/authorizer/parent_provider.go b/pkg/registry/apis/iam/authorizer/parent_provider.go index 4b7555d85ba..e64f73c2ce5 100644 --- a/pkg/registry/apis/iam/authorizer/parent_provider.go +++ b/pkg/registry/apis/iam/authorizer/parent_provider.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "net/http" "sync" "github.com/grafana/authlib/authn" @@ -15,8 +14,8 @@ import ( dashboardv1 "github.com/grafana/grafana/apps/dashboard/pkg/apis/dashboard/v1beta1" folderv1 "github.com/grafana/grafana/apps/folder/pkg/apis/folder/v1beta1" - "github.com/grafana/grafana/apps/provisioning/pkg/auth" "github.com/grafana/grafana/pkg/apimachinery/utils" + "github.com/grafana/grafana/pkg/clientauth" ) var ( @@ -73,10 +72,8 @@ func NewRemoteConfigProvider(cfg map[schema.GroupResource]DialConfig, exchangeCl for gr, dialConfig := range cfg { configProviders[gr] = func(ctx context.Context) (*rest.Config, error) { return &rest.Config{ - Host: dialConfig.Host, - WrapTransport: func(rt http.RoundTripper) http.RoundTripper { - return auth.NewRoundTripper(exchangeClient, rt, dialConfig.Audience) - }, + Host: dialConfig.Host, + WrapTransport: clientauth.NewStaticTokenExchangeTransportWrapper(exchangeClient, dialConfig.Audience, clientauth.WildcardNamespace), TLSClientConfig: rest.TLSClientConfig{ Insecure: dialConfig.Insecure, CAFile: dialConfig.CAFile, diff --git a/pkg/services/authz/rbac.go b/pkg/services/authz/rbac.go index c6c078a9b4a..a96ffc5d61d 100644 --- a/pkg/services/authz/rbac.go +++ b/pkg/services/authz/rbac.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "net/http" "time" "github.com/fullstorydev/grpchan/inprocgrpc" @@ -24,6 +23,7 @@ import ( authlib "github.com/grafana/authlib/types" "github.com/grafana/dskit/middleware" + "github.com/grafana/grafana/pkg/clientauth" "github.com/grafana/grafana/pkg/infra/db" "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/infra/tracing" @@ -262,9 +262,11 @@ func RegisterRBACAuthZService( folderStore = store.NewAPIFolderStore(tracer, reg, func(ctx context.Context) (*rest.Config, error) { return &rest.Config{ Host: cfg.Folder.Host, - WrapTransport: func(rt http.RoundTripper) http.RoundTripper { - return &tokenExhangeRoundTripper{te: exchangeClient, rt: rt} - }, + WrapTransport: clientauth.NewStaticTokenExchangeTransportWrapper( + exchangeClient, + "folder.grafana.app", + clientauth.WildcardNamespace, + ), TLSClientConfig: rest.TLSClientConfig{ Insecure: cfg.Folder.Insecure, CAFile: cfg.Folder.CAFile, @@ -291,27 +293,6 @@ func RegisterRBACAuthZService( authzv1.RegisterAuthzServiceServer(srv, server) } -var _ http.RoundTripper = tokenExhangeRoundTripper{} - -type tokenExhangeRoundTripper struct { - te authnlib.TokenExchanger - rt http.RoundTripper -} - -func (t tokenExhangeRoundTripper) RoundTrip(r *http.Request) (*http.Response, error) { - res, err := t.te.Exchange(r.Context(), authnlib.TokenExchangeRequest{ - Namespace: "*", - Audiences: []string{"folder.grafana.app"}, - }) - - if err != nil { - return nil, fmt.Errorf("create access token: %w", err) - } - - r.Header.Set("X-Access-Token", "Bearer "+res.Token) - return t.rt.RoundTrip(r) -} - type NoopCache struct{} func (lc *NoopCache) Get(ctx context.Context, key string) ([]byte, error) {