Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 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
260 changes: 247 additions & 13 deletions internal/mcp/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"errors"
"fmt"
"io"
"math"
"os"
"os/exec"
"strconv"
Expand Down Expand Up @@ -49,20 +50,38 @@ type Client struct {
stdin io.WriteCloser
reader *messageReader
writer *messageWriter
mu sync.Mutex
closeMu sync.Mutex
idMu sync.Mutex
nextID int
cleanup func()

writeMu sync.Mutex
writeQueue chan writeOp
writeClosed bool
writeSenders sync.WaitGroup
writerStop chan struct{}
writerDone chan struct{}
courtesyOverflow []writeOp

// 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 {
ctx context.Context
message rpcMessage
done chan error
}

const writeQueueCapacity = 32

var errMCPClientClosed = errors.New("MCP client closed")

// dispatchResult carries one matched JSON-RPC response (or a terminal reader
// error) to a waiting caller.
type dispatchResult struct {
Expand Down Expand Up @@ -252,7 +271,8 @@ func (client *Client) Close() error {
// Fail any callers still waiting on a response. The blocking read in the
// reader goroutine is released below when stdin closes and the process
// exits (or is killed), EOFing stdout.
client.failAll(errors.New("MCP client closed"))
client.failAll(errMCPClientClosed)
client.beginWriterShutdown()

var err error
stdin := client.stdin
Expand Down Expand Up @@ -292,6 +312,7 @@ func (client *Client) Close() error {
client.cleanup()
client.cleanup = nil
}
client.finishWriterShutdown()
return err
}

Expand All @@ -307,20 +328,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 so a caller with a canceled context can stop
// waiting even when the peer is not draining its input.
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 +350,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 +380,182 @@ func (client *Client) request(ctx context.Context, method string, params any, ta
}
}

func (client *Client) ensureWriter() {
_ = client.startWriter()
}

func (client *Client) startWriter() error {
client.writeMu.Lock()
defer client.writeMu.Unlock()
if client.writeClosed {
return errMCPClientClosed
}
if client.writeQueue != nil {
return nil
}
client.writeQueue = make(chan writeOp, writeQueueCapacity)
client.writerStop = make(chan struct{})
client.writerDone = make(chan struct{})
go client.writeLoop()
return nil
}

func (client *Client) writerStopped() bool {
if client.writerStop == nil {
return false
}
select {
case <-client.writerStop:
return true
default:
return false
}
}

func (client *Client) beginWriterShutdown() {
client.writeMu.Lock()
defer client.writeMu.Unlock()
if client.writeClosed {
return
}
client.writeClosed = true
if client.writerStop != nil {
close(client.writerStop)
}
for _, op := range client.courtesyOverflow {
if op.done != nil {
op.done <- errMCPClientClosed
}
}
client.courtesyOverflow = nil
}

func (client *Client) finishWriterShutdown() {
client.writeMu.Lock()
queue := client.writeQueue
done := client.writerDone
client.writeQueue = nil
client.writeMu.Unlock()
if queue == nil {
return
}
client.writeSenders.Wait()
close(queue)
if done != nil {
<-done
}
}

func (client *Client) writeLoop() {
defer close(client.writerDone)
for op := range client.writeQueue {
if client.writerStopped() {
if op.done != nil {
op.done <- errMCPClientClosed
}
continue
}
if op.ctx != nil {
select {
case <-op.ctx.Done():
if op.done != nil {
op.done <- op.ctx.Err()
}
continue
default:
}
}
err := client.writer.write(op.message)
if op.done != nil {
op.done <- err
}
client.drainCourtesyOverflow()
}
}

func (client *Client) drainCourtesyOverflow() {
client.writeMu.Lock()
defer client.writeMu.Unlock()
if client.writeClosed || client.writeQueue == nil {
client.courtesyOverflow = nil
return
}
for len(client.courtesyOverflow) > 0 {
select {
case client.writeQueue <- client.courtesyOverflow[0]:
client.courtesyOverflow = client.courtesyOverflow[1:]
default:
return
}
}
}

func (client *Client) enqueueCourtesy(message rpcMessage) {
if err := client.startWriter(); err != nil {
return
}
op := writeOp{message: message}
client.writeMu.Lock()
defer client.writeMu.Unlock()
if client.writeClosed || client.writeQueue == nil {
return
}
select {
case client.writeQueue <- op:
default:
client.courtesyOverflow = append(client.courtesyOverflow, op)
}
}

func (client *Client) writeMessage(ctx context.Context, message rpcMessage) error {
if err := ctx.Err(); err != nil {
return err
}
if err := client.startWriter(); err != nil {
return err
}
done := make(chan error, 1)
op := writeOp{ctx: ctx, message: message, done: done}

client.writeMu.Lock()
if client.writeClosed {
client.writeMu.Unlock()
return errMCPClientClosed
}
stop := client.writerStop
queue := client.writeQueue
client.writeSenders.Add(1)
client.writeMu.Unlock()

select {
case <-ctx.Done():
client.writeSenders.Done()
return ctx.Err()
case <-stop:
client.writeSenders.Done()
return errMCPClientClosed
case queue <- op:
client.writeSenders.Done()
}

select {
case <-ctx.Done():
return ctx.Err()
case <-stop:
select {
case err := <-done:
if err != nil {
return err
}
return errMCPClientClosed
case <-ctx.Done():
return ctx.Err()
}
case err := <-done:
return err
}
}

// 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 +580,20 @@ 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.methodPresent || message.Method != "" {
if message.ID != nil && jsonRPCIDEchoable(message.ID) {
client.enqueueCourtesy(rpcMessage{
ID: message.ID,
Error: &rpcError{
Code: -32601,
Message: fmt.Sprintf("Method %q not supported", message.Method),
},
})
}
continue
}
if message.ID == nil {
continue
}
Expand Down Expand Up @@ -471,12 +680,37 @@ func rpcIDMatches(value any, id int) bool {
}
}

// jsonRPCIDEchoable reports whether id is a valid JSON-RPC 2.0 identifier type
// (string or finite number) 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 !math.IsNaN(v) && !math.IsInf(v, 0)
case json.Number:
parsed, err := v.Float64()
if err != nil || math.IsNaN(parsed) || math.IsInf(parsed, 0) {
return false
}
_, err = json.Marshal(v)
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
Loading