From f17c17421fb43c99a6d908658e544f48f7625409 Mon Sep 17 00:00:00 2001 From: Sean Kenny Date: Mon, 24 Aug 2026 22:01:59 +0200 Subject: [PATCH 1/2] mcp: keep a long-running POST stream visibly alive The SSE response to a POST writes nothing until the request it carries completes. A tool call that runs for minutes therefore produces no bytes at all in the meantime, not even the response headers, so a client that applies a first-byte or idle timeout cannot tell a working call apart from a dead connection and hangs up on it. Intermediaries that buffer idle responses have the same problem. The standalone GET stream already handles this (#410) by committing the headers and writing an SSE comment up front. Do the same for streams created by servePOST, after a short delay: committing the headers fixes the HTTP status, and a SEP-2575 protocol-level error must still be able to set its own. Those errors are produced without any I/O, so a stream still silent after the delay is a genuinely long-running call. If the stream has already been flushed, such an error is delivered as an ordinary SSE event instead, which is the only option once the status is fixed. The flush holds the stream mutex, serialising it with deliverLocked and close (the other writers to s.w) and with release, which clears the writer and the per-request headersFlushed flag so a resumed stream can still set an error status. Fixes #1155 --- mcp/streamable.go | 81 ++++- mcp/streamable_earlyflush_test.go | 472 ++++++++++++++++++++++++++++++ 2 files changed, 551 insertions(+), 2 deletions(-) create mode 100644 mcp/streamable_earlyflush_test.go diff --git a/mcp/streamable.go b/mcp/streamable.go index 6ed5d3d3..8516b88c 100644 --- a/mcp/streamable.go +++ b/mcp/streamable.go @@ -1012,6 +1012,10 @@ type stream struct { // the duration of the subscription, and act as the target for // out-of-band notifications routed through this connection. isListen bool + + // headersFlushed records that the response headers have been committed, so + // the HTTP status can no longer be changed. See flushEarlyAfter. + headersFlushed bool } // close sends a 'close' event to the client (if protocolVersion >= 2025-11-25 @@ -1026,6 +1030,7 @@ func (s *stream) close(reconnectAfter time.Duration) { return // stream not connected or already closed } if s.protocolVersion >= protocolVersion20251125 && reconnectAfter > 0 { + s.headersFlushed = true reconnectStr := strconv.FormatInt(reconnectAfter.Milliseconds(), 10) if _, err := writeEvent(s.w, Event{ Name: "close", @@ -1044,7 +1049,54 @@ func (s *stream) release() { s.mu.Lock() defer s.mu.Unlock() s.w = nil - s.done = nil // may already be nil, if the stream is done or closed + s.done = nil // may already be nil, if the stream is done or closed + s.headersFlushed = false // per HTTP request; the stream object can be reused +} + +// earlyFlushDelay is how long an SSE response to a POST may stay completely +// silent before the server commits its headers and writes a keep-alive +// comment. +// +// It is a delay rather than an immediate flush because committing the headers +// fixes the HTTP status, and a protocol-level error must still be able to set +// its own (SEP-2575, see extractErrorStatus). Those errors are produced +// without any I/O, so anything still silent after this delay is a genuinely +// long-running call. +const earlyFlushDelay = 1 * time.Second + +// flushEarlyAfter commits the response headers and writes an SSE comment once +// the stream has been silent for d, so that a long-running call does not look +// like a dead connection to clients or intermediaries that apply a first-byte +// or idle timeout. +// +// Writing a comment rather than only flushing is deliberate: on HTTP/2 a proxy +// may hold the HEADERS frame until a DATA frame arrives, so a Flush alone can +// still leave the client with nothing. Comment lines are ignored by clients +// per the SSE spec. This is the same mechanism used for the standalone stream +// in acquireStream. +// +// Holding s.mu is what makes this safe: it serialises with deliverLocked and +// close (the other writers to s.w) and with release, which clears it. +func (s *stream) flushEarlyAfter(ctx context.Context, d time.Duration) { + t := time.NewTimer(d) + defer t.Stop() + select { + case <-t.C: + case <-ctx.Done(): + return + } + + s.mu.Lock() + defer s.mu.Unlock() + if s.w == nil || s.headersFlushed { + return + } + s.headersFlushed = true + s.w.WriteHeader(http.StatusOK) + fmt.Fprint(s.w, ": ok\n\n") + // Flushing is best-effort: a ResponseWriter that cannot flush simply keeps + // the previous behavior. + _ = http.NewResponseController(s.w).Flush() } // extractErrorStatus reports the HTTP status to send when the given @@ -1111,7 +1163,13 @@ func (s *stream) deliverLocked(data []byte, eventID string, responseTo jsonrpc.I // SEP-2575 protocol-level error override: write the error as a raw // JSON-RPC response with the spec-mandated HTTP status, bypassing any // SSE framing. - if overrideStatus != 0 { + // + // Only possible while the headers are uncommitted. If the stream has + // already been flushed to keep a long call alive (see flushEarlyAfter), + // the status is fixed at 200 and the error is delivered as an ordinary + // SSE event instead. + if overrideStatus != 0 && !s.headersFlushed { + s.headersFlushed = true s.w.Header().Set("Content-Type", "application/json") s.w.WriteHeader(overrideStatus) if _, err := s.w.Write(data); err != nil { @@ -1138,12 +1196,14 @@ func (s *stream) deliverLocked(data []byte, eventID string, responseTo jsonrpc.I return done, err } } + s.headersFlushed = true if _, err := s.w.Write(toWrite); err != nil { return done, err } } } else { // SSE mode: write event to response writer. + s.headersFlushed = true s.lastIdx++ if _, err := writeEvent(s.w, Event{Name: "message", Data: data, ID: eventID}); err != nil { return done, err @@ -1393,6 +1453,7 @@ func (c *streamableServerConn) acquireStream(ctx context.Context, w http.Respons rc := http.NewResponseController(w) // Ignore returned error as flushing is best-effort. _ = rc.Flush() + s.headersFlushed = true } for _, data := range toReplay { @@ -1404,6 +1465,7 @@ func (c *streamableServerConn) acquireStream(ctx context.Context, w http.Respons if _, err := writeEvent(w, e); err != nil { return nil, nil } + s.headersFlushed = true } if tempStream || s.doneLocked() { @@ -1753,6 +1815,21 @@ func (c *streamableServerConn) servePOST(w http.ResponseWriter, req *http.Reques if _, err := writeEvent(w, e); err != nil { c.logger.Warn(fmt.Sprintf("Writing priming event: %v", err)) } + stream.headersFlushed = true + } + + // The first byte of this response may be minutes away: a tool call runs + // to completion before its result is written, and nothing else is sent + // in the meantime. Clients that apply a first-byte or idle timeout (and + // intermediaries that do the same) cannot tell that silence apart from a + // dead connection, so they hang up on a call that is still running. + // + // Keep the stream visibly alive by committing the headers and writing an + // SSE comment once it has been silent for earlyFlushDelay. See + // flushEarlyAfter. Started after any priming write so the two cannot + // race on s.w. Skip it if priming already committed the headers. + if !stream.headersFlushed { + go stream.flushEarlyAfter(req.Context(), earlyFlushDelay) } } diff --git a/mcp/streamable_earlyflush_test.go b/mcp/streamable_earlyflush_test.go new file mode 100644 index 00000000..ed38f18a --- /dev/null +++ b/mcp/streamable_earlyflush_test.go @@ -0,0 +1,472 @@ +// Copyright 2025 The Go MCP SDK Authors. All rights reserved. +// Use of this source code is governed by an MIT-style +// license that can be found in the LICENSE file. + +package mcp + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2" + "github.com/modelcontextprotocol/go-sdk/jsonrpc" +) + +// TestPOSTStreamFlushesHeadersEarly checks that the SSE response to a POST +// produces its first bytes while the tool is still running, rather than when +// the tool call completes. +// +// Without this, a server whose tool runs for minutes writes nothing at all in +// the meantime, and clients (and intermediaries) that apply a first-byte +// timeout treat the silence as a dead connection and hang up. +func TestPOSTStreamFlushesHeadersEarly(t *testing.T) { + release := make(chan struct{}) + server := NewServer(testImpl, nil) + AddTool(server, &Tool{Name: "slow"}, func(ctx context.Context, req *CallToolRequest, args map[string]any) (*CallToolResult, any, error) { + select { + case <-release: + case <-ctx.Done(): + return nil, nil, ctx.Err() + } + return &CallToolResult{Content: []Content{&TextContent{Text: "done"}}}, nil, nil + }) + + httpServer := httptest.NewServer(NewStreamableHTTPHandler(func(*http.Request) *Server { return server }, nil)) + defer httpServer.Close() + + post := func(t *testing.T, sessionID, body string) *http.Response { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + req, err := http.NewRequestWithContext(ctx, "POST", httpServer.URL, strings.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + if sessionID != "" { + req.Header.Set(sessionIDHeader, sessionID) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("no response headers within the deadline: %v\n"+ + "(the stream writes nothing until the tool returns, so a client's "+ + "first-byte timer fires on a call that is still running)", err) + } + return resp + } + + initResp := post(t, "", `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"1"}}}`) + sessionID := initResp.Header.Get(sessionIDHeader) + io.Copy(io.Discard, initResp.Body) + initResp.Body.Close() + if sessionID == "" { + t.Fatal("no session ID from initialize") + } + notifResp := post(t, sessionID, `{"jsonrpc":"2.0","method":"notifications/initialized"}`) + io.Copy(io.Discard, notifResp.Body) + notifResp.Body.Close() + + callResp := post(t, sessionID, `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"slow","arguments":{}}}`) + defer callResp.Body.Close() + + if got := baseMediaType(callResp.Header.Get("Content-Type")); got != "text/event-stream" { + t.Fatalf("Content-Type = %q, want text/event-stream", got) + } + if callResp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200", callResp.StatusCode) + } + + firstByte := make(chan []byte, 1) + go func() { + buf := make([]byte, 64) + n, err := callResp.Body.Read(buf) + if err != nil && n == 0 { + firstByte <- nil + return + } + firstByte <- bytes.Clone(buf[:n]) + }() + + select { + case got := <-firstByte: + if len(got) == 0 { + t.Fatal("stream closed without writing anything") + } + if !bytes.HasPrefix(got, []byte(":")) { + t.Errorf("first bytes = %q, want an SSE comment", got) + } + case <-time.After(5 * time.Second): + t.Fatal("no bytes written while the tool was still running: the client's first-byte timer would fire here") + } + + close(release) + rest, err := io.ReadAll(callResp.Body) + if err != nil { + t.Fatalf("reading result: %v", err) + } + if !bytes.Contains(rest, []byte("done")) { + t.Errorf("after releasing the tool, body = %q, want it to contain the result", rest) + } +} + +// TestPOSTUnknownMethodKeeps404 checks that a protocol-level MethodNotFound +// produced before the stream exists still returns HTTP 404, not a flushed 200 +// SSE stream. +func TestPOSTUnknownMethodKeeps404(t *testing.T) { + server := NewServer(testImpl, nil) + httpServer := httptest.NewServer(NewStreamableHTTPHandler( + func(*http.Request) *Server { return server }, + &StreamableHTTPOptions{Stateless: true}, + )) + defer httpServer.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodPost, httpServer.URL, strings.NewReader( + `{"jsonrpc":"2.0","id":1,"method":"notamethod","params":{}}`)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + req.Header.Set(protocolVersionHeader, protocolVersion20260728) + + start := time.Now() + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + body, _ := io.ReadAll(resp.Body) + + if elapsed := time.Since(start); elapsed >= earlyFlushDelay { + t.Errorf("unknown method took %v, want well under earlyFlushDelay (%v)", elapsed, earlyFlushDelay) + } + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("status = %d, want 404; body = %s", resp.StatusCode, body) + } + if got := baseMediaType(resp.Header.Get("Content-Type")); got != "application/json" { + t.Errorf("Content-Type = %q, want application/json", got) + } +} + +// TestPOSTProtocolErrorKeepsOverrideStatus checks that a SEP-2575 protocol +// error produced after the stream exists, but before the delayed flush, still +// sets the spec-mandated HTTP status. +func TestPOSTProtocolErrorKeepsOverrideStatus(t *testing.T) { + server := NewServer(testImpl, nil) + httpServer := httptest.NewServer(NewStreamableHTTPHandler( + func(*http.Request) *Server { return server }, + &StreamableHTTPOptions{Stateless: true}, + )) + defer httpServer.Close() + + body, err := json.Marshal(map[string]any{ + "jsonrpc": "2.0", + "id": 1, + "method": "prompts/get", + "params": map[string]any{ + "_meta": map[string]any{ + MetaKeyProtocolVersion: protocolVersion20260728, + MetaKeyClientInfo: map[string]any{"name": "c", "version": "1"}, + MetaKeyClientCapabilities: map[string]any{"sampling": map[string]any{}}, + }, + "name": "no-such-prompt", + }, + }) + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodPost, httpServer.URL, bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + req.Header.Set(protocolVersionHeader, protocolVersion20260728) + req.Header.Set(methodHeader, "prompts/get") + req.Header.Set(nameHeader, "no-such-prompt") + + start := time.Now() + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + gotBody, _ := io.ReadAll(resp.Body) + + if elapsed := time.Since(start); elapsed >= earlyFlushDelay { + t.Errorf("protocol error took %v, want well under earlyFlushDelay (%v)", elapsed, earlyFlushDelay) + } + if resp.StatusCode != http.StatusBadRequest { + t.Fatalf("status = %d, want 400; body = %s", resp.StatusCode, gotBody) + } + if got := baseMediaType(resp.Header.Get("Content-Type")); got != "application/json" { + t.Errorf("Content-Type = %q, want application/json", got) + } + msg, err := jsonrpc2.DecodeMessage(gotBody) + if err != nil { + t.Fatalf("DecodeMessage: %v; body = %s", err, gotBody) + } + jresp, ok := msg.(*jsonrpc.Response) + if !ok || jresp.Error == nil { + t.Fatalf("response is not a JSON-RPC error: %s", gotBody) + } + var jerr *jsonrpc.Error + if !errors.As(jresp.Error, &jerr) { + t.Fatalf("error is not *jsonrpc.Error: %v", jresp.Error) + } + if jerr.Code != jsonrpc.CodeInvalidParams { + t.Errorf("error code = %d, want %d", jerr.Code, jsonrpc.CodeInvalidParams) + } +} + +func TestDeliverLockedOverrideStatusWinsBeforeFlush(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + + s.mu.Lock() + done, err := s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound) + s.mu.Unlock() + if err != nil { + t.Fatal(err) + } + if !done { + t.Fatal("stream not done after the only request was answered") + } + if rec.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", rec.Code) + } + if got := rec.codes; len(got) != 1 || got[0] != http.StatusNotFound { + t.Errorf("WriteHeader calls = %v, want [404]", got) + } + if got := baseMediaType(rec.Header().Get("Content-Type")); got != "application/json" { + t.Errorf("Content-Type = %q, want application/json", got) + } + if !s.headersFlushed { + t.Error("headersFlushed = false, want true") + } +} + +func TestFlushEarlyAfterSkippedOnceHeadersCommitted(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + + s.mu.Lock() + if _, err := s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound); err != nil { + s.mu.Unlock() + t.Fatal(err) + } + s.mu.Unlock() + + s.flushEarlyAfter(context.Background(), 0) + + if rec.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", rec.Code) + } + if got := rec.codes; len(got) != 1 || got[0] != http.StatusNotFound { + t.Errorf("WriteHeader calls = %v, want [404] (early flush must not write again)", got) + } + if bytes.Contains(rec.Body.Bytes(), []byte(": ok")) { + t.Errorf("body = %q, must not contain an SSE keep-alive after a protocol error", rec.Body.Bytes()) + } +} + +func TestDeliverLockedAfterEarlyFlushWritesSSENotStatus(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + + s.flushEarlyAfter(context.Background(), 0) + + s.mu.Lock() + if _, err := s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound); err != nil { + s.mu.Unlock() + t.Fatal(err) + } + s.mu.Unlock() + + if rec.Code != http.StatusOK { + t.Errorf("status = %d, want 200 (headers already committed)", rec.Code) + } + if got := rec.codes; len(got) != 1 || got[0] != http.StatusOK { + t.Errorf("WriteHeader calls = %v, want [200] (no superfluous override status)", got) + } + body := rec.Body.Bytes() + if !bytes.HasPrefix(body, []byte(":")) { + t.Errorf("body = %q, want an SSE comment first", body) + } + if !bytes.Contains(body, []byte("event: message")) { + t.Errorf("body = %q, want the error delivered as an SSE event", body) + } +} + +func TestFlushEarlyAfterCancelled(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + s.flushEarlyAfter(ctx, time.Hour) + + if len(rec.codes) != 0 { + t.Errorf("WriteHeader calls = %v, want none after cancel", rec.codes) + } + if rec.Body.Len() != 0 { + t.Errorf("body = %q, want empty after cancel", rec.Body.Bytes()) + } + if s.headersFlushed { + t.Error("headersFlushed = true, want false") + } +} + +func TestReleaseResetsHeadersFlushed(t *testing.T) { + rec1, w1 := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w1, id) + s.flushEarlyAfter(context.Background(), 0) + if rec1.Code != http.StatusOK { + t.Fatalf("first request status = %d, want 200", rec1.Code) + } + + s.release() + + rec2, w2 := newHeaderCounter() + s.w = w2 + s.done = make(chan struct{}) + s.requests = map[jsonrpc.ID]struct{}{id: {}} + s.mu.Lock() + if _, err := s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound); err != nil { + s.mu.Unlock() + t.Fatal(err) + } + s.mu.Unlock() + if rec2.Code != http.StatusNotFound { + t.Errorf("reused stream status = %d, want 404 (headersFlushed must reset on release)", rec2.Code) + } + if got := rec2.codes; len(got) != 1 || got[0] != http.StatusNotFound { + t.Errorf("WriteHeader calls = %v, want [404]", got) + } +} + +func TestJSONNotificationDoesNotBlockOverrideStatus(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + s.pendingJSONMessages = []json.RawMessage{} + + s.mu.Lock() + if _, err := s.deliverLocked([]byte(`{"jsonrpc":"2.0","method":"notifications/message"}`), "", jsonrpc.ID{}, 0); err != nil { + s.mu.Unlock() + t.Fatal(err) + } + if s.headersFlushed { + s.mu.Unlock() + t.Fatal("buffering a JSON notification committed headers") + } + if _, err := s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound); err != nil { + s.mu.Unlock() + t.Fatal(err) + } + s.mu.Unlock() + + if rec.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", rec.Code) + } +} + +func TestCloseMarksHeadersFlushed(t *testing.T) { + rec, w := newHeaderCounter() + id := jsonrpc2.Int64ID(1) + s := connectedTestStream(w, id) + s.protocolVersion = protocolVersion20251125 + s.close(time.Second) + + s.flushEarlyAfter(context.Background(), 0) + + if bytes.Contains(rec.Body.Bytes(), []byte(": ok")) { + t.Errorf("body = %q, close must not be followed by a keep-alive comment", rec.Body.Bytes()) + } +} + +func TestFlushEarlyAfterSerializesWithDeliverLocked(t *testing.T) { + id := jsonrpc2.Int64ID(1) + for i := 0; i < 50; i++ { + rec, w := newHeaderCounter() + s := connectedTestStream(w, id) + + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + s.flushEarlyAfter(context.Background(), 0) + }() + go func() { + defer wg.Done() + s.mu.Lock() + _, _ = s.deliverLocked(protocolErrorJSON, "", id, http.StatusNotFound) + s.mu.Unlock() + }() + wg.Wait() + + if len(rec.codes) != 1 { + t.Fatalf("iter %d: WriteHeader calls = %v, want exactly one", i, rec.codes) + } + switch rec.codes[0] { + case http.StatusNotFound: + if bytes.Contains(rec.Body.Bytes(), []byte(": ok")) { + t.Fatalf("iter %d: 404 body also contains an SSE comment: %q", i, rec.Body.Bytes()) + } + case http.StatusOK: + if !bytes.Contains(rec.Body.Bytes(), []byte("event: message")) { + t.Fatalf("iter %d: 200 body missing SSE error event: %q", i, rec.Body.Bytes()) + } + default: + t.Fatalf("iter %d: status %d, want 200 or 404", i, rec.codes[0]) + } + } +} + +var protocolErrorJSON = []byte(`{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"nope"}}`) + +func connectedTestStream(w http.ResponseWriter, id jsonrpc.ID) *stream { + return &stream{ + id: "s", + w: w, + done: make(chan struct{}), + requests: map[jsonrpc.ID]struct{}{id: {}}, + lastIdx: -1, + logger: ensureLogger(nil), + } +} + +type headerCounter struct { + *httptest.ResponseRecorder + codes []int +} + +func newHeaderCounter() (*headerCounter, http.ResponseWriter) { + h := &headerCounter{ResponseRecorder: httptest.NewRecorder()} + return h, h +} + +func (w *headerCounter) WriteHeader(code int) { + w.codes = append(w.codes, code) + w.ResponseRecorder.WriteHeader(code) +} From adf998c82362ed5aa48f3a1e9544ca3d73d7340b Mon Sep 17 00:00:00 2001 From: Sean Kenny Date: Mon, 24 Aug 2026 22:09:03 +0200 Subject: [PATCH 2/2] mcp: flush resumed GET streams that have nothing to replay acquireStream only wrote an SSE comment for the standalone GET (s.id == ""). A Last-Event-ID resume of an in-flight POST with an empty replay therefore hung with uncommitted headers, the same first-byte bug as the original POST. Flush once when claiming a live stream whose headers are still uncommitted. Also drop the idle-timeout overclaim (the write is one-shot), take the stream mutex around the priming write, and record implicit WriteHeader(200) in the test helper. --- mcp/streamable.go | 57 +++++---- mcp/streamable_earlyflush_test.go | 194 ++++++++++++++++++++++++++++++ 2 files changed, 226 insertions(+), 25 deletions(-) diff --git a/mcp/streamable.go b/mcp/streamable.go index 8516b88c..bc3fe80b 100644 --- a/mcp/streamable.go +++ b/mcp/streamable.go @@ -1059,24 +1059,30 @@ func (s *stream) release() { // // It is a delay rather than an immediate flush because committing the headers // fixes the HTTP status, and a protocol-level error must still be able to set -// its own (SEP-2575, see extractErrorStatus). Those errors are produced -// without any I/O, so anything still silent after this delay is a genuinely -// long-running call. +// its own (SEP-2575, see extractErrorStatus). Dispatch-time errors finish +// without handler I/O, so they win the race. A handler that later returns +// InvalidParams after real work may miss the HTTP status override; that is +// the tradeoff. const earlyFlushDelay = 1 * time.Second +// writeSSEComment commits 200 and writes an SSE comment so a DATA frame +// follows the HEADERS frame. On HTTP/2 a proxy may hold HEADERS until DATA +// arrives; Flush alone is not enough. Comment lines are ignored by clients +// per the SSE spec. See acquireStream and golang/go#31125. +func writeSSEComment(w http.ResponseWriter) { + w.WriteHeader(http.StatusOK) + fmt.Fprint(w, ": ok\n\n") + _ = http.NewResponseController(w).Flush() +} + // flushEarlyAfter commits the response headers and writes an SSE comment once // the stream has been silent for d, so that a long-running call does not look -// like a dead connection to clients or intermediaries that apply a first-byte -// or idle timeout. +// like a dead connection to clients that apply a first-byte timeout. // -// Writing a comment rather than only flushing is deliberate: on HTTP/2 a proxy -// may hold the HEADERS frame until a DATA frame arrives, so a Flush alone can -// still leave the client with nothing. Comment lines are ignored by clients -// per the SSE spec. This is the same mechanism used for the standalone stream -// in acquireStream. +// This is one-shot, not a keep-alive ping. Idle timeouts after the first byte +// are a different problem; JSON-RPC pings ride the standalone GET stream. // -// Holding s.mu is what makes this safe: it serialises with deliverLocked and -// close (the other writers to s.w) and with release, which clears it. +// Holding s.mu serialises this with deliverLocked, close, and priming. func (s *stream) flushEarlyAfter(ctx context.Context, d time.Duration) { t := time.NewTimer(d) defer t.Stop() @@ -1092,11 +1098,7 @@ func (s *stream) flushEarlyAfter(ctx context.Context, d time.Duration) { return } s.headersFlushed = true - s.w.WriteHeader(http.StatusOK) - fmt.Fprint(s.w, ": ok\n\n") - // Flushing is best-effort: a ResponseWriter that cannot flush simply keeps - // the previous behavior. - _ = http.NewResponseController(s.w).Flush() + writeSSEComment(s.w) } // extractErrorStatus reports the HTTP status to send when the given @@ -1448,11 +1450,7 @@ func (c *streamableServerConn) acquireStream(ctx context.Context, w http.Respons // proxy to forward both frames. See: // https://github.com/golang/go/issues/31125 // https://github.com/caddyserver/caddy/issues/4247 - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, ": ok\n\n") - rc := http.NewResponseController(w) - // Ignore returned error as flushing is best-effort. - _ = rc.Flush() + writeSSEComment(w) s.headersFlushed = true } @@ -1479,6 +1477,14 @@ func (c *streamableServerConn) acquireStream(ctx context.Context, w http.Respons s.done = make(chan struct{}) s.lastIdx = lastIdx s.protocolVersion = protocolVersion + // Same first-byte problem as a hanging POST: a GET that resumes an + // in-flight stream with nothing left to replay would otherwise write + // nothing until the next event. The standalone stream (s.id == "") + // already flushed above; this covers Last-Event-ID resume. + if !s.headersFlushed { + writeSSEComment(s.w) + s.headersFlushed = true + } return s, s.done } @@ -1812,17 +1818,18 @@ func (c *streamableServerConn) servePOST(w http.ResponseWriter, req *http.Reques } stream.lastIdx++ e := Event{Name: "prime", ID: formatEventID(stream.id, stream.lastIdx)} + stream.mu.Lock() if _, err := writeEvent(w, e); err != nil { c.logger.Warn(fmt.Sprintf("Writing priming event: %v", err)) } stream.headersFlushed = true + stream.mu.Unlock() } // The first byte of this response may be minutes away: a tool call runs // to completion before its result is written, and nothing else is sent - // in the meantime. Clients that apply a first-byte or idle timeout (and - // intermediaries that do the same) cannot tell that silence apart from a - // dead connection, so they hang up on a call that is still running. + // in the meantime. Clients that apply a first-byte timeout cannot tell + // that silence apart from a dead connection. // // Keep the stream visibly alive by committing the headers and writing an // SSE comment once it has been silent for earlyFlushDelay. See diff --git a/mcp/streamable_earlyflush_test.go b/mcp/streamable_earlyflush_test.go index ed38f18a..b22dd820 100644 --- a/mcp/streamable_earlyflush_test.go +++ b/mcp/streamable_earlyflush_test.go @@ -405,6 +405,145 @@ func TestCloseMarksHeadersFlushed(t *testing.T) { } } +func TestPOSTStreamPrimeSkipsEarlyFlush(t *testing.T) { + release := make(chan struct{}) + server := NewServer(testImpl, nil) + AddTool(server, &Tool{Name: "slow"}, func(ctx context.Context, req *CallToolRequest, args map[string]any) (*CallToolResult, any, error) { + select { + case <-release: + case <-ctx.Done(): + return nil, nil, ctx.Err() + } + return &CallToolResult{Content: []Content{&TextContent{Text: "done"}}}, nil, nil + }) + httpServer := httptest.NewServer(NewStreamableHTTPHandler( + func(*http.Request) *Server { return server }, + &StreamableHTTPOptions{EventStore: NewMemoryEventStore(nil)}, + )) + defer httpServer.Close() + defer close(release) + + sessionID := handshake(t, httpServer.URL, protocolVersion20251125) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + resp := postJSON(t, ctx, httpServer.URL, sessionID, protocolVersion20251125, + `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"slow","arguments":{}}}`) + defer resp.Body.Close() + + buf := make([]byte, 256) + n, err := resp.Body.Read(buf) + if err != nil && n == 0 { + t.Fatalf("reading primed stream: %v", err) + } + got := buf[:n] + if !bytes.Contains(got, []byte("event: prime")) { + t.Fatalf("first bytes = %q, want a prime event", got) + } + if bytes.Contains(got, []byte(": ok")) { + t.Errorf("primed stream also wrote an early-flush comment: %q", got) + } +} + +func TestGETResumeFlushesHeaders(t *testing.T) { + release := make(chan struct{}) + server := NewServer(testImpl, nil) + AddTool(server, &Tool{Name: "slow"}, func(ctx context.Context, req *CallToolRequest, args map[string]any) (*CallToolResult, any, error) { + select { + case <-release: + case <-ctx.Done(): + return nil, nil, ctx.Err() + } + return &CallToolResult{Content: []Content{&TextContent{Text: "done"}}}, nil, nil + }) + httpServer := httptest.NewServer(NewStreamableHTTPHandler( + func(*http.Request) *Server { return server }, + &StreamableHTTPOptions{EventStore: NewMemoryEventStore(nil)}, + )) + defer httpServer.Close() + defer func() { + select { + case <-release: + default: + close(release) + } + }() + + sessionID := handshake(t, httpServer.URL, protocolVersion20251125) + + postCtx, postCancel := context.WithCancel(context.Background()) + defer postCancel() + postResp := postJSON(t, postCtx, httpServer.URL, sessionID, protocolVersion20251125, + `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"slow","arguments":{}}}`) + + buf := make([]byte, 256) + n, err := postResp.Body.Read(buf) + if err != nil && n == 0 { + t.Fatalf("reading prime: %v", err) + } + eventID := sseField(buf[:n], "id:") + if eventID == "" { + t.Fatalf("prime event missing id: %q", buf[:n]) + } + + postCancel() + postResp.Body.Close() + + var getResp *http.Response + var getCancel context.CancelFunc + defer func() { + if getCancel != nil { + getCancel() + } + }() + deadline := time.Now().Add(2 * time.Second) + for { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, httpServer.URL, nil) + if err != nil { + cancel() + t.Fatal(err) + } + req.Header.Set("Accept", "text/event-stream") + req.Header.Set(sessionIDHeader, sessionID) + req.Header.Set(protocolVersionHeader, protocolVersion20251125) + req.Header.Set(lastEventIDHeader, eventID) + getResp, err = http.DefaultClient.Do(req) + if err != nil { + cancel() + if time.Now().After(deadline) { + t.Fatalf("resume GET: %v", err) + } + continue + } + if getResp.StatusCode == http.StatusConflict { + getResp.Body.Close() + cancel() + if time.Now().After(deadline) { + t.Fatal("resume GET still 409: original POST did not release") + } + time.Sleep(10 * time.Millisecond) + continue + } + getCancel = cancel + break + } + defer getResp.Body.Close() + if getResp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(getResp.Body) + t.Fatalf("resume GET status = %d, want 200; body = %s", getResp.StatusCode, body) + } + + first := make([]byte, 64) + n, err = getResp.Body.Read(first) + if err != nil && n == 0 { + t.Fatalf("resume GET body: %v", err) + } + if !bytes.HasPrefix(first[:n], []byte(":")) { + t.Errorf("resume GET first bytes = %q, want an SSE comment", first[:n]) + } +} + func TestFlushEarlyAfterSerializesWithDeliverLocked(t *testing.T) { id := jsonrpc2.Int64ID(1) for i := 0; i < 50; i++ { @@ -470,3 +609,58 @@ func (w *headerCounter) WriteHeader(code int) { w.codes = append(w.codes, code) w.ResponseRecorder.WriteHeader(code) } + +func (w *headerCounter) Write(p []byte) (int, error) { + if len(w.codes) == 0 { + w.WriteHeader(http.StatusOK) + } + return w.ResponseRecorder.Write(p) +} + +func handshake(t *testing.T, url, proto string) string { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + initBody := `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"` + proto + `","capabilities":{},"clientInfo":{"name":"test","version":"1"}}}` + initResp := postJSON(t, ctx, url, "", proto, initBody) + sessionID := initResp.Header.Get(sessionIDHeader) + io.Copy(io.Discard, initResp.Body) + initResp.Body.Close() + if sessionID == "" { + t.Fatal("no session ID from initialize") + } + notifResp := postJSON(t, ctx, url, sessionID, proto, `{"jsonrpc":"2.0","method":"notifications/initialized"}`) + io.Copy(io.Discard, notifResp.Body) + notifResp.Body.Close() + return sessionID +} + +func postJSON(t *testing.T, ctx context.Context, url, sessionID, proto, body string) *http.Response { + t.Helper() + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + if sessionID != "" { + req.Header.Set(sessionIDHeader, sessionID) + } + if proto != "" { + req.Header.Set(protocolVersionHeader, proto) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST: %v", err) + } + return resp +} + +func sseField(body []byte, prefix string) string { + for _, line := range strings.Split(string(body), "\n") { + if strings.HasPrefix(line, prefix) { + return strings.TrimSpace(strings.TrimPrefix(line, prefix)) + } + } + return "" +}