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

// defaultMaxEventSize bounds the number of bytes buffered for a single SSE
// event before the input gets rejected.
const defaultMaxEventSize = 16 << 20 // 16 MiB
Comment on lines +68 to +70

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should not we allow the possibility to override this value?


// 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.
//
// TODO(rfindley): consider a different API here that makes failure modes more
// apparent.
func scanEvents(r io.Reader) iter.Seq2[Event, error] {
return scanEventsLimited(r, defaultMaxEventSize)
}

// scanEventsLimited is [scanEvents] with an explicit per-event byte budget.
// When maxEventSize > 0, an event is reported as [errMalformedEvent] when
// its size exceed it. A non-positive maxEventSize disables the cap.
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 +105,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 +124,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 +344,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
133 changes: 133 additions & 0 deletions 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 @@ -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
71 changes: 66 additions & 5 deletions mcp/transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -117,13 +117,22 @@ type serverConnection interface {
sessionUpdated(ServerSessionState)
}

// DefaultMaxLineLength is the default maximum number of bytes buffered while
// decoding a single inbound JSON-RPC frame.
const DefaultMaxLineLength = 16 * 1024 * 1024

// A StdioTransport is a [Transport] that communicates over stdin/stdout using
// newline-delimited JSON.
type StdioTransport struct{}
type StdioTransport struct {
// MaxLineLength bounds the number of bytes that may be buffered while
// decoding a single inbound JSON-RPC frame. A value of 0 selects [DefaultMaxLineLength],
// a negative value disables the cap.
MaxLineLength int
}

// Connect implements the [Transport] interface.
func (*StdioTransport) Connect(context.Context) (Connection, error) {
return newIOConn(rwc{os.Stdin, nopCloserWriter{os.Stdout}}), nil
func (t *StdioTransport) Connect(context.Context) (Connection, error) {
return newIOConnLimited(rwc{os.Stdin, nopCloserWriter{os.Stdout}}, t.MaxLineLength), nil
}

// nopCloserWriter is an io.WriteCloser with a trivial Close method.
Expand All @@ -138,11 +147,15 @@ func (nopCloserWriter) Close() error { return nil }
type IOTransport struct {
Reader io.ReadCloser
Writer io.WriteCloser
// MaxLineLength bounds the number of bytes that may be buffered while
// decoding a single inbound JSON-RPC frame. A value of 0 selects [DefaultMaxLineLength],
// a negative value disables the cap.
MaxLineLength int
}

// Connect implements the [Transport] interface.
func (t *IOTransport) Connect(context.Context) (Connection, error) {
return newIOConn(rwc{t.Reader, t.Writer}), nil
return newIOConnLimited(rwc{t.Reader, t.Writer}, t.MaxLineLength), nil
}

// An InMemoryTransport is a [Transport] that communicates over an in-memory
Expand Down Expand Up @@ -476,6 +489,17 @@ type msgOrErr struct {
}

func newIOConn(rwc io.ReadWriteCloser) *ioConn {
return newIOConnLimited(rwc, DefaultMaxLineLength)
}

// newIOConnLimited builds an [ioConn] over rwc that bounds the number of bytes
// buffered while decoding a single inbound JSON-RPC frame to maxLineLength.
// maxLineLength == 0 selects [DefaultMaxLineLength], a negative value means no cap.
func newIOConnLimited(rwc io.ReadWriteCloser, maxLineLength int) *ioConn {
limit := maxLineLength
if limit == 0 {
limit = DefaultMaxLineLength
}
var (
incoming = make(chan msgOrErr)
closed = make(chan struct{})
Expand All @@ -487,7 +511,15 @@ func newIOConn(rwc io.ReadWriteCloser) *ioConn {
// but that is unavoidable since AFAIK there is no (easy and portable) way to
// guarantee that reads of stdin are unblocked when closed.
go func() {
dec := json.NewDecoder(rwc)
var (
reader io.Reader = rwc
limiter *frameLimitReader
)
if limit > 0 {
limiter = &frameLimitReader{r: rwc, limit: limit}
reader = limiter
}
dec := json.NewDecoder(reader)
for {
var raw json.RawMessage
err := dec.Decode(&raw)
Expand All @@ -513,6 +545,9 @@ func newIOConn(rwc io.ReadWriteCloser) *ioConn {
if err != nil {
return
}
if limiter != nil {
limiter.resetFrame()
}
}
}()
return &ioConn{
Expand All @@ -522,6 +557,32 @@ func newIOConn(rwc io.ReadWriteCloser) *ioConn {
}
}

// errFrameTooLarge means that a single inbound JSON-RPC frame exceeded the configured byte
// limit before the value was completely received.
var errFrameTooLarge = errors.New("inbound JSON-RPC frame exceeded the configured maximum line length")

// frameLimitReader bounds the number of bytes [json.Decoder] may buffer while
// decoding a single JSON value. Read returns [errFrameTooLarge] once the budget is exhausted.
type frameLimitReader struct {
r io.Reader
limit int
count int
}

func (r *frameLimitReader) Read(p []byte) (int, error) {
if r.count >= r.limit {
return 0, errFrameTooLarge
}
if len(p) > r.limit-r.count {
p = p[:r.limit-r.count]
}
n, err := r.r.Read(p)
r.count += n
return n, err
}

func (r *frameLimitReader) resetFrame() { r.count = 0 }

func (c *ioConn) SessionID() string { return "" }

func (c *ioConn) sessionUpdated(state ServerSessionState) {
Expand Down
Loading
Loading