Skip to content

Commit 6bc1a45

Browse files
committed
fix(mcp): preserve valid numeric request IDs
1 parent b7d9c66 commit 6bc1a45

3 files changed

Lines changed: 102 additions & 9 deletions

File tree

‎internal/mcp/client.go‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77
"errors"
88
"fmt"
99
"io"
10+
"math"
1011
"os"
1112
"os/exec"
1213
"strconv"
@@ -554,7 +555,7 @@ func rpcIDMatches(value any, id int) bool {
554555
}
555556

556557
// jsonRPCIDEchoable reports whether id is a valid JSON-RPC 2.0 identifier type
557-
// (string, integer, or float with no fractional part) that is safe to echo back.
558+
// (string or finite number) that is safe to echo back.
558559
func jsonRPCIDEchoable(id any) bool {
559560
if id == nil {
560561
return false
@@ -565,9 +566,13 @@ func jsonRPCIDEchoable(id any) bool {
565566
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
566567
return true
567568
case float64:
568-
return v == float64(int64(v))
569+
return !math.IsNaN(v) && !math.IsInf(v, 0)
569570
case json.Number:
570-
_, err := v.Int64()
571+
parsed, err := v.Float64()
572+
if err != nil || math.IsNaN(parsed) || math.IsInf(parsed, 0) {
573+
return false
574+
}
575+
_, err = json.Marshal(v)
571576
return err == nil
572577
default:
573578
return false

‎internal/mcp/client_test.go‎

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77
"errors"
88
"fmt"
99
"io"
10+
"math"
1011
"net/http"
1112
"net/http/httptest"
1213
"os"
@@ -1005,6 +1006,72 @@ func TestStdioClientDropsInvalidServerRequestIDs(t *testing.T) {
10051006
}
10061007
}
10071008

1009+
func TestStdioClientRepliesToValidNonIntegerServerRequestIDs(t *testing.T) {
1010+
requests := []struct {
1011+
wire string
1012+
want float64
1013+
}{
1014+
{wire: `{"jsonrpc":"2.0","id":1.5,"method":"roots/list","params":{}}` + "\n", want: 1.5},
1015+
{wire: `{"jsonrpc":"2.0","id":2e2,"method":"roots/list","params":{}}` + "\n", want: 200},
1016+
}
1017+
for _, request := range requests {
1018+
t.Run(fmt.Sprint(request.want), func(t *testing.T) {
1019+
inReader, inWriter := io.Pipe()
1020+
outReader, outWriter := io.Pipe()
1021+
client := &Client{
1022+
reader: newMessageReader(inReader),
1023+
writer: newMessageWriter(outWriter),
1024+
pending: make(map[int]chan dispatchResult),
1025+
}
1026+
t.Cleanup(func() {
1027+
_ = inWriter.Close()
1028+
_ = outReader.Close()
1029+
})
1030+
client.ensureReader()
1031+
if _, err := inWriter.Write([]byte(request.wire)); err != nil {
1032+
t.Fatalf("write server request: %v", err)
1033+
}
1034+
result := make(chan dispatchResult, 1)
1035+
go func() {
1036+
response, err := newMessageReader(outReader).read()
1037+
result <- dispatchResult{message: response, err: err}
1038+
}()
1039+
select {
1040+
case response := <-result:
1041+
if response.err != nil {
1042+
t.Fatalf("read method-not-found response: %v", response.err)
1043+
}
1044+
id, ok := response.message.ID.(float64)
1045+
if !ok || id != request.want || response.message.Error == nil || response.message.Error.Code != -32601 {
1046+
t.Fatalf("response = %#v, want id %v and error -32601", response.message, request.want)
1047+
}
1048+
case <-time.After(time.Second):
1049+
t.Fatal("timed out waiting for method-not-found response")
1050+
}
1051+
})
1052+
}
1053+
}
1054+
1055+
func TestJSONRPCIDEchoableAcceptsFiniteJSONNumbers(t *testing.T) {
1056+
for _, test := range []struct {
1057+
name string
1058+
id any
1059+
want bool
1060+
}{
1061+
{name: "fractional float", id: 1.5, want: true},
1062+
{name: "exponent json number", id: json.Number("2e2"), want: true},
1063+
{name: "not a number", id: math.NaN(), want: false},
1064+
{name: "positive infinity", id: math.Inf(1), want: false},
1065+
{name: "invalid json number", id: json.Number("not-a-number"), want: false},
1066+
} {
1067+
t.Run(test.name, func(t *testing.T) {
1068+
if got := jsonRPCIDEchoable(test.id); got != test.want {
1069+
t.Fatalf("jsonRPCIDEchoable(%v) = %v, want %v", test.id, got, test.want)
1070+
}
1071+
})
1072+
}
1073+
}
1074+
10081075
type gatedCaptureWriter struct {
10091076
started chan struct{}
10101077
release chan struct{}

‎internal/mcp/hang_test.go‎

Lines changed: 27 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -62,6 +62,15 @@ type blockingReader struct {
6262
release chan struct{}
6363
}
6464

65+
type signalingWriter struct {
66+
writes chan struct{}
67+
}
68+
69+
func (writer *signalingWriter) Write(p []byte) (int, error) {
70+
writer.writes <- struct{}{}
71+
return len(p), nil
72+
}
73+
6574
func newBlockingReader() *blockingReader {
6675
return &blockingReader{release: make(chan struct{})}
6776
}
@@ -86,10 +95,11 @@ func (reader *blockingReader) Close() error {
8695
func TestClientRequestUnblocksOnContextCancel(t *testing.T) {
8796
reader := newBlockingReader()
8897
defer reader.Close()
98+
writes := make(chan struct{}, 2)
8999

90100
client := &Client{
91101
reader: newMessageReader(reader),
92-
writer: newMessageWriter(io.Discard),
102+
writer: newMessageWriter(&signalingWriter{writes: writes}),
93103
nextID: 1,
94104
}
95105

@@ -99,6 +109,11 @@ func TestClientRequestUnblocksOnContextCancel(t *testing.T) {
99109
done <- client.request(ctx, "tools/list", map[string]any{}, nil)
100110
}()
101111

112+
select {
113+
case <-writes:
114+
case <-time.After(time.Second):
115+
t.Fatal("first request did not reach the transport")
116+
}
102117
// The request is now parked waiting for a response that never comes.
103118
cancel()
104119

@@ -111,15 +126,21 @@ func TestClientRequestUnblocksOnContextCancel(t *testing.T) {
111126
t.Fatal("request() hung on a non-responsive server")
112127
}
113128

114-
// The lock must be free: a second request under an already-cancelled
115-
// context must return immediately rather than block.
116-
cancelled, cancel2 := context.WithCancel(context.Background())
117-
cancel2()
129+
// The shared state must be free: prove a second live request reaches the
130+
// transport, then cancel it and verify cancellation releases the caller.
131+
secondCtx, cancel2 := context.WithCancel(context.Background())
132+
defer cancel2()
118133
second := make(chan error, 1)
119134
go func() {
120-
second <- client.request(cancelled, "tools/list", map[string]any{}, nil)
135+
second <- client.request(secondCtx, "tools/list", map[string]any{}, nil)
121136
}()
122137
select {
138+
case <-writes:
139+
case <-time.After(time.Second):
140+
t.Fatal("second request did not reach the transport")
141+
}
142+
cancel2()
143+
select {
123144
case err := <-second:
124145
if !errors.Is(err, context.Canceled) {
125146
t.Fatalf("second request() error = %v, want context.Canceled", err)

0 commit comments

Comments
 (0)