diff --git a/pkg/transport/proxy/streamable/streamable_proxy.go b/pkg/transport/proxy/streamable/streamable_proxy.go index 6d339d1b7e..b7d47f0ad3 100644 --- a/pkg/transport/proxy/streamable/streamable_proxy.go +++ b/pkg/transport/proxy/streamable/streamable_proxy.go @@ -723,18 +723,18 @@ func (p *HTTPProxy) handleSingleRequestSSE( // deliverCh is nil when req carried no progressToken (see // setupProgressRouting); a nil channel's case is never selected, // so this arm only fires when progress routing is active. - data, err := jsonrpc2.EncodeMessage(progressMsg) - if err != nil { - slog.Error("failed to encode progress notification", "error", err) - continue - } - if err := writeSSEData(w, flusher, data); err != nil { - slog.Debug("failed to write progress notification to SSE stream", "error", err) + if !writeSSEProgressMessage(w, flusher, progressMsg) { return } // Progress does not end the request; keep waiting for more // progress or the final response. case msg := <-waitCh: + // dispatchResponses preserves backend message order. Therefore, any + // progress sent before this final response is already buffered in + // deliverCh and must be emitted before returning from this stream. + if !drainPendingSSEProgress(w, flusher, deliverCh) { + return + } p.writeSingleRequestSSEFinalResponse(w, flusher, msg, ck) return case <-ctx.Done(): @@ -750,6 +750,40 @@ func (p *HTTPProxy) handleSingleRequestSSE( } } +// writeSSEProgressMessage encodes progressMsg and writes it as one SSE data +// frame. It reports whether the stream remains writable. +func writeSSEProgressMessage(w io.Writer, flusher http.Flusher, progressMsg jsonrpc2.Message) bool { + data, err := jsonrpc2.EncodeMessage(progressMsg) + if err != nil { + slog.Error("failed to encode progress notification", "error", err) + return true + } + if err := writeSSEData(w, flusher, data); err != nil { + slog.Debug("failed to write progress notification to SSE stream", "error", err) + return false + } + return true +} + +// drainPendingSSEProgress writes progress messages that were already routed to +// deliverCh before a final response is emitted. It reports whether the stream +// remains writable. +func drainPendingSSEProgress(w io.Writer, flusher http.Flusher, deliverCh <-chan jsonrpc2.Message) bool { + for { + select { + case progressMsg, ok := <-deliverCh: + if !ok { + return true + } + if !writeSSEProgressMessage(w, flusher, progressMsg) { + return false + } + default: + return true + } + } +} + // writeSingleRequestSSEFinalResponse restores msg's original client ID (if it // is a correlated *jsonrpc2.Response) and writes it as the final SSE data: // frame for handleSingleRequestSSE's request. Errors encoding/restoring are diff --git a/pkg/transport/proxy/streamable/utils_test.go b/pkg/transport/proxy/streamable/utils_test.go index 70eebb4be2..64926392b4 100644 --- a/pkg/transport/proxy/streamable/utils_test.go +++ b/pkg/transport/proxy/streamable/utils_test.go @@ -5,10 +5,13 @@ package streamable import ( "bytes" + "io" + "strings" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/exp/jsonrpc2" ) // TestIsSupportedMCPVersion verifies membership in supportedMCPVersions: every @@ -48,6 +51,10 @@ type recordingFlusher struct { func (f *recordingFlusher) Flush() { f.flushed++ } +type failingWriter struct{} + +func (failingWriter) Write([]byte) (int, error) { return 0, io.ErrClosedPipe } + // TestWriteSSEDataIncludesEventMessage verifies writeSSEData emits an explicit // SSE event name before the data line. Spec-lenient MCP clients (for example // @ai-sdk/mcp) only dispatch frames where event === "message" and drop @@ -74,3 +81,80 @@ func TestWriteSSEDataIncludesEventMessage(t *testing.T) { assert.True(t, bytes.HasPrefix([]byte(got), []byte("event: message\n")), "frame must start with event: message, not a bare data: line") } + +// TestDrainPendingSSEProgressWritesQueuedMessagesBeforeFinal verifies that +// progress already buffered when a POST-SSE request receives its final +// response is emitted first. This makes the select race in +// handleSingleRequestSSE deterministic: the helper drains all currently ready +// request-scoped progress messages without waiting for future ones. +func TestDrainPendingSSEProgressWritesQueuedMessagesBeforeFinal(t *testing.T) { + t.Parallel() + + progressCh := make(chan jsonrpc2.Message, 2) + first, err := jsonrpc2.NewNotification("notifications/progress", map[string]any{"progress": 1}) + require.NoError(t, err) + second, err := jsonrpc2.NewNotification("notifications/progress", map[string]any{"progress": 2}) + require.NoError(t, err) + final, err := jsonrpc2.NewResponse(jsonrpc2.StringID("request-1"), map[string]any{"status": "done"}, nil) + require.NoError(t, err) + progressCh <- first + progressCh <- second + + var buf bytes.Buffer + flusher := &recordingFlusher{} + require.True(t, drainPendingSSEProgress(&buf, flusher, progressCh)) + finalData, err := jsonrpc2.EncodeMessage(final) + require.NoError(t, err) + require.NoError(t, writeSSEData(&buf, flusher, finalData)) + + got := buf.String() + firstIndex := strings.Index(got, `"progress":1`) + secondIndex := strings.Index(got, `"progress":2`) + finalIndex := strings.Index(got, `"status":"done"`) + require.GreaterOrEqual(t, firstIndex, 0, "first progress frame was not written") + require.GreaterOrEqual(t, secondIndex, 0, "second progress frame was not written") + require.GreaterOrEqual(t, finalIndex, 0, "final response frame was not written") + assert.Less(t, firstIndex, secondIndex) + assert.Less(t, secondIndex, finalIndex) + assert.Equal(t, 3, flusher.flushed, "every progress and final frame must flush") +} + +// TestDrainPendingSSEProgressReturnsImmediatelyWithoutQueuedProgress verifies +// that requests without progress routing, and requests with an empty route, +// retain their prior final-response behavior without blocking. +func TestDrainPendingSSEProgressReturnsImmediatelyWithoutQueuedProgress(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + deliverCh <-chan jsonrpc2.Message + }{ + {name: "nil route", deliverCh: nil}, + {name: "empty route", deliverCh: make(chan jsonrpc2.Message, 1)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + var buf bytes.Buffer + flusher := &recordingFlusher{} + assert.True(t, drainPendingSSEProgress(&buf, flusher, tt.deliverCh)) + assert.Empty(t, buf.String()) + assert.Zero(t, flusher.flushed) + }) + } +} + +// TestWriteSSEProgressMessageStopsOnWriteFailure verifies a disconnected SSE +// client stops further draining rather than silently discarding queued frames. +func TestWriteSSEProgressMessageStopsOnWriteFailure(t *testing.T) { + t.Parallel() + + progress, err := jsonrpc2.NewNotification("notifications/progress", map[string]any{"progress": 1}) + require.NoError(t, err) + flusher := &recordingFlusher{} + + assert.False(t, writeSSEProgressMessage(failingWriter{}, flusher, progress)) + assert.Zero(t, flusher.flushed) +}