diff --git a/mcp/sse.go b/mcp/sse.go index 0a5bf7fe..f5e90b82 100644 --- a/mcp/sse.go +++ b/mcp/sse.go @@ -437,9 +437,11 @@ func (c *SSEClientTransport) Connect(ctx context.Context) (Connection, error) { maxEventSize = DefaultMaxEventSize } + // Reuse the scanner so events buffered after the endpoint are not lost. + events := scanEventsLimited(resp.Body, maxEventSize) msgEndpoint, err := func() (*url.URL, error) { var evt Event - for evt, err = range scanEventsLimited(resp.Body, maxEventSize) { + for evt, err = range events { break } if err != nil { @@ -468,7 +470,7 @@ func (c *SSEClientTransport) Connect(ctx context.Context) (Connection, error) { go func() { defer s.Close() // close the transport when the GET exits - for evt, err := range scanEventsLimited(resp.Body, maxEventSize) { + for evt, err := range events { if err != nil { return } diff --git a/mcp/sse_test.go b/mcp/sse_test.go index de10eb30..d9b7a2b2 100644 --- a/mcp/sse_test.go +++ b/mcp/sse_test.go @@ -13,10 +13,13 @@ import ( "net/http" "net/http/httptest" "slices" + "strings" "sync/atomic" "testing" + "time" "github.com/google/go-cmp/cmp" + "github.com/modelcontextprotocol/go-sdk/jsonrpc" ) // TestSSEServerTransport_SupportedVersions verifies that the deprecated @@ -432,3 +435,49 @@ func TestSSEServerTransportMaxRequestBodyBytes(t *testing.T) { }) } } + +func TestSSEClientTransportBufferedEvents(t *testing.T) { + const endpoint = "event: endpoint\ndata: /messages\n\n" + const message = "event: message\ndata: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/tools/list_changed\"}\n\n" + for _, coalesced := range []bool{false, true} { + t.Run(fmt.Sprintf("coalesced=%t", coalesced), func(t *testing.T) { + reader, writer := io.Pipe() + defer reader.Close() + defer writer.Close() + prefix := endpoint + if coalesced { + prefix += message + } + transport := &SSEClientTransport{ + Endpoint: "http://example.com/sse", + HTTPClient: &http.Client{Transport: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": {"text/event-stream"}}, + Body: struct { + io.Reader + io.Closer + }{io.MultiReader(strings.NewReader(prefix), reader), reader}, + }, nil + })}, + } + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + conn, err := transport.Connect(ctx) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if !coalesced { + go func() { _, _ = io.WriteString(writer, message) }() + } + msg, err := conn.Read(ctx) + if err != nil { + t.Fatalf("Read: %v", err) + } + if req, ok := msg.(*jsonrpc.Request); !ok || req.Method != "notifications/tools/list_changed" { + t.Fatalf("Read = %v, want tools/list_changed notification", msg) + } + }) + } +}