diff --git a/internal/mcp/network_client.go b/internal/mcp/network_client.go index b8b829af2..ce0b41162 100644 --- a/internal/mcp/network_client.go +++ b/internal/mcp/network_client.go @@ -216,6 +216,9 @@ func (client *networkClient) request(ctx context.Context, method string, params if err != nil { return err } + if message.isRequestOrNotification() { + return fmt.Errorf("MCP %s expected a response from server %s, got method %q instead", method, client.server.Name, message.Method) + } if !rpcIDMatches(message.ID, id) { return fmt.Errorf("MCP %s response id mismatch for server %s", method, client.server.Name) } diff --git a/internal/mcp/network_client_test.go b/internal/mcp/network_client_test.go index e4b406447..8cb0d2f13 100644 --- a/internal/mcp/network_client_test.go +++ b/internal/mcp/network_client_test.go @@ -316,6 +316,50 @@ func TestNonOAuthServerIsUnaffected(t *testing.T) { } } +func TestNetworkClientRejectsServerInitiatedRequestAsResponse(t *testing.T) { + // A streamable-HTTP server that answers a tools/call POST with a + // method-bearing frame (a server-initiated request) must surface a protocol + // error rather than an empty, error-free result. + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + message := readHTTPRPCMessage(t, r) + switch message.Method { + case "initialize": + writeHTTPRPCResponse(t, w, message.ID, map[string]any{"protocolVersion": "2024-11-05"}) + case "notifications/initialized": + w.WriteHeader(http.StatusAccepted) + case "tools/call": + w.Header().Set("Content-Type", "application/json") + if _, err := w.Write([]byte(`{"jsonrpc":"2.0","id":` + string(mustRaw(message.ID)) + `,"method":"roots/list","params":{}}`)); err != nil { + t.Errorf("write server request: %v", err) + } + default: + writeHTTPRPCResponse(t, w, message.ID, map[string]any{}) + } + })) + defer upstream.Close() + + client, err := Connect(ctx, Server{Name: "plain", Type: ServerTypeHTTP, URL: upstream.URL}) + if err != nil { + t.Fatalf("Connect() error = %v", err) + } + defer func() { + if err := client.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + }() + + result, err := client.CallTool(ctx, "lookup", map[string]any{"query": "zero"}) + if err == nil { + t.Fatalf("CallTool() error = nil, result = %#v, want protocol error", result) + } + if !strings.Contains(err.Error(), `"roots/list"`) { + t.Fatalf("CallTool() error = %q, want it to name the unexpected method", err) + } +} + func TestDecodeSSERPCMessageSkipsNotifications(t *testing.T) { // A leading server notification (has a method) on the POST's event stream must // be skipped so the actual response (no method) is returned, instead of the