Skip to content

mcp: prevent a failed request from breaking a stateless streamable client - #1353

Open
jeongukjae wants to merge 1 commit into
modelcontextprotocol:mainfrom
jeongukjae:fix-stateless-poisoning
Open

jeongukjae wants to merge 1 commit into
modelcontextprotocol:mainfrom
jeongukjae:fix-stateless-poisoning

Conversation

@jeongukjae

Copy link
Copy Markdown

Hi all,

I noticed one failed request permanently broke the streamable client connection. The later call failed with connection closed, or client is closing. I strongly believe 1) that the request for stateless client should be independent and 2) #723 fixed all 429 + 5xx cases to make connection broken, but still there's still some gaps like timeout, 4xx, and others.

So what I'm trying to patch in this pr is only making stateless streamable client not to be broken due to previous request. This doesn't change any public interface.

Partially address #683

Tested with
// Command reproduce shows a streamable HTTP client bug on stateless
// connections (no session ID): one failed call or notification permanently
// breaks the connection, so every later call fails with "connection closed".
//
// Run it from the repository root:
//
//	go run ./tmp/main
package main

import (
	"bytes"
	"context"
	"errors"
	"fmt"
	"io"
	"net/http"
	"net/http/httptest"
	"os"
	"strings"
	"sync/atomic"
	"time"

	"github.com/modelcontextprotocol/go-sdk/mcp"
)

// testDelays defines the sleep duration for each call
var testDelays = []time.Duration{
	500 * time.Millisecond, // Call 1: fast, should succeed
	3 * time.Second,        // Call 2: slow, will timeout
	500 * time.Millisecond, // Call 3: fast, should succeed but fails due to bug
}

// noDelays is for scenarios where the failure doesn't come from a slow tool.
var noDelays = []time.Duration{0, 0, 0}

func main() {
	scenarios := []struct {
		name string
		run  func() ([]callResult, error)
	}{
		{"call: HTTP client timeout before the response headers", timeoutBeforeHeaders},
		{"call: HTTP client timeout while reading the JSON body", timeoutDuringBody},
		{"call: plain-text 413 from a gateway", callRejected},
		{"call: HTML login page from a gateway", callLoginPage},
		{"notification: plain-text 400 for notifications/progress", notificationRejected},
		{"notification: plain-text 401 for notifications/cancelled", cancelNotificationRejected},
	}

	broken := 0
	for _, s := range scenarios {
		fmt.Printf("=== %s\n", s.name)
		results, err := s.run()
		if err != nil {
			fmt.Printf("    setup failed: %v\n", err)
			os.Exit(2)
		}
		if !reportResults(results) {
			broken++
		}
	}
	fmt.Printf("\n%d of %d scenarios broke the connection\n", broken, len(scenarios))
	if broken > 0 {
		os.Exit(1)
	}
}

// timeoutBeforeHeaders: call #2 is slower than the HTTP client timeout, which
// fires before the response headers arrive.
func timeoutBeforeHeaders() ([]callResult, error) {
	var callCount atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, testDelays), false, nil)
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 2*time.Second)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	return performCallSequence(session, nil, nil), nil
}

// timeoutDuringBody: a proxy sends the response headers early, so the HTTP
// client timeout fires while the JSON body of call #2 is still on its way.
func timeoutDuringBody() ([]callResult, error) {
	var callCount, faulted atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, testDelays), true, gateway(`"tools/call"`, 0, &faulted, earlyHeaders))
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 2*time.Second)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	return performCallSequence(session, nil, nil), nil
}

// callRejected: a gateway rejects call #2 with a plain-text 413.
func callRejected() ([]callResult, error) {
	var callCount, faulted atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, noDelays), false, gateway(`"tools/call"`, 2, &faulted, reject(http.StatusRequestEntityTooLarge)))
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 0)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	return performCallSequence(session, nil, nil), nil
}

// callLoginPage: a gateway answers call #2 with an HTML login page.
func callLoginPage() ([]callResult, error) {
	var callCount, faulted atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, noDelays), false, gateway(`"tools/call"`, 2, &faulted, loginPage))
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 0)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	return performCallSequence(session, nil, nil), nil
}

// notificationRejected: a gateway rejects a notification
// (notifications/progress) with a plain-text 400, before any call is made.
func notificationRejected() ([]callResult, error) {
	var callCount, faulted atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, noDelays), false, gateway(`"notifications/progress"`, 1, &faulted, reject(http.StatusBadRequest)))
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 0)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	err = session.NotifyProgress(context.Background(), &mcp.ProgressNotificationParams{ProgressToken: "token", Progress: 1})
	if err != nil {
		fmt.Printf("    Notification: FAILED - %v\n", err)
	} else {
		fmt.Printf("    Notification: SUCCESS\n")
	}
	if faulted.Load() != 1 {
		return nil, errors.New("the gateway did not see notifications/progress")
	}

	return performCallSequence(session, nil, nil), nil
}

// cancelNotificationRejected: call #2 times out, so the client sends
// notifications/cancelled, and a gateway rejects that notification with a
// plain-text 401.
func cancelNotificationRejected() ([]callResult, error) {
	var callCount, faulted atomic.Int32
	httpServer := startServer(createDelayTool(&callCount, testDelays), false, gateway(`"notifications/cancelled"`, 1, &faulted, reject(http.StatusUnauthorized)))
	defer httpServer.Close()
	session, err := connect(httpServer.URL, 0)
	if err != nil {
		return nil, err
	}
	defer session.Close()

	var waitErr error
	results := performCallSequence(session, map[int]time.Duration{2: time.Second}, func(num int) {
		if num == 2 {
			// The client sends notifications/cancelled in the background.
			if !waitFor(func() bool { return faulted.Load() == 1 }) {
				waitErr = errors.New("the gateway did not see notifications/cancelled")
			}
			time.Sleep(200 * time.Millisecond) // let the client handle the rejection
		}
	})
	return results, waitErr
}

// createDelayTool creates an MCP server with a tool that sleeps for delays[n-1] on call n.
func createDelayTool(callCount *atomic.Int32, delays []time.Duration) *mcp.Server {
	server := mcp.NewServer(&mcp.Implementation{
		Name:    "test-server",
		Version: "1.0.0",
	}, nil)

	tool := &mcp.Tool{
		Name:        "delay_tool",
		Description: "Tool with configurable delays for testing",
	}

	handler := func(ctx context.Context, req *mcp.CallToolRequest, args any) (*mcp.CallToolResult, any, error) {
		callNum := int(callCount.Add(1))
		delay := delays[0] // default
		if callNum <= len(delays) {
			delay = delays[callNum-1]
		}

		time.Sleep(delay)

		return &mcp.CallToolResult{
			Content: []mcp.Content{
				&mcp.TextContent{
					Text: fmt.Sprintf("Call #%d completed", callNum),
				},
			},
		}, nil, nil
	}

	mcp.AddTool(server, tool, handler)
	return server
}

// startServer starts an HTTP test server with a stateless MCP handler. If
// front is set, it wraps the handler, simulating a proxy or gateway in front
// of the server.
func startServer(server *mcp.Server, jsonResponse bool, front func(http.Handler) http.Handler) *httptest.Server {
	var handler http.Handler = mcp.NewStreamableHTTPHandler(func(req *http.Request) *mcp.Server {
		return server
	}, &mcp.StreamableHTTPOptions{Stateless: true, JSONResponse: jsonResponse})
	if front != nil {
		handler = front(handler)
	}
	return httptest.NewServer(handler)
}

// gateway returns a front handler that calls fault for the nth request whose
// JSON-RPC body contains match (every matching request if n is 0), and
// forwards everything else. faulted counts the requests passed to fault.
func gateway(match string, n int32, faulted *atomic.Int32, fault func(w http.ResponseWriter, r *http.Request, next http.Handler)) func(http.Handler) http.Handler {
	return func(next http.Handler) http.Handler {
		var seen atomic.Int32
		return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
			body, _ := io.ReadAll(r.Body)
			r.Body = io.NopCloser(bytes.NewReader(body))
			if strings.Contains(string(body), match) {
				if k := seen.Add(1); n == 0 || k == n {
					faulted.Add(1)
					fault(w, r, next)
					return
				}
			}
			next.ServeHTTP(w, r)
		})
	}
}

// reject answers with a plain-text error status, like an auth gateway.
func reject(status int) func(http.ResponseWriter, *http.Request, http.Handler) {
	return func(w http.ResponseWriter, r *http.Request, next http.Handler) {
		http.Error(w, "rejected by gateway", status)
	}
}

// loginPage answers with an HTML page, like an auth proxy showing a login page.
func loginPage(w http.ResponseWriter, r *http.Request, next http.Handler) {
	w.Header().Set("Content-Type", "text/html")
	io.WriteString(w, "<html><body>Please log in</body></html>")
}

// earlyHeaders sends "200 OK" headers before the server has answered, as some
// proxies do, then copies the server's response body.
func earlyHeaders(w http.ResponseWriter, r *http.Request, next http.Handler) {
	w.Header().Set("Content-Type", "application/json")
	w.WriteHeader(http.StatusOK)
	w.(http.Flusher).Flush()
	rec := httptest.NewRecorder()
	next.ServeHTTP(rec, r)
	w.Write(rec.Body.Bytes())
}

// connect connects a client to serverURL and checks that the connection is
// stateless. A clientTimeout of 0 means no HTTP timeout.
func connect(serverURL string, clientTimeout time.Duration) (*mcp.ClientSession, error) {
	client := mcp.NewClient(&mcp.Implementation{
		Name:    "test-client",
		Version: "1.0.0",
	}, nil)

	transport := &mcp.StreamableClientTransport{
		Endpoint:   serverURL,
		HTTPClient: &http.Client{Timeout: clientTimeout},
	}

	session, err := client.Connect(context.Background(), transport, nil)
	if err != nil {
		return nil, fmt.Errorf("failed to connect: %v", err)
	}
	if session.ID() != "" {
		session.Close()
		return nil, fmt.Errorf("session ID %q: want a stateless connection", session.ID())
	}
	fmt.Printf("    Connected: protocol %s, no session ID\n", session.InitializeResult().ProtocolVersion)
	return session, nil
}

// callResult represents the result of a CallTool invocation
type callResult struct {
	num    int
	err    error
	result *mcp.CallToolResult
}

// performCallSequence executes three sequential tool calls and returns results.
// timeouts optionally sets a context timeout per call number, and afterCall,
// if set, runs after each call.
func performCallSequence(session *mcp.ClientSession, timeouts map[int]time.Duration, afterCall func(num int)) []callResult {
	results := make([]callResult, 3)

	for i := range 3 {
		ctx := context.Background() // Fresh context for each call
		if d, ok := timeouts[i+1]; ok {
			var cancel context.CancelFunc
			ctx, cancel = context.WithTimeout(ctx, d)
			defer cancel()
		}
		result, err := session.CallTool(ctx, &mcp.CallToolParams{
			Name:      "delay_tool",
			Arguments: map[string]any{},
		})
		results[i] = callResult{num: i + 1, err: err, result: result}
		if afterCall != nil {
			afterCall(i + 1)
		}
	}

	return results
}

// reportResults prints the results and reports whether the connection
// survived, which means call #3 succeeded.
func reportResults(results []callResult) bool {
	for _, r := range results {
		if r.err != nil {
			fmt.Printf("    Call #%d: FAILED - %v\n", r.num, r.err)
		} else {
			fmt.Printf("    Call #%d: SUCCESS\n", r.num)
		}
	}

	if results[2].err != nil {
		fmt.Println("    => BUG: call #3 failed, the connection is broken")
		return false
	}
	fmt.Println("    => OK: call #3 succeeded")
	return true
}

// waitFor reports whether cond became true within 5 seconds.
func waitFor(cond func() bool) bool {
	deadline := time.Now().Add(5 * time.Second)
	for !cond() {
		if time.Now().After(deadline) {
			return false
		}
		time.Sleep(10 * time.Millisecond)
	}
	return true
}

@jeongukjae
jeongukjae force-pushed the fix-stateless-poisoning branch from c3da297 to 8547b3e Compare October 8, 2026 16:00
…ient

One failed call broke the streamable client connection, even with stateless, even after the server recovered.

Signed-off-by: Ukjae Jeong <jeongukjae@gmail.com>
@jeongukjae
jeongukjae force-pushed the fix-stateless-poisoning branch from 8547b3e to d80f8d3 Compare October 9, 2026 15:13

@chrikrah chrikrah left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

@jeongukjae your repro breaks the connection in 5 of 6 scenarios with mcp/streamable.go and internal/jsonrpc2/conn.go at the merge base, and in 0 of 6 on your tree. I would merge the fix once the #1118 question below is settled.

$ go run ./tmp/main   # your <details> program, go1.25.3, tree at d80f8d3
0 of 6 scenarios broke the connection
$ go run ./tmp/main   # both source files at 3c08147 (merge base)
5 of 6 scenarios broke the connection
$ go test ./mcp/ ./internal/jsonrpc2/ && go vet ./mcp/ ./internal/jsonrpc2/   # d80f8d3
ok  	github.com/modelcontextprotocol/go-sdk/mcp	32.303s
ok  	github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2	0.004s
$ sed -i 's/if forCall != nil \&\& c.SessionID() == "" {/if false \&\& forCall != nil {/' mcp/streamable.go   # d80f8d3, 7 matches
$ # plus the content-type NonFatal return removed
$ go vet ./mcp/ && go test -count=1 ./mcp/
ok  	github.com/modelcontextprotocol/go-sdk/mcp	32.425s
$ go run ./tmp/main   # same tree
2 of 6 scenarios broke the connection

non-blocking: the table only reaches the Write path. The seven guarded branches in handleJSON, handleSSE and processStream and the content-type return have no test. The HTML login page and the JSON-body timeout from your repro would each fit as a row.

non-blocking: NonFatal sits beside ErrRejected and does not replace it. On #683 findleyr proposed a flaky mode in jsonrpc2 that removes ErrRejected.

@maciej-kisiel these six scenarios answer your April question on #683 about concrete errors. @guglielmo-san the PR inverts three rows you added in #1118 (noprotocolerrorbody=1, plain-text 400, bare 404). Is a bare 404 on a 2026-07-28 connection still meant to end it?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants