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
203 changes: 203 additions & 0 deletions internal/jsonrpc2/identity_test.go
Original file line number Diff line number Diff line change
@@ -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
}
138 changes: 136 additions & 2 deletions internal/jsonrpc2/messages.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ import (
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"

internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
)
Expand Down Expand Up @@ -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"`
Expand All @@ -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
}
Expand Down Expand Up @@ -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
Expand Down
Loading