Skip to content
Open
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
48 changes: 41 additions & 7 deletions pkg/transport/proxy/streamable/streamable_proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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
Expand Down
84 changes: 84 additions & 0 deletions pkg/transport/proxy/streamable/utils_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)
}