Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 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
116 changes: 104 additions & 12 deletions internal/mcp/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,20 +49,30 @@ type Client struct {
stdin io.WriteCloser
reader *messageReader
writer *messageWriter
mu sync.Mutex
closeMu sync.Mutex
idMu sync.Mutex
nextID int
cleanup func()

writeQueue chan writeOp
writerOnce sync.Once

// dispatchMu guards the response-dispatch state shared with the single
// reader goroutine. It is never held across a blocking read.
// reader goroutine. It is never held across a blocking read or write.
dispatchMu sync.Mutex
readerOnce sync.Once
pending map[int]chan dispatchResult
readErr error
readDone bool
}

type writeOp struct {
message rpcMessage
done chan error
}

const writeQueueCapacity = 32

// dispatchResult carries one matched JSON-RPC response (or a terminal reader
// error) to a waiting caller.
type dispatchResult struct {
Expand Down Expand Up @@ -307,20 +317,20 @@ func (client *Client) request(ctx context.Context, method string, params any, ta
return err
}

// Allocate an id, register a response channel, and write the request while
// holding client.mu. The mutex serializes writes and id allocation but is
// released before the (potentially unbounded) wait for the response, so a
// hung server never holds the lock and blocks other callers/Close.
client.mu.Lock()
// Allocate an id and register a response channel. ID allocation and dispatch
// registrations are fast non-blocking operations. Message transmission is
// handled via writeMessage, ensuring a hung server never blocks other callers
// or prevents a caller with a deadline from giving up.
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
client.idMu.Lock()
id := client.nextID
client.nextID++
client.idMu.Unlock()

responses := make(chan dispatchResult, 1)
client.dispatchMu.Lock()
if client.readDone {
readErr := client.readErr
client.dispatchMu.Unlock()
client.mu.Unlock()
if readErr != nil {
return readErr
}
Expand All @@ -329,16 +339,14 @@ func (client *Client) request(ctx context.Context, method string, params any, ta
client.pending[id] = responses
client.dispatchMu.Unlock()

if err := client.writer.write(rpcMessage{
if err := client.writeMessage(ctx, rpcMessage{
ID: id,
Method: method,
Params: rawParams,
}); err != nil {
client.removePending(id)
client.mu.Unlock()
return err
}
client.mu.Unlock()

select {
case <-ctx.Done():
Expand All @@ -361,6 +369,43 @@ func (client *Client) request(ctx context.Context, method string, params any, ta
}
}

// ensureWriter lazily starts the single writer goroutine.
func (client *Client) ensureWriter() {
client.writerOnce.Do(func() {
if client.writeQueue == nil {
client.writeQueue = make(chan writeOp, writeQueueCapacity)
}
go client.writeLoop()
})
}

func (client *Client) writeLoop() {
for op := range client.writeQueue {
err := client.writer.write(op.message)
if op.done != nil {
op.done <- err
}
}
}

func (client *Client) writeMessage(ctx context.Context, message rpcMessage) error {
client.ensureWriter()
done := make(chan error, 1)
op := writeOp{message: message, done: done}
select {
case <-ctx.Done():
return ctx.Err()
case client.writeQueue <- op:
}

select {
case <-ctx.Done():
return ctx.Err()
case err := <-done:
return err
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
}
}

// ensureReader lazily starts the single reader goroutine. It runs once per
// client; subsequent calls are no-ops.
func (client *Client) ensureReader() {
Expand All @@ -385,6 +430,32 @@ func (client *Client) readLoop() {
client.failAll(err)
return
}
// A message with a Method is a server-initiated request or notification.
// It must never be routed as a response to a pending client request.
if message.Method != "" {
if message.ID != nil && jsonRPCIDEchoable(message.ID) {
// Send the courtesy -32601 reply via the bounded writer queue. If the
// queue is saturated (e.g. an undrained server pipe), drop the reply
// immediately so it never stalls readLoop, holds a mutex, or blocks callers.
client.ensureWriter()
id := message.ID
method := message.Method
courtesy := writeOp{
message: rpcMessage{
ID: id,
Error: &rpcError{
Code: -32601,
Message: fmt.Sprintf("Method %q not supported", method),
},
},
}
select {
case client.writeQueue <- courtesy:
default:
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
}
continue
}
if message.ID == nil {
continue
}
Expand Down Expand Up @@ -471,12 +542,33 @@ func rpcIDMatches(value any, id int) bool {
}
}

// jsonRPCIDEchoable reports whether id is a valid JSON-RPC 2.0 identifier type
// (string, integer, or float with no fractional part) that is safe to echo back.
func jsonRPCIDEchoable(id any) bool {
if id == nil {
return false
}
switch v := id.(type) {
case string:
return true
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
return true
case float64:
return v == float64(int64(v))
case json.Number:
_, err := v.Int64()
return err == nil
Comment thread
coderabbitai[bot] marked this conversation as resolved.
default:
return false
}
}

func (client *Client) notify(method string, params any) error {
rawParams, err := json.Marshal(params)
if err != nil {
return err
}
return client.writer.write(rpcMessage{
return client.writeMessage(context.Background(), rpcMessage{
Method: method,
Params: rawParams,
})
Expand Down
190 changes: 190 additions & 0 deletions internal/mcp/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -842,3 +842,193 @@ func TestBoundedBufferCapsRetainedBytes(t *testing.T) {
t.Fatalf("retained %q, want %q (capped at 8 bytes, head kept)", got, "hellowor")
}
}

func TestStdioClientIgnoresServerInitiatedRequestsInPendingResponses(t *testing.T) {
inReader, inWriter := io.Pipe()
outReader, outWriter := io.Pipe()

client := &Client{
reader: newMessageReader(inReader),
writer: newMessageWriter(outWriter),
pending: make(map[int]chan dispatchResult),
}
defer func() {
_ = inWriter.Close()
_ = outReader.Close()
}()

// Drain output responses from client
go func() {
buf := make([]byte, 1024)
for {
if _, err := outReader.Read(buf); err != nil {
return
}
}
}()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated

client.ensureReader()

// Register a pending response for ID 1
responses := make(chan dispatchResult, 1)
client.dispatchMu.Lock()
client.pending[1] = responses
client.dispatchMu.Unlock()

// Simulate server sending a request with ID 1 ("roots/list")
serverReq := `{"jsonrpc":"2.0","id":1,"method":"roots/list","params":{}}` + "\n"
go func() {
_, _ = inWriter.Write([]byte(serverReq))
}()

// The pending channel for client ID 1 should NOT receive the server request.
select {
case res := <-responses:
t.Fatalf("pending request 1 received server request: %#v", res.message)
case <-time.After(100 * time.Millisecond):
// Expected: server request was not misdelivered to pending client request
}

// Now send the actual response for ID 1
serverResp := `{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}` + "\n"
go func() {
_, _ = inWriter.Write([]byte(serverResp))
}()

select {
case res := <-responses:
if res.message.Method != "" || len(res.message.Result) == 0 {
t.Fatalf("expected valid response, got %#v", res.message)
}
case <-time.After(500 * time.Millisecond):
t.Fatal("timed out waiting for actual response")
}
}

// TestStdioClientUndrainedServerDoesNotStallReadLoop proves that when a server
// stops draining its stdin, an asynchronous courtesy -32601 reply does not stall
// the read loop or block client.mu from servicing legitimate responses.
func TestStdioClientUndrainedServerDoesNotStallReadLoop(t *testing.T) {
inReader, inWriter := io.Pipe()
outReader, outWriter := io.Pipe()

client := &Client{
reader: newMessageReader(inReader),
writer: newMessageWriter(outWriter),
pending: make(map[int]chan dispatchResult),
}
defer func() {
_ = inWriter.Close()
_ = outReader.Close()
}()

client.ensureReader()

// 1. Register a pending response for call ID 1
responses := make(chan dispatchResult, 1)
client.dispatchMu.Lock()
client.pending[1] = responses
client.dispatchMu.Unlock()

// 2. Server sends request ID 2 (we DO NOT drain outReader so server stdin pipe is blocked)
serverReq := `{"jsonrpc":"2.0","id":2,"method":"roots/list","params":{}}` + "\n"
go func() {
_, _ = inWriter.Write([]byte(serverReq))
// 3. Immediately send response for ID 1
serverResp := `{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}` + "\n"
_, _ = inWriter.Write([]byte(serverResp))
}()

// 4. Verify the response for ID 1 is delivered without being stalled by the undrained pipe
select {
case res := <-responses:
if res.message.Method != "" || len(res.message.Result) == 0 {
t.Fatalf("expected valid response for ID 1, got %#v", res.message)
}
case <-time.After(1 * time.Second):
t.Fatal("readLoop stalled on undrained error write; pending response was not delivered")
}
}

func TestStdioClientDropsInvalidServerRequestIDs(t *testing.T) {
inReader, inWriter := io.Pipe()
outReader, outWriter := io.Pipe()

client := &Client{
reader: newMessageReader(inReader),
writer: newMessageWriter(outWriter),
pending: make(map[int]chan dispatchResult),
}
defer func() {
_ = inWriter.Close()
_ = outReader.Close()
}()

client.ensureReader()

// Server sends request with boolean ID (invalid JSON-RPC id)
serverReq := `{"jsonrpc":"2.0","id":true,"method":"roots/list","params":{}}` + "\n"
go func() {
_, _ = inWriter.Write([]byte(serverReq))
}()

// Read should not produce any output response for invalid ID
buf := make([]byte, 1024)
readChan := make(chan int, 1)
go func() {
n, _ := outReader.Read(buf)
readChan <- n
}()

select {
case n := <-readChan:
t.Fatalf("unexpected reply for invalid ID: %s", string(buf[:n]))
case <-time.After(100 * time.Millisecond):
// Expected: dropped cleanly
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

// TestStdioClientUndrainedServerDoesNotBlockCallerWithDeadline verifies that
// when the server's input pipe is completely blocked and courtesy replies are
// queued/dropped, a caller invoking request() with a deadline aborts cleanly
// when the context expires rather than hanging indefinitely on a write or mutex.
func TestStdioClientUndrainedServerDoesNotBlockCallerWithDeadline(t *testing.T) {
inReader, inWriter := io.Pipe()
outReader, outWriter := io.Pipe()

client := &Client{
reader: newMessageReader(inReader),
writer: newMessageWriter(outWriter),
pending: make(map[int]chan dispatchResult),
}
defer func() {
_ = inWriter.Close()
_ = outReader.Close()
}()

client.ensureReader()

// Flood server requests to saturate write queue while outReader is NOT drained.
for i := 1; i <= 50; i++ {
serverReq := fmt.Sprintf(`{"jsonrpc":"2.0","id":%d,"method":"roots/list","params":{}}`+"\n", i+100)
_, err := inWriter.Write([]byte(serverReq))
if err != nil {
t.Fatalf("failed to write server request: %v", err)
}
}

// Caller with short deadline invokes request()
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()

start := time.Now()
err := client.request(ctx, "tools/list", map[string]any{}, nil)
elapsed := time.Since(start)

if !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, context.Canceled) {
t.Fatalf("expected deadline exceeded, got: %v", err)
}
if elapsed > 1*time.Second {
t.Fatalf("request took too long to abort on deadline: %v", elapsed)
}
}