Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions api/v1alpha1/mcp_route.go
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,19 @@ type MCPRouteOAuth struct {
//
// +kubebuilder:validation:Required
ProtectedResourceMetadata ProtectedResourceMetadata `json:"protectedResourceMetadata"`

// ClaimToHeaders specifies JWT claims to extract and forward as HTTP headers to backend MCP servers.
// This enables backends to access user identity for authorization, auditing, or personalization.
//
// Security considerations:
// - Any client-provided headers matching the configured header names will be stripped to prevent forgery
// - Only the specified claims are extracted; the full JWT is not forwarded to backends
// - Consider using a header prefix (e.g., "X-Jwt-Claim-") to avoid conflicts with other headers
//
// +kubebuilder:validation:Optional
// +kubebuilder:validation:MaxItems=16
// +optional
ClaimToHeaders []egv1a1.ClaimToHeader `json:"claimToHeaders,omitempty"`
}

// MCPRouteAuthorization defines the authorization configuration for a MCPRoute.
Expand Down
5 changes: 5 additions & 0 deletions api/v1alpha1/zz_generated.deepcopy.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

6 changes: 6 additions & 0 deletions internal/controller/gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -566,6 +566,12 @@ func mcpConfig(mcpRoutes []aigv1a1.MCPRoute) (_ *filterapi.MCPConfig, hasEffecti
mcpRoute.Authorization.Rules = append(mcpRoute.Authorization.Rules, mcpRule)
}
}
// Add headers to forward from the incoming request to backend MCP servers.
if route.Spec.SecurityPolicy != nil && route.Spec.SecurityPolicy.OAuth != nil {
for _, ctoh := range route.Spec.SecurityPolicy.OAuth.ClaimToHeaders {
mcpRoute.ForwardHeaders = append(mcpRoute.ForwardHeaders, ctoh.Header)
}
}
mc.Routes = append(mc.Routes, mcpRoute)
}
return mc, hasEffectiveRoute
Expand Down
5 changes: 5 additions & 0 deletions internal/controller/mcp_route_security_policy.go
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,11 @@ func (c *MCPRouteController) ensureSecurityPolicy(ctx context.Context, mcpRoute
}
}

// Add ClaimToHeaders to extract JWT claims and set them as HTTP headers.
// Envoy's JWT filter will extract these claims and add them to the request headers,
// which can then be forwarded to backend MCP servers.
jwtProvider.ClaimToHeaders = append(jwtProvider.ClaimToHeaders, oauth.ClaimToHeaders...)

securityPolicySpec.JWT = &egv1a1.JWT{
Providers: []egv1a1.JWTProvider{jwtProvider},
}
Expand Down
52 changes: 52 additions & 0 deletions internal/controller/mcp_route_security_policy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -564,6 +564,58 @@ func TestMCPRouteController_syncMCPRouteSecurityPolicy_DisableOAuthKeepsAPIKey(t
require.True(t, apierrors.IsNotFound(err))
}

func TestMCPRouteController_syncMCPRouteSecurityPolicy_ClaimToHeaders(t *testing.T) {
// Test that ClaimToHeaders from MCPRoute OAuth config are correctly configured
// in the SecurityPolicy's JWTProvider.
fakeClient := requireNewFakeClientWithIndexesForMCP(t)
eventCh := internaltesting.NewControllerEventChan[*gwapiv1.Gateway]()
c := NewMCPRouteController(fakeClient, nil, logr.Discard(), eventCh.Ch)

mcpRoute := &aigv1a1.MCPRoute{
ObjectMeta: metav1.ObjectMeta{Name: "test-route", Namespace: "default"},
Spec: aigv1a1.MCPRouteSpec{
SecurityPolicy: &aigv1a1.MCPRouteSecurityPolicy{
OAuth: &aigv1a1.MCPRouteOAuth{
Issuer: "https://auth.example.com",
Audiences: []string{"test-audience"},
JWKS: &aigv1a1.JWKS{
RemoteJWKS: &egv1a1.RemoteJWKS{
URI: "https://auth.example.com/.well-known/jwks.json",
},
},
ProtectedResourceMetadata: aigv1a1.ProtectedResourceMetadata{
Resource: "https://api.example.com/mcp",
},
ClaimToHeaders: []egv1a1.ClaimToHeader{
{Claim: "sub", Header: "X-User-Id"},
{Claim: "email", Header: "X-User-Email"},
{Claim: "realm_access.roles", Header: "X-User-Roles"},
},
},
},
},
}

require.NoError(t, fakeClient.Create(t.Context(), mcpRoute))

httpRouteName := "test-http-route"
require.NoError(t, c.syncMCPRouteSecurityPolicy(t.Context(), mcpRoute, httpRouteName))

// Verify SecurityPolicy was created with ClaimToHeaders.
securityPolicyName := internalapi.MCPGeneratedResourceCommonPrefix + mcpRoute.Name
var sp egv1a1.SecurityPolicy
require.NoError(t, fakeClient.Get(t.Context(), client.ObjectKey{Name: securityPolicyName, Namespace: mcpRoute.Namespace}, &sp))

require.NotNil(t, sp.Spec.JWT)
require.Len(t, sp.Spec.JWT.Providers, 1)

provider := sp.Spec.JWT.Providers[0]
require.Len(t, provider.ClaimToHeaders, 3)
require.Equal(t, egv1a1.ClaimToHeader{Claim: "sub", Header: "X-User-Id"}, provider.ClaimToHeaders[0])
require.Equal(t, egv1a1.ClaimToHeader{Claim: "email", Header: "X-User-Email"}, provider.ClaimToHeaders[1])
require.Equal(t, egv1a1.ClaimToHeader{Claim: "realm_access.roles", Header: "X-User-Roles"}, provider.ClaimToHeaders[2])
}

func Test_buildOAuthProtectedResourceMetadataJSON(t *testing.T) {
auth := &aigv1a1.MCPRouteOAuth{
Issuer: "https://auth.example.com",
Expand Down
3 changes: 3 additions & 0 deletions internal/filterapi/mcpconfig.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,9 @@ type MCPRoute struct {

// Authorization is the authorization configuration for this route.
Authorization *MCPRouteAuthorization `json:"authorization,omitempty"`

// ForwardHeaders specifies HTTP headers to extract from the incoming request and forward to backend MCP servers.
ForwardHeaders []string `json:"forwardHeaders,omitempty"`
}

// MCPBackend is the MCP backend configuration.
Expand Down
14 changes: 8 additions & 6 deletions internal/mcpproxy/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,10 @@ type (
}

mcpProxyConfigRoute struct {
backends map[filterapi.MCPBackendName]filterapi.MCPBackend
toolSelectors map[filterapi.MCPBackendName]*toolSelector
authorization *compiledAuthorization
backends map[filterapi.MCPBackendName]filterapi.MCPBackend
toolSelectors map[filterapi.MCPBackendName]*toolSelector
authorization *compiledAuthorization
forwardHeaders []string
}

// toolSelector filters tools using include patterns with exact matches or regular expressions.
Expand Down Expand Up @@ -151,9 +152,10 @@ func (p *ProxyConfig) LoadConfig(_ context.Context, config *filterapi.Config) er
}

r := &mcpProxyConfigRoute{
backends: make(map[filterapi.MCPBackendName]filterapi.MCPBackend, len(route.Backends)),
toolSelectors: make(map[filterapi.MCPBackendName]*toolSelector, len(route.Backends)),
authorization: compiledAuth,
backends: make(map[filterapi.MCPBackendName]filterapi.MCPBackend, len(route.Backends)),
toolSelectors: make(map[filterapi.MCPBackendName]*toolSelector, len(route.Backends)),
authorization: compiledAuth,
forwardHeaders: route.ForwardHeaders,
}
for _, backend := range route.Backends {
r.backends[backend.Name] = backend
Expand Down
19 changes: 19 additions & 0 deletions internal/mcpproxy/handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -1184,6 +1184,25 @@ func extractSubject(r *http.Request) string {
return claims.Subject
}

// extractForwardHeaders reads the configured headers from the incoming request to forward to backends.
func extractForwardHeaders(reqHeaders http.Header, headers []string) map[string]string {
if len(headers) == 0 {
return nil
}

result := make(map[string]string)
for _, header := range headers {
if value := reqHeaders.Get(header); value != "" {
result[header] = value
}
}

if len(result) == 0 {
return nil
}
return result
}

// handlePromptGetRequest handles the "prompts/get" JSON-RPC method.
func (m *mcpRequestContext) handlePromptGetRequest(ctx context.Context, s *session, w http.ResponseWriter, req *jsonrpc.Request, p *mcp.GetPromptParams) error {
backendName, promptName, err := upstreamResourceName(p.Name)
Expand Down
77 changes: 77 additions & 0 deletions internal/mcpproxy/handlers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1226,6 +1226,83 @@ func TestExtractSubject(t *testing.T) {
})
}

func TestExtractForwardHeaders(t *testing.T) {
// Test that extractForwardHeaders correctly reads headers from the request.
tests := []struct {
name string
requestHeaders map[string]string
forwardHeaders []string
wantHeaders map[string]string
}{
{
name: "extract configured headers",
requestHeaders: map[string]string{
"X-User-Id": "user123",
"X-User-Email": "user@example.com",
},
forwardHeaders: []string{"X-User-Id", "X-User-Email"},
wantHeaders: map[string]string{
"X-User-Id": "user123",
"X-User-Email": "user@example.com",
},
},
{
name: "extract nested claim header",
requestHeaders: map[string]string{
"X-User-Roles": `["admin","user"]`,
},
forwardHeaders: []string{"X-User-Roles"},
wantHeaders: map[string]string{
"X-User-Roles": `["admin","user"]`,
},
},
{
name: "missing header returns nil",
requestHeaders: map[string]string{
// X-Missing is not set
},
forwardHeaders: []string{"X-Missing"},
wantHeaders: nil,
},
{
name: "mixed existing and missing headers",
requestHeaders: map[string]string{
"X-User-Id": "user123",
// X-Missing is not set
},
forwardHeaders: []string{"X-User-Id", "X-Missing"},
wantHeaders: map[string]string{
"X-User-Id": "user123",
},
},
{
name: "empty forward headers",
requestHeaders: map[string]string{"X-User-Id": "user123"},
forwardHeaders: []string{},
wantHeaders: nil,
},
{
name: "no headers on request",
requestHeaders: map[string]string{},
forwardHeaders: []string{"X-User-Id"},
wantHeaders: nil,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
headers := make(http.Header)
// Set headers (simulating what Envoy's JWT filter does)
for header, value := range tt.requestHeaders {
headers.Set(header, value)
}

result := extractForwardHeaders(headers, tt.forwardHeaders)
require.Equal(t, tt.wantHeaders, result)
})
}
}

func secureID(t *testing.T, proxy *mcpRequestContext, sessionID string) string {
secure, err := proxy.sessionCrypto.Encrypt(sessionID)
require.NoError(t, err)
Expand Down
45 changes: 38 additions & 7 deletions internal/mcpproxy/mcpproxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -137,16 +137,18 @@ func extractMetaFromJSONRPCMessage(msg jsonrpc.Message) map[string]any {
func (m *mcpRequestContext) newSession(ctx context.Context, p *mcp.InitializeParams, routeName filterapi.MCPRouteName, subject string, span tracingapi.MCPSpan) (*session, error) {
m.l.Debug("creating new MCP session")

backends := m.routes[routeName]
if backends == nil {
return nil, fmt.Errorf("no backends found for route %s", routeName)
}

forwardHeaders := extractForwardHeaders(m.requestHeaders, backends.forwardHeaders)

var (
wg sync.WaitGroup
entries []compositeSessionEntry
counter int
)

backends := m.routes[routeName]
if backends == nil {
return nil, fmt.Errorf("no backends found for route %s", routeName)
}
entries = make([]compositeSessionEntry, len(backends.backends))

if m.l.Enabled(ctx, slog.LevelDebug) {
Expand Down Expand Up @@ -198,7 +200,21 @@ func (m *mcpRequestContext) newSession(ctx context.Context, p *mcp.InitializePar
if err != nil {
return nil, fmt.Errorf("failed to encrypt session ID: %w", err)
}
return &session{reqCtx: m, id: secureClientToGatewaySessionID(encrypted)}, nil

// Build perBackendSessions map from finalEntries
perBackendSessions := make(map[filterapi.MCPBackendName]*compositeSessionEntry, len(finalEntries))
for i := range finalEntries {
entry := &finalEntries[i]
perBackendSessions[entry.backendName] = entry
}

return &session{
reqCtx: m,
id: secureClientToGatewaySessionID(encrypted),
route: routeName,
perBackendSessions: perBackendSessions,
extraHeaders: forwardHeaders,
}, nil
}

// sessionFromID returns the session with the given ID, or error if not found or invalid.
Expand Down Expand Up @@ -226,7 +242,13 @@ func (m *mcpRequestContext) sessionFromID(id secureClientToGatewaySessionID, las
}
}

return &session{id: id, route: route, reqCtx: m, perBackendSessions: perBackendSessionIDs}, nil
// Extract forward headers from the current request based on the route's forwardHeaders config.
var extraHeaders map[string]string
if routeConfig := m.routes[route]; routeConfig != nil {
extraHeaders = extractForwardHeaders(m.requestHeaders, routeConfig.forwardHeaders)
}

return &session{id: id, route: route, reqCtx: m, perBackendSessions: perBackendSessionIDs, extraHeaders: extraHeaders}, nil
}

type initializeResult struct {
Expand Down Expand Up @@ -375,6 +397,15 @@ func (m *mcpRequestContext) invokeJSONRPCRequest(ctx context.Context, routeName
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")

// Forward configured headers to backend.
if routeConfig := m.routes[routeName]; routeConfig != nil {
for _, header := range routeConfig.forwardHeaders {
if value := m.requestHeaders.Get(header); value != "" {
req.Header.Set(header, value)
}
}
}

resp, err := m.client.Do(req)
if err != nil {
return nil, fmt.Errorf("failed to send MCP notifications/initialized request: %w", err)
Expand Down
13 changes: 13 additions & 0 deletions internal/mcpproxy/session.go
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,11 @@ type session struct {
reqCtx *mcpRequestContext
mu sync.RWMutex
perBackendSessions map[filterapi.MCPBackendName]*compositeSessionEntry
// extraHeaders contains header values extracted from the current HTTP request to be forwarded to backends.
// These are derived from the route's configured forward headers and the current request's headers.
// The key is the HTTP header name, the value is the header value.
// Note: extraHeaders is NOT encoded in the session ID. It is re-extracted from each incoming request.
extraHeaders map[string]string
}

// Close implements [io.Closer.Close].
Expand Down Expand Up @@ -340,6 +345,14 @@ func (s *session) sendRequestPerBackend(ctx context.Context, eventChan chan<- *s
req.Header.Set("Accept", "text/event-stream, application/json")
req.Header.Set("Accept-encoding", "gzip, br, zstd, deflate")

// Forward configured headers to the backend.
// First, strip any client-provided headers that match configured forward headers to prevent forgery.
// Then set the values extracted from the original request.
for header, value := range s.extraHeaders {
req.Header.Del(header) // Prevent forgery by stripping client-provided headers.
req.Header.Set(header, value)
}

if lastEventID := cse.lastEventID; lastEventID != "" {
req.Header.Set(lastEventIDHeader, lastEventID)
}
Expand Down
Loading