diff --git a/mcp/mrtr.go b/mcp/mrtr.go index fdf98ece..7d358c8c 100644 --- a/mcp/mrtr.go +++ b/mcp/mrtr.go @@ -54,11 +54,7 @@ func validateMultiRoundTripResult(logger *slog.Logger, res multiRoundTripRespons } func clientSupportsMultiRoundTrip(ss *ServerSession) bool { - protocolVersion := latestProtocolVersion - if iparams := ss.InitializeParams(); iparams != nil { - protocolVersion = iparams.ProtocolVersion - } - return protocolVersion >= protocolVersion20260728 + return !ss.negotiatedLegacyProtocol() } func clientMultiRoundTripMiddleware() Middleware { diff --git a/mcp/mrtr_test.go b/mcp/mrtr_test.go index d9662198..82b372c3 100644 --- a/mcp/mrtr_test.go +++ b/mcp/mrtr_test.go @@ -8,6 +8,7 @@ package mcp import ( "context" + "encoding/json" "fmt" "slices" "strings" @@ -16,6 +17,8 @@ import ( "github.com/google/go-cmp/cmp" "github.com/google/jsonschema-go/jsonschema" + "github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2" + "github.com/modelcontextprotocol/go-sdk/jsonrpc" ) func TestMultiRoundTrip_ManualRetry(t *testing.T) { @@ -863,6 +866,167 @@ func TestSetMultiRoundTripRetryParams(t *testing.T) { }) } +func TestClientSupportsMultiRoundTrip(t *testing.T) { + tests := []struct { + name string + state ServerSessionState + want bool + }{ + { + // A session that ran no handshake speaks the new protocol: every + // request carries its own version in _meta (SEP-2575). + name: "no handshake", + want: true, + }, + { + name: "discover, new protocol", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20260728}, + }, + want: true, + }, + { + name: "initialize, legacy version", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20251125}, + NegotiatedProtocolVersion: protocolVersion20251125, + }, + want: false, + }, + { + // initialize is deprecated in protocolVersion20260728, so a client + // asking for it there is negotiated down and must be served the + // legacy interaction whatever it declared. + name: "initialize, negotiated down", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20260728}, + NegotiatedProtocolVersion: protocolVersion20251125, + }, + want: false, + }, + { + name: "stateless request, legacy version", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20250618}, + }, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ss := &ServerSession{state: tt.state} + if got := clientSupportsMultiRoundTrip(ss); got != tt.want { + t.Errorf("clientSupportsMultiRoundTrip() = %t, want %t", got, tt.want) + } + }) + } +} + +// TestMultiRoundTrip_NegotiatedDownFromNewProtocol asserts that a session +// negotiated down to protocolVersion20251125 gets elicitation/create, not an +// input-required result that its version does not define. +func TestMultiRoundTrip_NegotiatedDownFromNewProtocol(t *testing.T) { + ctx := context.Background() + + srv := NewServer(testImpl, nil) + srv.AddTool( + &Tool{Name: "act", InputSchema: &jsonschema.Schema{Type: "object"}}, + func(ctx context.Context, req *CallToolRequest) (*CallToolResult, error) { + if len(req.Params.InputResponses) == 0 { + return &CallToolResult{ + InputRequests: InputRequestMap{"confirm": &ElicitParams{Message: "OK?"}}, + RequestState: "state-1", + }, nil + } + return &CallToolResult{Content: []Content{&TextContent{Text: "confirmed"}}}, nil + }, + ) + + ct, st := NewInMemoryTransports() + ss, err := srv.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server.Connect() error = %v", err) + } + defer ss.Close() + + conn, err := ct.Connect(ctx) + if err != nil { + t.Fatalf("transport.Connect() error = %v", err) + } + defer conn.Close() + + write := func(msg jsonrpc.Message, err error) { + t.Helper() + if err != nil { + t.Fatalf("building message: %v", err) + } + if err := conn.Write(ctx, msg); err != nil { + t.Fatalf("conn.Write() error = %v", err) + } + } + read := func() jsonrpc.Message { + t.Helper() + msg, err := conn.Read(ctx) + if err != nil { + t.Fatalf("conn.Read() error = %v", err) + } + return msg + } + + write(jsonrpc2.NewCall(jsonrpc2.Int64ID(1), methodInitialize, &InitializeParams{ + ProtocolVersion: protocolVersion20260728, + ClientInfo: testImpl, + Capabilities: &ClientCapabilities{Elicitation: &ElicitationCapabilities{}}, + })) + initResp, ok := read().(*jsonrpc2.Response) + if !ok { + t.Fatalf("initialize: got %T, want *jsonrpc2.Response", initResp) + } + if initResp.Error != nil { + t.Fatalf("initialize failed: %v", initResp.Error) + } + var initRes InitializeResult + if err := json.Unmarshal(initResp.Result, &initRes); err != nil { + t.Fatalf("unmarshalling initialize result: %v", err) + } + if initRes.ProtocolVersion != protocolVersion20251125 { + t.Fatalf("negotiated protocol version = %q, want %q", initRes.ProtocolVersion, protocolVersion20251125) + } + write(jsonrpc2.NewNotification(notificationInitialized, &InitializedParams{})) + + write(jsonrpc2.NewCall(jsonrpc2.Int64ID(2), methodCallTool, &CallToolParams{Name: "act"})) + msg := read() + elicitReq, ok := msg.(*jsonrpc2.Request) + if !ok { + resp := msg.(*jsonrpc2.Response) + t.Fatalf("tools/call was answered without an %q request: result = %s, error = %v", + methodElicit, resp.Result, resp.Error) + } + if elicitReq.Method != methodElicit { + t.Fatalf("server request method = %q, want %q", elicitReq.Method, methodElicit) + } + write(jsonrpc2.NewResponse(elicitReq.ID, &ElicitResult{Action: "accept"}, nil)) + + callResp, ok := read().(*jsonrpc2.Response) + if !ok { + t.Fatalf("tools/call: got %T, want *jsonrpc2.Response", callResp) + } + if callResp.Error != nil { + t.Fatalf("tools/call failed: %v", callResp.Error) + } + var callRes CallToolResult + if err := json.Unmarshal(callResp.Result, &callRes); err != nil { + t.Fatalf("unmarshalling tools/call result: %v", err) + } + if len(callRes.Content) != 1 { + t.Fatalf("len(result.Content) = %d, want 1", len(callRes.Content)) + } + if got := callRes.Content[0].(*TextContent).Text; got != "confirmed" { + t.Errorf("result text = %q, want %q", got, "confirmed") + } +} + func mustConnect(t *testing.T, s *Server, clientOpts *ClientOptions) *ClientSession { t.Helper() diff --git a/mcp/server.go b/mcp/server.go index 16aad6a9..2978f2c6 100644 --- a/mcp/server.go +++ b/mcp/server.go @@ -791,7 +791,7 @@ func (s *Server) notifySessions(n string) { // shared session channel without opt-in; collect them while we hold the lock. var legacySessions []*ServerSession for _, sess := range s.sessions { - if sess.InitializeParams().isNil() || sess.InitializeParams().ProtocolVersion < protocolVersion20260728 { + if sess.negotiatedLegacyProtocol() { legacySessions = append(legacySessions, sess) } } @@ -1237,7 +1237,7 @@ func (s *Server) ResourceUpdated(ctx context.Context, params *ResourceUpdatedNot var legacySessions []*ServerSession newSessions := make(map[*ServerSession]jsonrpc.ID) for sess, reqID := range subscribedSessions { - if sess.InitializeParams().isNil() || sess.InitializeParams().ProtocolVersion < protocolVersion20260728 { + if sess.negotiatedLegacyProtocol() { legacySessions = append(legacySessions, sess) } else { newSessions[sess] = reqID @@ -1676,12 +1676,11 @@ func (ss *ServerSession) ID() string { // in an [InputRequiredResult] returned from a handler for one of the multi // round-trip methods (`tools/call`, `prompts/get`, `resources/read`). func (ss *ServerSession) assertServerInitiatedRequestAllowed(method string) error { - if iparams := ss.InitializeParams(); iparams != nil && - iparams.ProtocolVersion >= protocolVersion20260728 { + if version := ss.protocolVersion(); version >= protocolVersion20260728 { return fmt.Errorf( "%q cannot be sent while serving a request on protocol version %s: "+ "return an InputRequests map instead (multi round-trip requests, SEP-2322)", - method, iparams.ProtocolVersion) + method, version) } return nil } @@ -2149,6 +2148,29 @@ func (ss *ServerSession) InitializeParams() *InitializeParams { return ss.state.InitializeParams } +// protocolVersion returns the version the session speaks: the negotiated one, +// or the declared one when the state recorded none (a caller-supplied State, or +// state saved by an older release). It returns "" when neither is known. +func (ss *ServerSession) protocolVersion() string { + ss.mu.Lock() + defer ss.mu.Unlock() + if v := ss.state.NegotiatedProtocolVersion; v != "" { + return v + } + if ss.state.InitializeParams != nil { + return ss.state.InitializeParams.ProtocolVersion + } + return "" +} + +// negotiatedLegacyProtocol reports whether the session speaks a version older +// than protocolVersion20260728. A session with no recorded version is not +// legacy: without an 'initialize' handshake it is a SEP-2575 session. +func (ss *ServerSession) negotiatedLegacyProtocol() bool { + version := ss.protocolVersion() + return version != "" && version < protocolVersion20260728 +} + func (ss *ServerSession) initialize(ctx context.Context, params *InitializeParams) (*InitializeResult, error) { if params == nil { return nil, fmt.Errorf("%w: \"params\" must be be provided", jsonrpc2.ErrInvalidParams) diff --git a/mcp/server_test.go b/mcp/server_test.go index 64e7de0d..b935fc24 100644 --- a/mcp/server_test.go +++ b/mcp/server_test.go @@ -2692,3 +2692,192 @@ func TestServerUnknownProtocolVersion_NewProtocol(t *testing.T) { }) } } + +// rawSession is a JSON-RPC connection to a server, driven message by message +// so a test can run a handshake the SDK's own client does not offer. +type rawSession struct { + write func(jsonrpc.Message, error) + read func() jsonrpc.Message +} + +// connectNegotiatedDown runs initialize declaring protocolVersion20260728 and +// asserts it settles on protocolVersion20251125, so InitializeParams and the +// negotiated version disagree. +func connectNegotiatedDown(t *testing.T, ctx context.Context, srv *Server) *rawSession { + t.Helper() + + ct, st := NewInMemoryTransports() + ss, err := srv.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server.Connect() error = %v", err) + } + t.Cleanup(func() { ss.Close() }) + + conn, err := ct.Connect(ctx) + if err != nil { + t.Fatalf("transport.Connect() error = %v", err) + } + t.Cleanup(func() { conn.Close() }) + + sess := &rawSession{ + write: func(msg jsonrpc.Message, err error) { + t.Helper() + if err != nil { + t.Fatalf("building message: %v", err) + } + if err := conn.Write(ctx, msg); err != nil { + t.Fatalf("conn.Write() error = %v", err) + } + }, + read: func() jsonrpc.Message { + t.Helper() + msg, err := conn.Read(ctx) + if err != nil { + t.Fatalf("conn.Read() error = %v", err) + } + return msg + }, + } + + sess.write(jsonrpc2.NewCall(jsonrpc2.Int64ID(1), methodInitialize, &InitializeParams{ + ProtocolVersion: protocolVersion20260728, + ClientInfo: testImpl, + Capabilities: &ClientCapabilities{}, + })) + initResp, ok := sess.read().(*jsonrpc2.Response) + if !ok { + t.Fatalf("initialize was not answered with a response") + } + if initResp.Error != nil { + t.Fatalf("initialize failed: %v", initResp.Error) + } + var initRes InitializeResult + if err := json.Unmarshal(initResp.Result, &initRes); err != nil { + t.Fatalf("unmarshalling initialize result: %v", err) + } + if initRes.ProtocolVersion != protocolVersion20251125 { + t.Fatalf("negotiated protocol version = %q, want %q", initRes.ProtocolVersion, protocolVersion20251125) + } + sess.write(jsonrpc2.NewNotification(notificationInitialized, &InitializedParams{})) + return sess +} + +// TestNotifySessions_NegotiatedDownFromNewProtocol asserts that a session +// negotiated down from protocolVersion20260728 receives a list-changed +// notification on the shared session channel. +func TestNotifySessions_NegotiatedDownFromNewProtocol(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + srv := NewServer(testImpl, nil) + sess := connectNegotiatedDown(t, ctx, srv) + + srv.AddTool(&Tool{Name: "act", InputSchema: &jsonschema.Schema{Type: "object"}}, nil) + + msg := sess.read() + note, ok := msg.(*jsonrpc2.Request) + if !ok { + t.Fatalf("got %T, want the %q notification", msg, notificationToolListChanged) + } + if note.Method != notificationToolListChanged { + t.Fatalf("notification method = %q, want %q", note.Method, notificationToolListChanged) + } +} + +// TestResourceUpdated_NegotiatedDownFromNewProtocol asserts the same for a +// resource update, which must also carry no subscription id in _meta, since +// the negotiated version does not define one. +func TestResourceUpdated_NegotiatedDownFromNewProtocol(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + const uri = "test://resource" + srv := NewServer(testImpl, &ServerOptions{ + SubscribeHandler: func(context.Context, *SubscribeRequest) error { return nil }, + UnsubscribeHandler: func(context.Context, *UnsubscribeRequest) error { return nil }, + }) + srv.AddResource(&Resource{URI: uri, Name: "res"}, + func(context.Context, *ReadResourceRequest) (*ReadResourceResult, error) { + return &ReadResourceResult{Contents: []*ResourceContents{{URI: uri, Text: "data"}}}, nil + }) + + sess := connectNegotiatedDown(t, ctx, srv) + + sess.write(jsonrpc2.NewCall(jsonrpc2.Int64ID(2), methodSubscribe, &SubscribeParams{URI: uri})) + subResp, ok := sess.read().(*jsonrpc2.Response) + if !ok { + t.Fatalf("subscribe was not answered with a response") + } + if subResp.Error != nil { + t.Fatalf("subscribe failed: %v", subResp.Error) + } + + if err := srv.ResourceUpdated(ctx, &ResourceUpdatedNotificationParams{URI: uri}); err != nil { + t.Fatalf("ResourceUpdated() error = %v", err) + } + + msg := sess.read() + note, ok := msg.(*jsonrpc2.Request) + if !ok { + t.Fatalf("got %T, want the %q notification", msg, notificationResourceUpdated) + } + if note.Method != notificationResourceUpdated { + t.Fatalf("notification method = %q, want %q", note.Method, notificationResourceUpdated) + } + var params ResourceUpdatedNotificationParams + if err := json.Unmarshal(note.Params, ¶ms); err != nil { + t.Fatalf("unmarshalling notification params: %v", err) + } + if params.URI != uri { + t.Errorf("notification URI = %q, want %q", params.URI, uri) + } + if _, ok := params.GetMeta()[MetaKeySubscriptionID]; ok { + t.Errorf("notification carries %q, which protocol version %s does not define", + MetaKeySubscriptionID, protocolVersion20251125) + } +} + +// TestSpeaksLegacyProtocol_NoHandshakeIsNotLegacy pins that a session with no +// recorded protocol version reads as the current protocol, since a session +// without an initialize handshake is a SEP-2575 session. +func TestSpeaksLegacyProtocol_NoHandshakeIsNotLegacy(t *testing.T) { + tests := []struct { + name string + state ServerSessionState + want bool + }{ + {name: "no handshake", want: false}, + { + name: "discover, new protocol", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20260728}, + }, + want: false, + }, + { + name: "initialize, legacy version", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20251125}, + NegotiatedProtocolVersion: protocolVersion20251125, + }, + want: true, + }, + { + name: "initialize, negotiated down", + state: ServerSessionState{ + InitializeParams: &InitializeParams{ProtocolVersion: protocolVersion20260728}, + NegotiatedProtocolVersion: protocolVersion20251125, + }, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ss := &ServerSession{state: tt.state} + if got := ss.negotiatedLegacyProtocol(); got != tt.want { + t.Errorf("negotiatedLegacyProtocol() = %t, want %t", got, tt.want) + } + }) + } +}