diff --git a/internal/jsonrpc2/identity_test.go b/internal/jsonrpc2/identity_test.go new file mode 100644 index 00000000..6b394874 --- /dev/null +++ b/internal/jsonrpc2/identity_test.go @@ -0,0 +1,203 @@ +// Copyright 2025 The Go MCP SDK Authors. All rights reserved. +// Use of this source code is governed by the license +// that can be found in the LICENSE file. + +package jsonrpc2 + +import ( + "context" + "errors" + "io" + "sync" + "sync/atomic" + "testing" +) + +func TestDecodeMessageIdentity(t *testing.T) { + tests := []struct { + name string + wire string + want any + wantErr bool + }{ + {name: "string", wire: `"request-1"`, want: "request-1"}, + {name: "numeric string", wire: `"9007199254740993"`, want: "9007199254740993"}, + {name: "zero", wire: `0`, want: int64(0)}, + {name: "negative", wire: `-17`, want: int64(-17)}, + {name: "large adjacent low", wire: `9007199254740992`, want: int64(9007199254740992)}, + {name: "large adjacent high", wire: `9007199254740993`, want: int64(9007199254740993)}, + {name: "maximum", wire: `9223372036854775807`, want: int64(9223372036854775807)}, + {name: "minimum", wire: `-9223372036854775808`, want: int64(-9223372036854775808)}, + {name: "integral fraction", wire: `9007199254740993.0`, want: int64(9007199254740993)}, + {name: "integral exponent", wire: `900719925474099300e-2`, want: int64(9007199254740993)}, + {name: "fractional", wire: `1.5`, wantErr: true}, + {name: "fractional exponent", wire: `1e-1`, wantErr: true}, + {name: "positive overflow", wire: `9223372036854775808`, wantErr: true}, + {name: "negative overflow", wire: `-9223372036854775809`, wantErr: true}, + {name: "large exponent", wire: `1e999999999999999999`, wantErr: true}, + {name: "boolean", wire: `true`, wantErr: true}, + {name: "object", wire: `{}`, wantErr: true}, + {name: "array", wire: `[]`, wantErr: true}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + for _, message := range []string{ + `{"jsonrpc":"2.0","id":` + test.wire + `,"method":"test"}`, + `{"jsonrpc":"2.0","id":` + test.wire + `,"result":null}`, + } { + got, err := DecodeMessage([]byte(message)) + if test.wantErr { + if !errors.Is(err, ErrParse) { + t.Fatalf("DecodeMessage() error = %v, want ErrParse", err) + } + continue + } + if err != nil { + t.Fatal(err) + } + var id ID + switch got := got.(type) { + case *Request: + id = got.ID + case *Response: + id = got.ID + default: + t.Fatalf("DecodeMessage() type = %T", got) + } + if id.Raw() != test.want { + t.Fatalf("ID = %#v, want %#v", id.Raw(), test.want) + } + encoded, err := EncodeMessage(got) + if err != nil { + t.Fatal(err) + } + roundTrip, err := DecodeMessage(encoded) + if err != nil { + t.Fatal(err) + } + var roundTripID ID + switch got := roundTrip.(type) { + case *Request: + roundTripID = got.ID + case *Response: + roundTripID = got.ID + } + if roundTripID != id { + t.Fatalf("round-trip ID = %#v, want %#v", roundTripID.Raw(), id.Raw()) + } + } + }) + } +} + +func TestDecodeMessageNotificationIdentity(t *testing.T) { + for _, wire := range []string{ + `{"jsonrpc":"2.0","method":"notify"}`, + `{"jsonrpc":"2.0","id":null,"method":"notify"}`, + } { + msg, err := DecodeMessage([]byte(wire)) + if err != nil { + t.Fatal(err) + } + if msg.(*Request).IsCall() { + t.Fatalf("DecodeMessage(%s) produced a call, want notification", wire) + } + } + if _, err := DecodeMessage([]byte(`{"jsonrpc":"2.0","id":null,"result":null}`)); !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("null response ID error = %v, want ErrInvalidRequest", err) + } +} + +func TestDecodedNumericAndStringIDsAreDistinct(t *testing.T) { + numeric := mustDecodeIdentityMessage(t, `{"jsonrpc":"2.0","id":9007199254740993,"method":"test"}`).(*Request).ID + text := mustDecodeIdentityMessage(t, `{"jsonrpc":"2.0","id":"9007199254740993","method":"test"}`).(*Request).ID + if numeric == text { + t.Fatalf("numeric ID %#v aliases string ID %#v", numeric.Raw(), text.Raw()) + } +} + +func TestConnectionCorrelatesAdjacentLargeIDsOutOfOrder(t *testing.T) { + reader := newIdentityTestReader() + writer := &identityTestWriter{messages: make(chan Message, 2)} + conn := NewConnection(context.Background(), ConnectionConfig{ + Reader: reader, + Writer: writer, + Closer: reader, + Bind: func(*Connection) Handler { + return HandlerFunc(func(context.Context, *Request) (any, error) { + return nil, ErrNotHandled + }) + }, + }) + t.Cleanup(func() { _ = conn.Close() }) + + atomic.StoreInt64(&conn.seq, 9007199254740991) + low := conn.Call(context.Background(), "low", nil) + high := conn.Call(context.Background(), "high", nil) + for range 2 { + <-writer.messages + } + + reader.messages <- mustDecodeIdentityMessage(t, `{"jsonrpc":"2.0","id":9007199254740993,"result":"high"}`) + reader.messages <- mustDecodeIdentityMessage(t, `{"jsonrpc":"2.0","id":9007199254740992,"result":"low"}`) + + var lowResult, highResult string + if err := high.Await(context.Background(), &highResult); err != nil { + t.Fatal(err) + } + if err := low.Await(context.Background(), &lowResult); err != nil { + t.Fatal(err) + } + if lowResult != "low" || highResult != "high" { + t.Fatalf("results = (%q, %q), want (low, high)", lowResult, highResult) + } +} + +type identityTestReader struct { + messages chan Message + done chan struct{} + once sync.Once +} + +func newIdentityTestReader() *identityTestReader { + return &identityTestReader{messages: make(chan Message, 4), done: make(chan struct{})} +} + +func (r *identityTestReader) Read(ctx context.Context) (Message, error) { + select { + case msg := <-r.messages: + return msg, nil + case <-r.done: + return nil, io.EOF + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +func (r *identityTestReader) Close() error { + r.once.Do(func() { close(r.done) }) + return nil +} + +type identityTestWriter struct { + messages chan Message +} + +func (w *identityTestWriter) Write(ctx context.Context, msg Message) error { + select { + case w.messages <- msg: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func mustDecodeIdentityMessage(t *testing.T, wire string) Message { + t.Helper() + msg, err := DecodeMessage([]byte(wire)) + if err != nil { + t.Fatal(err) + } + return msg +} diff --git a/internal/jsonrpc2/messages.go b/internal/jsonrpc2/messages.go index 8b967706..b7ab21ae 100644 --- a/internal/jsonrpc2/messages.go +++ b/internal/jsonrpc2/messages.go @@ -9,6 +9,8 @@ import ( "encoding/json" "errors" "fmt" + "strconv" + "strings" internaljson "github.com/modelcontextprotocol/go-sdk/internal/json" ) @@ -176,7 +178,7 @@ func EncodeIndent(msg Message, prefix, indent string) ([]byte, error) { // when its value is the empty string (see go-sdk#976). type wireDecode struct { VersionTag string `json:"jsonrpc"` - ID any `json:"id,omitempty"` + ID json.RawMessage `json:"id"` Method json.RawMessage `json:"method"` Params json.RawMessage `json:"params,omitempty"` Result json.RawMessage `json:"result,omitempty"` @@ -191,7 +193,7 @@ func DecodeMessage(data []byte) (Message, error) { if msg.VersionTag != wireVersion { return nil, fmt.Errorf("invalid message version tag %q; expected %q", msg.VersionTag, wireVersion) } - id, err := MakeID(msg.ID) + id, err := DecodeID(msg.ID) if err != nil { return nil, err } @@ -222,6 +224,138 @@ func DecodeMessage(data []byte) (Message, error) { return resp, nil } +const maxIDNumberBytes = 128 + +// DecodeID decodes a JSON-RPC request identity without converting numeric IDs +// through float64. It is internal to the SDK and shared with MCP cancellation +// decoding. +func DecodeID(raw json.RawMessage) (ID, error) { + if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { + return ID{}, nil + } + if raw[0] == '"' { + var id string + if err := internaljson.Unmarshal(raw, &id); err != nil { + return ID{}, fmt.Errorf("%w: invalid string ID: %v", ErrParse, err) + } + return StringID(id), nil + } + if len(raw) > maxIDNumberBytes { + return ID{}, fmt.Errorf("%w: numeric ID exceeds %d bytes", ErrParse, maxIDNumberBytes) + } + id, err := parseIntegerID(string(raw)) + if err != nil { + return ID{}, fmt.Errorf("%w: invalid numeric ID %q: %v", ErrParse, raw, err) + } + return Int64ID(id), nil +} + +func parseIntegerID(raw string) (int64, error) { + if raw == "" { + return 0, errors.New("empty number") + } + + negative := raw[0] == '-' + if negative { + raw = raw[1:] + if raw == "" { + return 0, errors.New("missing digits") + } + } + + mantissa, exponentText, hasExponent := strings.Cut(raw, "e") + if !hasExponent { + mantissa, exponentText, hasExponent = strings.Cut(raw, "E") + } + exponent := 0 + if hasExponent { + var err error + exponent, err = parseIDExponent(exponentText) + if err != nil { + return 0, err + } + } + + integerPart, fractionalPart, hasFraction := strings.Cut(mantissa, ".") + if integerPart == "" || !allDecimalDigits(integerPart) || (hasFraction && (fractionalPart == "" || !allDecimalDigits(fractionalPart))) { + return 0, errors.New("invalid number syntax") + } + if len(integerPart) > 1 && integerPart[0] == '0' { + return 0, errors.New("invalid leading zero") + } + + digits := strings.TrimLeft(integerPart+fractionalPart, "0") + if digits == "" { + return 0, nil + } + scale := exponent - len(fractionalPart) + if scale < 0 { + trim := -scale + if trim >= len(digits) || !allZeroes(digits[len(digits)-trim:]) { + return 0, errors.New("ID is not an integer") + } + digits = digits[:len(digits)-trim] + scale = 0 + } + if len(digits)+scale > 19 { + return 0, errors.New("ID is outside the signed 64-bit range") + } + digits += strings.Repeat("0", scale) + if negative { + digits = "-" + digits + } + id, err := strconv.ParseInt(digits, 10, 64) + if err != nil { + return 0, errors.New("ID is outside the signed 64-bit range") + } + return id, nil +} + +func parseIDExponent(raw string) (int, error) { + if raw == "" { + return 0, errors.New("missing exponent") + } + negative := raw[0] == '-' + if negative || raw[0] == '+' { + raw = raw[1:] + } + if raw == "" || !allDecimalDigits(raw) { + return 0, errors.New("invalid exponent") + } + const limit = maxIDNumberBytes + 20 + exponent := 0 + for _, digit := range raw { + value := int(digit - '0') + if exponent > (limit-value)/10 { + exponent = limit + break + } + exponent = exponent*10 + value + } + if negative { + exponent = -exponent + } + return exponent, nil +} + +func allDecimalDigits(s string) bool { + for _, c := range s { + if c < '0' || c > '9' { + return false + } + } + return true +} + +func allZeroes(s string) bool { + for _, c := range s { + if c != '0' { + return false + } + } + return true +} + func marshalToRaw(obj any) (json.RawMessage, error) { if obj == nil { return nil, nil diff --git a/mcp/cancellation_identity_test.go b/mcp/cancellation_identity_test.go new file mode 100644 index 00000000..6eabe260 --- /dev/null +++ b/mcp/cancellation_identity_test.go @@ -0,0 +1,191 @@ +// Copyright 2025 The Go MCP SDK Authors. All rights reserved. +// Use of this source code is governed by the license +// that can be found in the LICENSE file. + +package mcp + +import ( + "context" + "fmt" + "io" + "sync" + "testing" + "time" + + "github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2" + "github.com/modelcontextprotocol/go-sdk/jsonrpc" +) + +func TestCancellationPreservesAdjacentLargeRequestIDs(t *testing.T) { + reader := newCancellationTestReader() + writer := &cancellationTestWriter{messages: make(chan jsonrpc.Message, 4)} + started := make(chan int64, 2) + cancelled := make(chan int64, 2) + release := make(chan struct{}) + preempter := &canceller{} + + conn := jsonrpc2.NewConnection(context.Background(), jsonrpc2.ConnectionConfig{ + Reader: reader, + Writer: writer, + Closer: reader, + Preempter: preempter, + Bind: func(*jsonrpc2.Connection) jsonrpc2.Handler { + return jsonrpc2.HandlerFunc(func(ctx context.Context, req *jsonrpc.Request) (any, error) { + if !req.IsCall() { + return nil, nil + } + jsonrpc2.Async(ctx) + id := req.ID.Raw().(int64) + started <- id + select { + case <-ctx.Done(): + cancelled <- id + return nil, ctx.Err() + case <-release: + return map[string]any{"id": id}, nil + } + }) + }, + }) + preempter.conn = conn + + for _, wire := range []string{ + `{"jsonrpc":"2.0","id":9007199254740992,"method":"slow"}`, + `{"jsonrpc":"2.0","id":9007199254740993,"method":"slow"}`, + } { + reader.messages <- mustDecodeCancellationMessage(t, wire) + } + waitForIDs(t, started, 9007199254740992, 9007199254740993) + + reader.messages <- mustDecodeCancellationMessage(t, `{"jsonrpc":"2.0","method":"notifications/cancelled","params":{"_meta":{"trace":"kept"},"reason":"test","requestId":9007199254740993}}`) + select { + case got := <-cancelled: + if got != 9007199254740993 { + t.Fatalf("cancelled ID = %d, want 9007199254740993", got) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for cancellation") + } + select { + case got := <-cancelled: + t.Fatalf("unexpected cancellation of ID %d", got) + case <-time.After(20 * time.Millisecond): + } + + close(release) + if err := conn.Close(); err != nil { + t.Fatal(err) + } +} + +func TestCancellationRejectsInvalidRequestIDs(t *testing.T) { + preempter := &canceller{} + for _, requestID := range []string{`1.5`, `9223372036854775808`, `true`, `{}`} { + t.Run(requestID, func(t *testing.T) { + req := mustDecodeCancellationMessage(t, fmt.Sprintf(`{"jsonrpc":"2.0","method":"notifications/cancelled","params":{"requestId":%s}}`, requestID)).(*jsonrpc.Request) + if _, err := preempter.Preempt(context.Background(), req); err == nil { + t.Fatal("Preempt() succeeded, want invalid request ID error") + } + }) + } +} + +func TestDecodeCancelledRequestID(t *testing.T) { + tests := []struct { + name string + value string + want any + }{ + {name: "missing", value: ``, want: nil}, + {name: "null", value: `null`, want: nil}, + {name: "string", value: `"9007199254740993"`, want: "9007199254740993"}, + {name: "zero", value: `0`, want: int64(0)}, + {name: "negative", value: `-17`, want: int64(-17)}, + {name: "large adjacent low", value: `9007199254740992`, want: int64(9007199254740992)}, + {name: "large adjacent high", value: `9007199254740993`, want: int64(9007199254740993)}, + {name: "maximum", value: `9223372036854775807`, want: int64(9223372036854775807)}, + {name: "minimum", value: `-9223372036854775808`, want: int64(-9223372036854775808)}, + {name: "integral exponent", value: `900719925474099300e-2`, want: int64(9007199254740993)}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + params := `{"_meta":{"trace":"kept"},"reason":"test"}` + if test.value != "" { + params = `{"_meta":{"trace":"kept"},"reason":"test","requestId":` + test.value + `}` + } + id, err := decodeCancelledRequestID([]byte(params)) + if err != nil { + t.Fatal(err) + } + if id.Raw() != test.want { + t.Fatalf("ID = %#v, want %#v", id.Raw(), test.want) + } + }) + } +} + +type cancellationTestReader struct { + messages chan jsonrpc.Message + done chan struct{} + once sync.Once +} + +func newCancellationTestReader() *cancellationTestReader { + return &cancellationTestReader{messages: make(chan jsonrpc.Message, 4), done: make(chan struct{})} +} + +func (r *cancellationTestReader) Read(ctx context.Context) (jsonrpc.Message, error) { + select { + case msg := <-r.messages: + return msg, nil + case <-r.done: + return nil, io.EOF + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +func (r *cancellationTestReader) Close() error { + r.once.Do(func() { close(r.done) }) + return nil +} + +type cancellationTestWriter struct { + messages chan jsonrpc.Message +} + +func (w *cancellationTestWriter) Write(ctx context.Context, msg jsonrpc.Message) error { + select { + case w.messages <- msg: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func mustDecodeCancellationMessage(t *testing.T, wire string) jsonrpc.Message { + t.Helper() + msg, err := jsonrpc.DecodeMessage([]byte(wire)) + if err != nil { + t.Fatal(err) + } + return msg +} + +func waitForIDs(t *testing.T, ids <-chan int64, want ...int64) { + t.Helper() + got := make(map[int64]bool) + for range want { + select { + case id := <-ids: + got[id] = true + case <-time.After(time.Second): + t.Fatal("timed out waiting for calls to start") + } + } + for _, id := range want { + if !got[id] { + t.Fatalf("request %d did not start", id) + } + } +} diff --git a/mcp/transport.go b/mcp/transport.go index 8b6cd14f..24180d5e 100644 --- a/mcp/transport.go +++ b/mcp/transport.go @@ -258,11 +258,7 @@ type canceller struct { // Preempt implements [jsonrpc2.Preempter]. func (c *canceller) Preempt(ctx context.Context, req *jsonrpc.Request) (result any, err error) { if req.Method == notificationCancelled { - var params CancelledParams - if err := internaljson.Unmarshal(req.Params, ¶ms); err != nil { - return nil, err - } - id, err := jsonrpc2.MakeID(params.RequestID) + id, err := decodeCancelledRequestID(req.Params) if err != nil { return nil, err } @@ -271,6 +267,18 @@ func (c *canceller) Preempt(ctx context.Context, req *jsonrpc.Request) (result a return nil, jsonrpc2.ErrNotHandled } +func decodeCancelledRequestID(data json.RawMessage) (jsonrpc2.ID, error) { + var params struct { + Meta `json:"_meta,omitempty"` + Reason string `json:"reason,omitempty"` + RequestID json.RawMessage `json:"requestId"` + } + if err := internaljson.Unmarshal(data, ¶ms); err != nil { + return jsonrpc2.ID{}, err + } + return jsonrpc2.DecodeID(params.RequestID) +} + // callSubscriptionsListen issues a "subscriptions/listen" call (SEP-2575) // without awaiting its JSON-RPC response. The call's logical lifetime is the // stream of notifications that follow on the same channel — the empty