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
4 changes: 3 additions & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ go 1.26.4

require (
github.com/BurntSushi/toml v1.6.0
github.com/modelcontextprotocol/go-sdk v1.6.1
github.com/modelcontextprotocol/go-sdk v1.7.0
github.com/spf13/cobra v1.10.2
golang.org/x/term v0.43.0
)
Expand Down Expand Up @@ -43,7 +43,9 @@ require (
go.opentelemetry.io/proto/otlp v1.10.0 // indirect
golang.org/x/net v0.55.0 // indirect
golang.org/x/oauth2 v0.36.0 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/sys v0.45.0 // indirect
golang.org/x/time v0.15.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/grpc v1.82.1 // indirect
Expand Down
8 changes: 6 additions & 2 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,8 @@ github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/modelcontextprotocol/go-sdk v1.6.1 h1:0zOSupjKUxPKSocPT1Wtago+mUHU2/uZ4xSOY0FGReU=
github.com/modelcontextprotocol/go-sdk v1.6.1/go.mod h1:kzm3kzFL1/+AziGOE0nUs3gvPoNxMCvkxokMkuFapXQ=
github.com/modelcontextprotocol/go-sdk v1.7.0 h1:yqjY2dsbKAC0LSuWZVBMrHgiG8ukXv6NRo0JiALay44=
github.com/modelcontextprotocol/go-sdk v1.7.0/go.mod h1:dL7u98E/zjJTGzEq+j30jQ8K2k1mb6LeAH4inEcSGts=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
Expand Down Expand Up @@ -84,12 +84,16 @@ golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8=
golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww=
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY=
golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4=
golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk=
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc=
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
golang.org/x/tools v0.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c=
golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
Expand Down
25 changes: 24 additions & 1 deletion internal/mcp/connect_timeout_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package mcp
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"testing"
Expand All @@ -24,9 +25,31 @@ func newMinimalTestServer(t *testing.T) *httptest.Server {
w.WriteHeader(http.StatusOK)
return
case http.MethodPost:
var req map[string]interface{}
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
http.Error(w, fmt.Sprintf("failed to decode request: %v", err), http.StatusBadRequest)
return
}
method, _ := req["method"].(string)
if method == "server/discover" {
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"jsonrpc": "2.0",
"id": req["id"],
"error": map[string]interface{}{
"code": -32601,
"message": `method not found: "server/discover"`,
},
})
return
}
if method == "notifications/initialized" {
w.WriteHeader(http.StatusAccepted)
return
}
resp := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down
103 changes: 90 additions & 13 deletions internal/mcp/http_connection_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,41 @@ import (
"github.com/stretchr/testify/require"
)

// handleDiscoveryProbe answers the SDK's "server/discover" probe (introduced
// in go-sdk v1.7.0 for the 2026-07-28 stateless protocol) with a JSON-RPC
// "method not found" error and acknowledges "notifications/initialized" with
// a 202 Accepted, so that handlers written for the pre-1.7.0 SDK still fall
// through to the initialize response for the streamable transport. It
// returns true if the request was fully handled by this helper.
func handleDiscoveryProbe(w http.ResponseWriter, r *http.Request, req map[string]interface{}) bool {
method, _ := req["method"].(string)
switch method {
case "server/discover":
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"jsonrpc": "2.0",
"id": req["id"],
"error": map[string]interface{}{
"code": -32601,
"message": `method not found: "server/discover"`,
},
})
return true
case "notifications/initialized":
w.WriteHeader(http.StatusAccepted)
return true
}
return false
}

// decodeJSONRPCRequest reads and decodes the JSON-RPC request body without
// consuming it for later use by the caller.
func decodeJSONRPCRequest(r *http.Request) map[string]interface{} {
var req map[string]interface{}
_ = json.NewDecoder(r.Body).Decode(&req)
return req
}

// TestNewHTTPConnection_WithCustomHeaders tests that custom headers are injected into the
// SDK-managed Streamable HTTP transport (not bypassed to plain JSON-RPC).
func TestNewHTTPConnection_WithCustomHeaders(t *testing.T) {
Expand All @@ -30,10 +65,15 @@ func TestNewHTTPConnection_WithCustomHeaders(t *testing.T) {
assert.Equal("test-auth-token", r.Header.Get("Authorization"))
assert.Equal("custom-value", r.Header.Get("X-Custom-Header"))

req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}

// Return a valid initialize response
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{
Expand Down Expand Up @@ -82,9 +122,13 @@ func TestNewHTTPConnection_WithoutHeaders_FallbackSequence(t *testing.T) {
testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Accept all POST requests with valid JSON-RPC response
if r.Method == "POST" {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{
Expand Down Expand Up @@ -279,11 +323,15 @@ func TestHTTPConnection_SSEFormattedResponse(t *testing.T) {

// Create test server that returns SSE-formatted initialize response
testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
// Return SSE-formatted response (like Tavily backend)
response := `event: message
data: {"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"test-server","version":"1.0.0"}}}

`
id, _ := json.Marshal(req["id"])
response := "event: message\ndata: " +
`{"jsonrpc":"2.0","id":` + string(id) + `,"result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"test-server","version":"1.0.0"}}}` +
"\n\n"
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Mcp-Session-Id", "sse-session-456")
w.WriteHeader(http.StatusOK)
Expand Down Expand Up @@ -314,9 +362,13 @@ func TestHTTPConnection_NoSessionIDInResponse(t *testing.T) {

// Create test server that doesn't return Mcp-Session-Id header
testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{
Expand Down Expand Up @@ -359,9 +411,14 @@ func TestNewHTTPConnection_HeadersPropagation(t *testing.T) {
receivedHeaders["X-Custom-1"] = r.Header.Get("X-Custom-1")
receivedHeaders["X-Custom-2"] = r.Header.Get("X-Custom-2")

req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}

response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down Expand Up @@ -401,9 +458,13 @@ func TestNewHTTPConnection_EmptyHeaders(t *testing.T) {
testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Accept POST requests with valid JSON-RPC response
if r.Method == "POST" {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down Expand Up @@ -443,9 +504,13 @@ func TestNewHTTPConnection_NilHeaders(t *testing.T) {

testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method == "POST" && r.URL.Path == "/" {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down Expand Up @@ -478,11 +543,15 @@ func TestNewHTTPConnection_HTTPClientTimeoutUnset(t *testing.T) {

// Create test server with delayed response.
testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
time.Sleep(50 * time.Millisecond)

response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down Expand Up @@ -532,9 +601,13 @@ func TestNewHTTPConnection_GettersAfterCreation(t *testing.T) {
require := require.New(t)

testServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{"name": "test"},
Expand Down Expand Up @@ -598,9 +671,13 @@ func TestNewHTTPConnection_StreamableTransport_BadSSEEndpoint(t *testing.T) {
}

// POST: respond with a valid JSON-RPC initialize result.
req := decodeJSONRPCRequest(r)
if handleDiscoveryProbe(w, r, req) {
return
}
response := map[string]interface{}{
"jsonrpc": "2.0",
"id": 1,
"id": req["id"],
"result": map[string]interface{}{
"protocolVersion": "2024-11-05",
"serverInfo": map[string]interface{}{
Expand Down
26 changes: 25 additions & 1 deletion internal/mcp/http_transport_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1373,6 +1373,18 @@ func TestMaxRetriesSentinelCanary(t *testing.T) {
var req map[string]interface{}
_ = json.NewDecoder(r.Body).Decode(&req)
method, _ := req["method"].(string)
if method == "server/discover" {
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"jsonrpc": "2.0",
"id": req["id"],
"error": map[string]interface{}{
"code": -32601,
"message": `method not found: "server/discover"`,
},
})
return
}
if method == "initialize" {
resp := map[string]interface{}{
"jsonrpc": "2.0",
Expand All @@ -1388,7 +1400,7 @@ func TestMaxRetriesSentinelCanary(t *testing.T) {
close(initializeDone)
return
}
w.WriteHeader(http.StatusOK)
w.WriteHeader(http.StatusAccepted)
}
}))
defer srv.Close()
Expand Down Expand Up @@ -1477,6 +1489,18 @@ func TestDisableStandaloneSSECanary(t *testing.T) {
var req map[string]interface{}
_ = json.NewDecoder(r.Body).Decode(&req)
method, _ := req["method"].(string)
if method == "server/discover" {
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"jsonrpc": "2.0",
"id": req["id"],
"error": map[string]interface{}{
"code": -32601,
"message": `method not found: "server/discover"`,
},
})
return
}
if method == "initialize" {
resp := map[string]interface{}{
"jsonrpc": "2.0",
Expand Down
2 changes: 1 addition & 1 deletion internal/server/register_tools_from_backend_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -236,7 +236,7 @@ func TestRegisterToolsFromBackend_BackendError(t *testing.T) {
// Attempt to register tools should fail with backend error
err = us.registerToolsFromBackend("error-backend")
require.Error(err, "Should fail when backend returns error")
require.ErrorContains(err, "failed to list tools", "Error should mention failed to list tools")
require.ErrorContains(err, "backend error listing tools", "Error should mention the backend error")
require.ErrorContains(err, "unable to list tools", "Error should include backend error message")
}

Expand Down
19 changes: 19 additions & 0 deletions internal/server/session_auto_init.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,14 @@ const autoInitProtocolVersion = "2025-11-25"
// autoInitClientInfo is the JSON snippet for the clientInfo field in the initialize request.
const autoInitClientInfo = `{"name":"mcpg-auto-init","version":"1.0"}`

// mcpProtocolVersionHeader is the HTTP header carrying the negotiated MCP protocol version.
const mcpProtocolVersionHeader = "Mcp-Protocol-Version"

// statelessProtocolVersion is the first MCP protocol version (SEP-2577/SEP-2575)
// that intentionally omits Mcp-Session-Id on stateless requests. Auto-init must
// not treat these requests as missing a legacy handshake.
const statelessProtocolVersion = "2026-07-28"

// WrapWithSessionAutoInit wraps an MCP streamable HTTP handler to automatically
// initialize sessions for clients that send tools/call before completing the MCP
// session handshake.
Expand Down Expand Up @@ -46,6 +54,17 @@ func WrapWithSessionAutoInit(streamableHandler http.Handler) http.Handler {
return
}

// SDK v1.7.0+ defaults to the stateless "2026-07-28" protocol, under
// which requests intentionally omit Mcp-Session-Id (see SEP-2577).
// Auto-init only targets legacy stateful clients (e.g. Gemini CLI
// v0.37.x) that skip the initialize handshake, so bypass it here to
// avoid allocating an unused stateful session for every stateless
// tools/call.
if r.Header.Get(mcpProtocolVersionHeader) >= statelessProtocolVersion {
streamableHandler.ServeHTTP(w, r)
return
}

// Peek at the request body to detect tools/call.
bodyBytes, err := readAndRestoreRequestBody(r)
if err != nil || len(bodyBytes) == 0 {
Expand Down
Loading
Loading