Skip to content
Merged
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
58 changes: 51 additions & 7 deletions mcp/event.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,13 +65,21 @@ func writeEvent(w http.ResponseWriter, evt Event) (int, error) {
return n, err
}

// scanEvents iterates SSE events in the given scanner. The iterated error is
// terminal: if encountered, the stream is corrupt or broken and should no
// longer be used.
// DefaultMaxEventSize is the default maximum number of bytes buffered while
// reading a single server-sent event before the input is rejected.
const DefaultMaxEventSize = 16 << 20 // 16 MiB

// scanEventsLimited iterates SSE events in the given reader, buffering at most
// maxEventSize bytes per event. When maxEventSize > 0, an event is reported as
// [errMalformedEvent] once its size exceeds it; a non-positive maxEventSize
// disables the cap.
//
// The iterated error is terminal: if encountered, the stream is corrupt or
// broken and should no longer be used.
//
// TODO(rfindley): consider a different API here that makes failure modes more
// apparent.
func scanEvents(r io.Reader) iter.Seq2[Event, error] {
func scanEventsLimited(r io.Reader, maxEventSize int) iter.Seq2[Event, error] {
reader := bufio.NewReader(r)

// TODO: investigate proper behavior when events are out of order, or have
Expand All @@ -94,8 +102,9 @@ func scanEvents(r io.Reader) iter.Seq2[Event, error] {
// - Lines starting with ":" are ignored.
// - Records are terminated with two consecutive newlines.
var (
evt Event
dataBuf *bytes.Buffer // if non-nil, preceding field was also data
evt Event
dataBuf *bytes.Buffer // if non-nil, preceding field was also data
eventBytes int // bytes read for the current, in-progress event
)
yieldEvent := func() bool {
if dataBuf != nil {
Expand All @@ -112,18 +121,28 @@ func scanEvents(r io.Reader) iter.Seq2[Event, error] {
return true
}
for {
line, err := reader.ReadBytes('\n')
budget := -1
if maxEventSize > 0 {
budget = maxEventSize - eventBytes
}
line, err := readEventLine(reader, budget)
if errors.Is(err, errEventTooLarge) {
yield(Event{}, fmt.Errorf("%w: SSE event exceeded %d bytes without terminating", errMalformedEvent, maxEventSize))
return
}
if err != nil && !errors.Is(err, io.EOF) {
yield(Event{}, fmt.Errorf("error reading event: %v", err))
return
}
eventBytes += len(line)
line = bytes.TrimRight(line, "\r\n")
isEOF := errors.Is(err, io.EOF)

if len(line) == 0 {
if !yieldEvent() {
return
}
eventBytes = 0 // reset the budget between events
if isEOF {
return
}
Expand Down Expand Up @@ -322,6 +341,31 @@ var ErrEventsPurged = errors.New("data purged")
// transient I/O errors which may be retryable.
var errMalformedEvent = errors.New("malformed event")

// errEventTooLarge is returned by [readEventLine] when a single line would
// exceed the remaining per-event byte budget.
var errEventTooLarge = errors.New("SSE event exceeded maximum size")

// readEventLine reads a single '\n'-terminated line from r. When budget >= 0 it
// reads at most budget bytes, returning [errEventTooLarge] once that budget is
// exceeded before a newline arrives.
func readEventLine(r *bufio.Reader, budget int) ([]byte, error) {
if budget < 0 {
return r.ReadBytes('\n')
}
var line []byte
for {
frag, err := r.ReadSlice('\n')
if len(line)+len(frag) > budget {
return nil, errEventTooLarge
}
line = append(line, frag...)
if errors.Is(err, bufio.ErrBufferFull) {
continue // looking for delim
}
return line, err
}
}

// After implements [EventStore.After].
func (s *MemoryEventStore) After(_ context.Context, sessionID, streamID string, index int) iter.Seq2[[]byte, error] {
// Return the data items to yield.
Expand Down
135 changes: 134 additions & 1 deletion mcp/event_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,12 @@
package mcp

import (
"bytes"
"context"
"crypto/rand"
"errors"
"fmt"
"io"
"slices"
"strings"
"testing"
Expand Down Expand Up @@ -134,7 +137,7 @@ func TestScanEvents(t *testing.T) {
r := strings.NewReader(tt.input)
var got []Event
var err error
for e, err2 := range scanEvents(r) {
for e, err2 := range scanEventsLimited(r, DefaultMaxEventSize) {
if err2 != nil {
err = err2
break
Expand Down Expand Up @@ -175,6 +178,136 @@ func TestScanEvents(t *testing.T) {
}
}

// endlessReader streams a fixed prefix once, then repeats fill forever.
type endlessReader struct {
prefix string
sent bool
repeat string
repIndex int
}

func (r *endlessReader) Read(p []byte) (int, error) {
if !r.sent {
n := copy(p, r.prefix)
r.sent = true
return n, nil
}
for i := range p {
p[i] = r.repeat[r.repIndex%len(r.repeat)]
r.repIndex++
}
return len(p), nil
}

func (r *endlessReader) Close() error { return nil }

// TestScanEventsMaxEventSize verifies that scanEvents bounds the bytes
// buffered for a single event.
func TestScanEventsMaxEventSize(t *testing.T) {
wantByte := byte('A')
eventOfSize := func(size int) []byte {
var buf bytes.Buffer
buf.WriteString("data: ")
buf.Write(bytes.Repeat([]byte{wantByte}, size))
buf.WriteString("\n\n")
return buf.Bytes()
}

tests := []struct {
name string
reader io.Reader
maxEventSize int
wantErr bool
wantDataLengths []int
}{
{
name: "negative budget disables the cap",
reader: bytes.NewReader(eventOfSize(512)),
maxEventSize: -1,
wantDataLengths: []int{512},
},
{
name: "long line under cap", // bufio.Reader default buf size is 1 << 12
reader: bytes.NewReader(eventOfSize(1 << 14)),
maxEventSize: 1 << 15,
wantDataLengths: []int{1 << 14},
},
{
name: "unbounded rejected",
reader: &endlessReader{prefix: "data: ", repeat: string(wantByte)},
maxEventSize: 1024,
wantErr: true,
},

{
name: "consecutive data fields longer than limit",
reader: strings.NewReader(func() string {
var b strings.Builder
for range 2048 {
b.WriteString("data: \n")
}
return b.String()
}()),
maxEventSize: 1024,
wantErr: true,
},
{
name: "budget resets between events",
reader: bytes.NewReader(append(eventOfSize(3*1024), eventOfSize(3*1024)...)),
maxEventSize: 4 * 1024,
wantDataLengths: []int{3 * 1024, 3 * 1024},
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
type result struct {
events []Event
err error
}

done := make(chan result, 1)
go func() {
var events []Event
for e, err := range scanEventsLimited(tt.reader, tt.maxEventSize) {
if err != nil {
done <- result{events, err}
return
}
events = append(events, e)
}
done <- result{events, nil}
}()

var res result
select {
case res = <-done:
case <-time.After(5 * time.Second):
t.Fatal("scanEvents did not return")
}

if tt.wantErr {
if !errors.Is(res.err, errMalformedEvent) {
t.Fatalf("got error %v, want errMalformedEvent", res.err)
}
return
}
if res.err != nil {
t.Fatalf("unexpected error: %v", res.err)
}
if len(res.events) != len(tt.wantDataLengths) {
t.Fatalf("got %d events, want %d", len(res.events), len(tt.wantDataLengths))
}
for i, wantLen := range tt.wantDataLengths {
wantRepeated := string(wantByte)
if got := res.events[i].Data; string(got) != strings.Repeat(wantRepeated, wantLen) {
t.Errorf("event %d: got %d data bytes, want %d bytes of %s", i, len(got), wantLen, wantRepeated)
}
}
})
}
}

func TestMemoryEventStoreState(t *testing.T) {
ctx := context.Background()

Expand Down
14 changes: 12 additions & 2 deletions mcp/sse.go
Original file line number Diff line number Diff line change
Expand Up @@ -397,6 +397,11 @@ type SSEClientTransport struct {
// HTTPClient is the client to use for making HTTP requests. If nil,
// http.DefaultClient is used.
HTTPClient *http.Client

// MaxEventSize bounds the number of bytes buffered while reading a single
// server-sent event. A value of 0 selects [DefaultMaxEventSize], a negative
// value disables the cap.
MaxEventSize int
}

// Connect connects through the client endpoint.
Expand Down Expand Up @@ -427,9 +432,14 @@ func (c *SSEClientTransport) Connect(ctx context.Context) (Connection, error) {
return nil, fmt.Errorf("failed to connect: %s", http.StatusText(resp.StatusCode))
}

maxEventSize := c.MaxEventSize
if maxEventSize == 0 {
maxEventSize = DefaultMaxEventSize
}

msgEndpoint, err := func() (*url.URL, error) {
var evt Event
for evt, err = range scanEvents(resp.Body) {
for evt, err = range scanEventsLimited(resp.Body, maxEventSize) {
break
}
if err != nil {
Expand Down Expand Up @@ -458,7 +468,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 scanEvents(resp.Body) {
for evt, err := range scanEventsLimited(resp.Body, maxEventSize) {
if err != nil {
return
}
Expand Down
12 changes: 11 additions & 1 deletion mcp/streamable.go
Original file line number Diff line number Diff line change
Expand Up @@ -2018,6 +2018,11 @@ type StreamableClientTransport struct {
// OAuthHandler is an optional field that, if provided, will be used to authorize the requests.
OAuthHandler auth.OAuthHandler

// MaxEventSize bounds the number of bytes buffered while reading a single
// server-sent event. A value of 0 selects [DefaultMaxEventSize], a negative
// value disables the cap.
MaxEventSize int

// TODO(rfindley): propose exporting these.
// If strict is set, the transport is in 'strict mode', where any violation
// of the MCP spec causes a failure.
Expand Down Expand Up @@ -2105,6 +2110,7 @@ func (t *StreamableClientTransport) Connect(ctx context.Context) (Connection, er
failed: make(chan struct{}),
disableStandaloneSSE: t.DisableStandaloneSSE,
oauthHandler: t.OAuthHandler,
maxEventSize: t.MaxEventSize,
}
return conn, nil
}
Expand All @@ -2126,6 +2132,10 @@ type streamableClientConn struct {
// oauthHandler is the OAuth handler for the connection.
oauthHandler auth.OAuthHandler // from [StreamableClientTransport.OAuthHandler]

// maxEventSize bounds the number of bytes buffered while reading a single
// server-sent event. Resolved from [StreamableClientTransport.MaxEventSize].
maxEventSize int

// Guard calls to Close, as it may be called multiple times.
closeOnce sync.Once
closeErr error
Expand Down Expand Up @@ -2635,7 +2645,7 @@ func (c *streamableClientConn) processStream(ctx context.Context, requestSummary
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
}()
for evt, err := range scanEvents(resp.Body) {
for evt, err := range scanEventsLimited(resp.Body, c.maxEventSize) {
if err != nil {
if ctx.Err() != nil {
return "", 0, true // don't reconnect: client cancelled
Expand Down
Loading
Loading