diff --git a/mcp/streamable.go b/mcp/streamable.go index 1f07ea61..f642db1d 100644 --- a/mcp/streamable.go +++ b/mcp/streamable.go @@ -2403,13 +2403,12 @@ func (c *streamableClientConn) setMCPHeaders(req *http.Request, msg jsonrpc.Mess } } } - switch { - case protocolVersionFromMessage(msg) != "": - req.Header.Set(protocolVersionHeader, protocolVersionFromMessage(msg)) - case c.initializedResult != nil: + 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) - case protocolVersionFromContext(req.Context()) != "": - req.Header.Set(protocolVersionHeader, protocolVersionFromContext(req.Context())) } if c.sessionID != "" { req.Header.Set(sessionIDHeader, c.sessionID) @@ -2424,7 +2423,7 @@ func (c *streamableClientConn) setMCPHeaders(req *http.Request, msg jsonrpc.Mess // nil msg. func protocolVersionFromMessage(msg jsonrpc.Message) string { req, ok := msg.(*jsonrpc.Request) - if !ok { + if !ok || req == nil { return "" } meta := extractRequestMeta(req.Params) diff --git a/mcp/streamable_client_test.go b/mcp/streamable_client_test.go index a92f2e4e..c8ab0669 100644 --- a/mcp/streamable_client_test.go +++ b/mcp/streamable_client_test.go @@ -1386,6 +1386,12 @@ func TestStreamableClientConnSetMCPHeaders_ProtocolVersion(t *testing.T) { msg jsonrpc.Message want string }{ + { + name: "nil checked", + initializedResult: nil, + msg: (*jsonrpc.Request)(nil), + want: "", + }, { name: "message meta wins when initializedResult unset", initializedResult: nil,