Chore: Unify token exchange round trippers (#115609)

* Chore: Unify token exchange rount trippers

* Remove the conditional provider for now

* Remove unecessary strategy

* test cleanup

* Lint
This commit is contained in:
Gabriel MABILLE
2026-01-05 11:23:35 +01:00
committed by GitHub
parent 76a6db818e
commit 93566ce4ef
7 changed files with 485 additions and 60 deletions
+43
View File
@@ -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
}
+78
View File
@@ -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)
}
+90
View File
@@ -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 = "*"
+259
View File
@@ -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)
})
}
}
+6 -29
View File
@@ -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)
}
@@ -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,
+6 -25
View File
@@ -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) {