Skip to content
Merged
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
14 changes: 10 additions & 4 deletions mcp/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -1989,17 +1989,23 @@ 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
})
}
}

// 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)
Comment thread
guglielmo-san marked this conversation as resolved.
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
Expand Down
156 changes: 156 additions & 0 deletions mcp/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading