diff --git a/mcp/server.go b/mcp/server.go index 442d7bc6..b2f89f48 100644 --- a/mcp/server.go +++ b/mcp/server.go @@ -1989,10 +1989,6 @@ func (ss *ServerSession) handle(ctx context.Context, req *jsonrpc.Request) (any, } } default: - if !initialized && !validatedMeta.usesNewProtocol && req.IsCall() { - ss.server.opts.Logger.Error("method invalid during initialization", "method", req.Method) - return nil, fmt.Errorf("method %q is invalid during session initialization", req.Method) - } if !initialized && validatedMeta.usesNewProtocol && validatedMeta.initializeParams != nil { ss.updateState(func(state *ServerSessionState) { state.InitializeParams = validatedMeta.initializeParams @@ -2000,6 +1996,16 @@ func (ss *ServerSession) handle(ctx context.Context, req *jsonrpc.Request) (any, } } + // In legacy protocol versions, a client cannot send requests other than + // pings before the server has responded to the initialize request. A + // new-protocol request is exempt: a SEP-2575 session has no 'initialize' + // to wait for, and the request itself says which protocol it speaks. + if !initialized && !validatedMeta.usesNewProtocol && req.IsCall() && + req.Method != methodInitialize && req.Method != methodPing { + ss.server.opts.Logger.Error("method invalid during initialization", "method", req.Method) + return nil, fmt.Errorf("method %q is invalid during session initialization", req.Method) + } + // modelcontextprotocol/go-sdk#26: handle calls asynchronously, and // notifications synchronously, except for 'initialize' which shouldn't be // asynchronous to other diff --git a/mcp/server_test.go b/mcp/server_test.go index 2c7bbe05..4d09cbca 100644 --- a/mcp/server_test.go +++ b/mcp/server_test.go @@ -2198,6 +2198,162 @@ func TestServerSupportedProtocolVersions_NewProtocol(t *testing.T) { } } +// TestServerHandle_LegacyCallBeforeInitialize asserts that a call carrying no +// new-protocol `_meta` is refused on a session that has not completed +// 'initialize', whichever branch of the method switch it lands in. The +// handshake itself is what ends that state, and the lifecycle spec exempts +// pings, so those two are served. +// +// The gate used to sit in the default branch, so 'logging/setLevel', +// 'resources/subscribe' and 'resources/unsubscribe', which share a case that +// only refuses new-protocol requests, were served before the handshake. For +// subscribe that registered a session that never handshook as a subscriber, +// so the subtest also asserts that the refusal leaves no subscriber and runs +// no handler. +func TestServerHandle_LegacyCallBeforeInitialize(t *testing.T) { + ctx := context.Background() + const uri = "test://resource" + + tests := []struct { + name string + method string + params string + served bool + }{ + { + name: "initialize", + method: methodInitialize, + params: fmt.Sprintf(`{"protocolVersion":%q,"capabilities":{},"clientInfo":{"name":"c","version":"1"}}`, protocolVersion20251125), + served: true, + }, + {name: "ping", method: methodPing, params: `{}`, served: true}, + {name: "tools/list", method: methodListTools, params: `{}`}, + {name: "logging/setLevel", method: methodSetLevel, params: `{"level":"info"}`}, + {name: "resources/subscribe", method: methodSubscribe, params: fmt.Sprintf(`{"uri":%q}`, uri)}, + {name: "resources/unsubscribe", method: methodUnsubscribe, params: fmt.Sprintf(`{"uri":%q}`, uri)}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + var handlerRan bool + server := NewServer(testImpl, &ServerOptions{ + SubscribeHandler: func(context.Context, *SubscribeRequest) error { + handlerRan = true + return nil + }, + UnsubscribeHandler: func(context.Context, *UnsubscribeRequest) error { + handlerRan = true + return nil + }, + }) + _, st := NewInMemoryTransports() + ss, err := server.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server.Connect: %v", err) + } + defer ss.Close() + + res, err := ss.handle(ctx, &jsonrpc.Request{ + ID: jsonrpc2.Int64ID(1), + Method: tc.method, + Params: json.RawMessage(tc.params), + }) + if tc.served { + if err != nil { + t.Fatalf("handle(%q) before initialize returned %v, want it served", tc.method, err) + } + if res == nil { + t.Fatalf("handle(%q) before initialize returned no result", tc.method) + } + if tc.method == methodInitialize && ss.InitializeParams() == nil { + t.Errorf("initialize was answered without recording InitializeParams") + } + return + } + if err == nil { + t.Fatalf("handle(%q) served the call before initialize, want a refusal", tc.method) + } + if want := fmt.Sprintf("method %q is invalid during session initialization", tc.method); err.Error() != want { + t.Errorf("handle(%q) error = %q, want %q", tc.method, err.Error(), want) + } + if handlerRan { + t.Errorf("handle(%q) ran the server's handler for a call it refused", tc.method) + } + server.mu.Lock() + _, subscribed := server.resourceSubscriptions[uri][ss] + server.mu.Unlock() + if subscribed { + t.Errorf("handle(%q) left the session subscribed to %q", tc.method, uri) + } + }) + } +} + +// TestServerHandle_NewProtocolCallWithoutInitialize asserts that the gate +// stays off the SEP-2575 path. A fresh session that never sends 'initialize' +// is served a new-protocol 'tools/list', and the methods the gate now covers +// are answered, when a new-protocol request names them, with the error the +// new protocol already gave them (the method does not exist there), and not +// with the initialization refusal. +func TestServerHandle_NewProtocolCallWithoutInitialize(t *testing.T) { + ctx := context.Background() + const uri = "test://resource" + + tests := []struct { + name string + method string + params map[string]any + served bool + }{ + {name: "tools/list", method: methodListTools, params: newProtocolParams(nil), served: true}, + {name: "logging/setLevel", method: methodSetLevel, params: newProtocolParams(map[string]any{"level": "info"})}, + {name: "resources/subscribe", method: methodSubscribe, params: newProtocolParams(map[string]any{"uri": uri})}, + {name: "resources/unsubscribe", method: methodUnsubscribe, params: newProtocolParams(map[string]any{"uri": uri})}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + server := NewServer(testImpl, &ServerOptions{ + SubscribeHandler: func(context.Context, *SubscribeRequest) error { return nil }, + UnsubscribeHandler: func(context.Context, *UnsubscribeRequest) error { return nil }, + }) + _, st := NewInMemoryTransports() + ss, err := server.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server.Connect: %v", err) + } + defer ss.Close() + + params, err := json.Marshal(tc.params) + if err != nil { + t.Fatalf("marshalling params: %v", err) + } + res, err := ss.handle(ctx, &jsonrpc.Request{ + ID: jsonrpc2.Int64ID(1), + Method: tc.method, + Params: params, + }) + if tc.served { + if err != nil { + t.Fatalf("handle(%q) on a new-protocol session returned %v, want it served", tc.method, err) + } + if res == nil { + t.Fatalf("handle(%q) on a new-protocol session returned no result", tc.method) + } + return + } + var jerr *jsonrpc.Error + if !errors.As(err, &jerr) { + t.Fatalf("handle(%q) returned %v, want a *jsonrpc.Error", tc.method, err) + } + if jerr.Code != jsonrpc.CodeMethodNotFound { + t.Errorf("handle(%q) error code = %d, want %d (removed in the new protocol)", tc.method, jerr.Code, jsonrpc.CodeMethodNotFound) + } + if strings.Contains(err.Error(), "invalid during session initialization") { + t.Errorf("handle(%q) applied the initialization gate to a new-protocol request: %v", tc.method, err) + } + }) + } +} + // TestServerUnknownProtocolVersion_NewProtocol verifies that a request whose // `_meta.protocolVersion` names a version the SDK does not know is rejected // with [CodeUnsupportedProtocolVersion], and not served as a legacy handshake