Skip to content
Open
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
28 changes: 16 additions & 12 deletions go/adk/pkg/mcp/registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -259,13 +259,15 @@ func createTransport(ctx context.Context, params mcpServerParams) (mcpsdk.Transp

// headerRoundTripper wraps an http.RoundTripper to add custom headers to all
// requests. It supports four sources of headers, applied in this order so that
// higher-priority sources win on collision:
// 1. propagateToken: when true, Authorization is read from the incoming A2A
// later sources win on collision:
// 1. headers: static key/value pairs configured on the MCP server spec. These
// are applied first, as defaults — a dynamic source with the same header
// name overrides them.
// 2. propagateToken: when true, Authorization is read from the incoming A2A
// CallContext and forwarded unconditionally (independent of allowedHeaders).
// 2. allowedHeaders: explicit per-header forwarding from the A2A CallContext.
// 3. headerProvider: runtime headers derived from ADK context, such as STS tokens.
// 4. headers: static key/value pairs configured on the MCP server spec (highest
// priority — always wins).
// 3. allowedHeaders: explicit per-header forwarding from the A2A CallContext.
// 4. headerProvider: runtime headers derived from ADK context, such as STS
// tokens (highest priority — always wins).
type headerRoundTripper struct {
base http.RoundTripper
headers map[string]string
Expand All @@ -277,6 +279,12 @@ type headerRoundTripper struct {
func (rt *headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())

// Apply static headers first, as defaults — any dynamic source below that
// sets the same header name overrides them.
for key, value := range rt.headers {
req.Header.Set(key, value)
}

// When KAGENT_PROPAGATE_TOKEN is set, forward Authorization from the incoming
// A2A request independently of allowedHeaders.
if rt.propagateToken {
Expand All @@ -294,18 +302,14 @@ func (rt *headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, erro
req.Header.Set(k, v)
}

// Dynamic headers (e.g., STS access tokens) override propagated/allowed headers.
// Dynamic headers (e.g., STS access tokens) override propagated/allowed
// headers, and are applied last so they win over the static default too.
if rt.headerProvider != nil {
for key, value := range rt.headerProvider(req.Context()) {
req.Header.Set(key, value)
}
}

// Apply static headers last — they take precedence over all dynamic sources.
for key, value := range rt.headers {
req.Header.Set(key, value)
}

return rt.base.RoundTrip(req)
}

Expand Down
99 changes: 91 additions & 8 deletions go/adk/pkg/mcp/registry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -79,9 +79,49 @@ func TestAllowedRequestHeaders_ForwardsMatchingHeaders(t *testing.T) {
}
}

// TestAllowedRequestHeaders_StaticOverridesDynamic verifies that a statically
// configured header wins over the same header forwarded from the A2A request.
func TestAllowedRequestHeaders_StaticOverridesDynamic(t *testing.T) {
// TestUnlistedRequestHeader_DoesNotOverrideStatic verifies the override in
// TestAllowedRequestHeaders_DynamicOverridesStatic is gated on the header
// being explicitly opted in. An incoming request carrying the same header
// name as a static default, with neither allowedHeaders nor propagateToken
// configured for it, must not override the static value — a request cannot
// silently clobber a static header just by sending one with a matching name.
func TestUnlistedRequestHeader_DoesNotOverrideStatic(t *testing.T) {
t.Parallel()
var capturedAuth string

srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
capturedAuth = r.Header.Get("Authorization")
w.WriteHeader(http.StatusOK)
}))
defer srv.Close()

ctx := a2aCtx(map[string][]string{
"Authorization": {"Bearer incoming"},
})

rt := &headerRoundTripper{
base: newTestTransport(t),
headers: map[string]string{"Authorization": "Bearer static"},
// Deliberately no allowedHeaders, no propagateToken, no headerProvider:
// nothing opts Authorization into being forwarded from the request.
}

req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL, nil)
resp, err := rt.RoundTrip(req)
if err != nil {
t.Fatalf("RoundTrip failed: %v", err)
}
resp.Body.Close()

if capturedAuth != "Bearer static" {
t.Errorf("Authorization: got %q, want %q", capturedAuth, "Bearer static")
}
Comment thread
onematchfox marked this conversation as resolved.
}

// TestAllowedRequestHeaders_DynamicOverridesStatic verifies that a header
// forwarded from the A2A request (via allowedHeaders) wins over a statically
// configured header of the same name — the static header is only a default.
func TestAllowedRequestHeaders_DynamicOverridesStatic(t *testing.T) {
t.Parallel()
var capturedAuth string

Expand All @@ -108,8 +148,8 @@ func TestAllowedRequestHeaders_StaticOverridesDynamic(t *testing.T) {
}
resp.Body.Close()

if capturedAuth != "Bearer static" {

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This assertion was wrong to begin with. Header round-tripper is configured above with allowedHeaders: []string{"Authorization"}, indicating that the Authorization header (the user's token) should be forwarded.

t.Errorf("Authorization: got %q, want %q", capturedAuth, "Bearer static")
if capturedAuth != "Bearer incoming" {
t.Errorf("Authorization: got %q, want %q", capturedAuth, "Bearer incoming")
}
}

Expand Down Expand Up @@ -535,9 +575,12 @@ func TestDynamicHeaders_OverridePropagatedAndAllowedHeaders(t *testing.T) {
}
}

// TestStaticHeaders_OverrideDynamic verifies static configured headers remain
// the highest-precedence source.
func TestStaticHeaders_OverrideDynamic(t *testing.T) {
// TestDynamicHeaders_OverrideStatic verifies a headerProvider-sourced header
// (e.g. an STS-exchanged or propagated Authorization) wins over a static header
// of the same name configured on the MCP server spec. This is what lets a
// per-user token override a RemoteMCPServer's static headersFrom Authorization
// that the controller needs for its own tool-discovery handshake.
func TestDynamicHeaders_OverrideStatic(t *testing.T) {
t.Parallel()
var capturedAuth string

Expand All @@ -562,7 +605,47 @@ func TestStaticHeaders_OverrideDynamic(t *testing.T) {
}
resp.Body.Close()

if capturedAuth != "Bearer dynamic" {
t.Errorf("Authorization: got %q, want %q", capturedAuth, "Bearer dynamic")
}
}

// TestStaticHeaders_ApplyWhenNoDynamicSource verifies the static header is
// still used when no dynamic source sets the same header name — e.g. an
// autonomous run with no caller token, or an unrelated static header.
func TestStaticHeaders_ApplyWhenNoDynamicSource(t *testing.T) {
t.Parallel()
var capturedAuth, capturedOther string

srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
capturedAuth = r.Header.Get("Authorization")
capturedOther = r.Header.Get("X-Other")
w.WriteHeader(http.StatusOK)
}))
defer srv.Close()

rt := &headerRoundTripper{
base: newTestTransport(t),
headers: map[string]string{
"Authorization": "Bearer static",
"X-Other": "static-other",
},
headerProvider: func(context.Context) map[string]string {
return map[string]string{"X-Different": "dynamic-value"}
},
}

req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL, nil)
resp, err := rt.RoundTrip(req)
if err != nil {
t.Fatalf("RoundTrip failed: %v", err)
}
resp.Body.Close()

if capturedAuth != "Bearer static" {
t.Errorf("Authorization: got %q, want %q", capturedAuth, "Bearer static")
}
if capturedOther != "static-other" {
t.Errorf("X-Other: got %q, want %q", capturedOther, "static-other")
}
}
Loading