diff --git a/mcp/streamable.go b/mcp/streamable.go index b2190fa0..b761f9cf 100644 --- a/mcp/streamable.go +++ b/mcp/streamable.go @@ -2446,10 +2446,10 @@ func (c *streamableClientConn) setMCPHeaders(req *http.Request, msg jsonrpc.Mess } if pv := protocolVersionFromMessage(msg); pv != "" { req.Header.Set(protocolVersionHeader, pv) - } else if pv := protocolVersionFromContext(req.Context()); pv != "" { - req.Header.Set(protocolVersionHeader, pv) } else if c.initializedResult != nil { req.Header.Set(protocolVersionHeader, c.initializedResult.ProtocolVersion) + } else if pv := protocolVersionFromContext(req.Context()); pv != "" { + req.Header.Set(protocolVersionHeader, pv) } if c.sessionID != "" { req.Header.Set(sessionIDHeader, c.sessionID) diff --git a/mcp/streamable_client_test.go b/mcp/streamable_client_test.go index a5957bf6..389f87ec 100644 --- a/mcp/streamable_client_test.go +++ b/mcp/streamable_client_test.go @@ -1384,6 +1384,7 @@ func TestStreamableClientConnSetMCPHeaders_ProtocolVersion(t *testing.T) { name string initializedResult *InitializeResult msg jsonrpc.Message + ctxVersion string want string }{ { @@ -1422,6 +1423,30 @@ func TestStreamableClientConnSetMCPHeaders_ProtocolVersion(t *testing.T) { msg: req(1, methodListTools, &ListToolsParams{}), want: "", }, + { + // A process that is both a server and a client (a proxy) carries the + // inbound request's version on the context of the outgoing call. The + // version this connection negotiated must win over it. + name: "initializedResult preferred over request context", + initializedResult: &InitializeResult{ProtocolVersion: protocolVersion20251125}, + msg: req(1, methodListTools, &ListToolsParams{}), + ctxVersion: protocolVersion20260728, + want: protocolVersion20251125, + }, + { + name: "request context used when initializedResult unset", + initializedResult: nil, + msg: req(1, methodListTools, &ListToolsParams{}), + ctxVersion: protocolVersion20260728, + want: protocolVersion20260728, + }, + { + name: "message meta preferred over request context", + initializedResult: nil, + msg: req(1, methodListTools, &ListToolsParams{Meta: Meta{MetaKeyProtocolVersion: protocolVersion20251125}}), + ctxVersion: protocolVersion20260728, + want: protocolVersion20251125, + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -1430,6 +1455,9 @@ func TestStreamableClientConnSetMCPHeaders_ProtocolVersion(t *testing.T) { if err != nil { t.Fatal(err) } + if tt.ctxVersion != "" { + httpReq = httpReq.WithContext(context.WithValue(httpReq.Context(), protocolVersionContextKey{}, tt.ctxVersion)) + } if err := conn.setMCPHeaders(httpReq, tt.msg); err != nil { t.Fatalf("setMCPHeaders: %v", err) }