From 5f04aec3ee3885484b8de569cd70612bf26d0b99 Mon Sep 17 00:00:00 2001 From: bodapatisaikrishna Date: Sat, 3 Oct 2026 17:50:07 +0530 Subject: [PATCH] node: follow MCP tool-list pagination in discovery and remote tool lookups Drain the MCP Tools iterator instead of single-page ListTools across discovery and find_remote_tools, and stream describe_remote_tool with early return. Cap per-service tool drain at 256 to bound against hostile endless cursor loops under zero trust. --- internal/node/mcp_discovery_test.go | 82 +++++++++++ internal/node/mcp_handlers.go | 25 ++-- internal/node/mcp_handlers_test.go | 211 +++++++++++++++++++++++++++- internal/node/mcp_service.go | 32 ++++- 4 files changed, 332 insertions(+), 18 deletions(-) diff --git a/internal/node/mcp_discovery_test.go b/internal/node/mcp_discovery_test.go index 127304b7..b1e3ee0d 100644 --- a/internal/node/mcp_discovery_test.go +++ b/internal/node/mcp_discovery_test.go @@ -16,11 +16,16 @@ package node import ( "context" + "encoding/json" "fmt" + "io" + "net/http" "net/http/httptest" "reflect" "sort" + "sync/atomic" "testing" + "time" "github.com/google/sam/api" samdiscovery "github.com/google/sam/internal/node/discovery" @@ -58,6 +63,83 @@ func TestMCPService_Tools(t *testing.T) { } } +func TestMCPService_Tools_Pagination(t *testing.T) { + backend := httptest.NewServer(newFakeMCPHandlerWithOptions(t, []*mcp.Tool{ + {Name: "zeta", Description: "z", InputSchema: map[string]any{"type": "object"}}, + {Name: "alpha", Description: "a", InputSchema: map[string]any{"type": "object"}}, + {Name: "beta", Description: "b", InputSchema: map[string]any{"type": "object"}}, + }, &mcp.ServerOptions{PageSize: 1})) + defer backend.Close() + + svc := &MCPService{baseService: baseService{ + info: &api.ServiceInfo{Type: api.ServiceType_SERVICE_TYPE_MCP, Name: "tools-svc"}, + backend: &api.RegisterServiceRequest_TargetUrl{TargetUrl: backend.URL}, + }} + if err := svc.Init(context.Background()); err != nil { + t.Fatalf("Init: %v", err) + } + + got, err := svc.Tools(context.Background()) + if err != nil { + t.Fatalf("Tools: %v", err) + } + if want := []string{"alpha", "beta", "zeta"}; !reflect.DeepEqual(got, want) { + t.Errorf("Tools: got %v, want %v (sorted across pages)", got, want) + } +} + +func TestListAllTools_BoundedAgainstEndlessCursor(t *testing.T) { + var pages atomic.Int64 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + var req struct { + ID any `json:"id"` + Method string `json:"method"` + } + _ = json.Unmarshal(body, &req) + if req.ID == nil { + w.WriteHeader(http.StatusAccepted) + return + } // notification + var result any = map[string]any{} + switch req.Method { + case "initialize": + result = map[string]any{ + "protocolVersion": "2026-07-28", + "capabilities": map[string]any{"tools": map[string]any{}}, + "serverInfo": map[string]any{"name": "hostile", "version": "0"}, + } + case "tools/list": + pages.Add(1) + result = map[string]any{ + "tools": []any{map[string]any{"name": "x", "inputSchema": map[string]any{"type": "object"}}}, + "nextCursor": "again", + "ttlMs": 60000, "cacheScope": "public", // makes the SDK serve repeats from cache + } + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"jsonrpc": "2.0", "id": req.ID, "result": result}) + })) + defer srv.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + session, err := mcp.NewClient(&mcp.Implementation{Name: "t", Version: "0"}, nil). + Connect(ctx, &mcp.StreamableClientTransport{Endpoint: srv.URL}, nil) + if err != nil { + t.Fatal(err) + } + defer session.Close() + + tools, err := listAllTools(ctx, session) + if err != nil { + t.Fatal(err) + } + if len(tools) != maxToolsPerService { + t.Fatalf("got %d tools, want cap %d", len(tools), maxToolsPerService) + } +} + // fakeToolService is a local MCP Service reporting served tool names. type fakeToolService struct { testService diff --git a/internal/node/mcp_handlers.go b/internal/node/mcp_handlers.go index e52cd401..1da5234d 100644 --- a/internal/node/mcp_handlers.go +++ b/internal/node/mcp_handlers.go @@ -527,7 +527,7 @@ func (n *SamNode) fetchToolsForRemoteService( } defer cleanup() - listRes, err := session.ListTools(ctx, nil) + tools, err := listAllTools(ctx, session) if err != nil { if serviceNameFilter == "" || connectService == serviceNameFilter { return []remoteToolRow{{ @@ -538,11 +538,8 @@ func (n *SamNode) fetchToolsForRemoteService( } return nil } - if listRes == nil { - return nil - } var rows []remoteToolRow - for _, t := range listRes.Tools { + for _, t := range tools { if t == nil { continue } @@ -643,15 +640,11 @@ func (n *SamNode) fetchRemoteToolDescription(ctx context.Context, pid peer.ID, t } defer cleanup() - listRes, err := session.ListTools(ctx, nil) - if err != nil { - return nil, err - } - if listRes == nil { - return nil, fmt.Errorf("list tools response was nil") - } - - for _, tool := range listRes.Tools { + count := 0 + for tool, err := range session.Tools(ctx, nil) { + if err != nil { + return nil, err + } if tool == nil { continue } @@ -664,6 +657,10 @@ func (n *SamNode) fetchRemoteToolDescription(ctx context.Context, pid peer.ID, t OutputSchema: tool.OutputSchema, }, nil } + count++ + if count >= maxToolsPerService { + break + } } return nil, fmt.Errorf("tool not found on peer") diff --git a/internal/node/mcp_handlers_test.go b/internal/node/mcp_handlers_test.go index 7fab2568..b4e3fa21 100644 --- a/internal/node/mcp_handlers_test.go +++ b/internal/node/mcp_handlers_test.go @@ -217,6 +217,71 @@ func TestHandleFindRemoteTools_SinglePeer(t *testing.T) { } } +func TestHandleFindRemoteTools_Pagination(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + tools := []*mcp.Tool{ + {Name: "review_pr", Description: "Run a code review", InputSchema: map[string]any{"type": "object"}}, + {Name: "add_comment", Description: "Add a comment", InputSchema: map[string]any{"type": "object"}}, + } + hostedSrv := httptest.NewServer(newFakeMCPHandlerWithOptions(t, tools, &mcp.ServerOptions{PageSize: 1})) + defer hostedSrv.Close() + + nodeA, cleanupA := startBareNode(t, ctx) + defer cleanupA() + nodeB, cleanupB := startBareNode(t, ctx) + defer cleanupB() + + if err := nodeA.Host.Connect(ctx, peer.AddrInfo{ID: nodeB.Host.ID(), Addrs: nodeB.Host.Addrs()}); err != nil { + t.Fatalf("connect: %v", err) + } + + enrollUnderRoot(t, nodeA, nodeB) + + // Register an MCP service on B with two tools paginated across pages. + regReq := &api.RegisterServiceRequest{ + Service: &api.ServiceInfo{Type: api.ServiceType_SERVICE_TYPE_MCP, Name: "code-reviewer"}, + Backend: &api.RegisterServiceRequest_TargetUrl{TargetUrl: hostedSrv.URL}, + } + if err := nodeB.RegisterService(ctx, regReq); err != nil { + t.Fatalf("RegisterService: %v", err) + } + + res, _, err := nodeA.handleFindRemoteTools(ctx, &mcp.CallToolRequest{}, FindRemoteToolsParams{ + PeerID: nodeB.Host.ID().String(), + }) + if err != nil { + t.Fatalf("handleFindRemoteTools: %v", err) + } + tc, ok := res.Content[0].(*mcp.TextContent) + if !ok { + t.Fatalf("expected TextContent, got %T", res.Content[0]) + } + var rows []remoteToolRow + if err := json.Unmarshal([]byte(tc.Text), &rows); err != nil { + t.Fatalf("unmarshal: %v (text: %q)", err, tc.Text) + } + + wantNames := map[string]bool{ + "mcp://code-reviewer/review_pr": false, + "mcp://code-reviewer/add_comment": false, + } + for _, row := range rows { + if row.PeerID != nodeB.Host.ID().String() { + t.Errorf("row has peer_id %q, want %q", row.PeerID, nodeB.Host.ID().String()) + } + if _, ok := wantNames[row.ToolName]; ok { + wantNames[row.ToolName] = true + } + } + for name, found := range wantNames { + if !found { + t.Errorf("expected tool %q in response, not found across pages; rows=%+v", name, rows) + } + } +} + // TestHandleFindRemoteTools_BackendPredatesDiscover is a regression test for // a go-sdk v1.7.0 (SEP-2575) incompatibility: mcp.Client.Connect() sends a // "server/discover" preflight before "initialize", falling back to the @@ -1030,6 +1095,143 @@ func TestHandleDescribeRemoteTool_RoundTrip(t *testing.T) { } } +func TestHandleDescribeRemoteTool_Pagination(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + nodeA, cleanupA := startBareNode(t, ctx) + defer cleanupA() + nodeB, cleanupB := startBareNode(t, ctx) + defer cleanupB() + + if err := nodeA.Host.Connect(ctx, peer.AddrInfo{ID: nodeB.Host.ID(), Addrs: nodeB.Host.Addrs()}); err != nil { + t.Fatalf("connect: %v", err) + } + + enrollUnderRoot(t, nodeA, nodeB) + + // Tools are served with PageSize: 1, with "alpha" ahead of "review_pr" so + // review_pr lives on a subsequent page. + tools := []*mcp.Tool{ + {Name: "alpha", Description: "First tool", InputSchema: map[string]any{"type": "object"}}, + { + Name: "review_pr", + Description: "Run a code review", + InputSchema: map[string]any{ + "type": "object", + "required": []any{"pr_url"}, + "properties": map[string]any{ + "pr_url": map[string]any{"type": "string"}, + }, + }, + OutputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "summary": map[string]any{"type": "string"}, + }, + }, + }, + } + hostedSrv := httptest.NewServer(newFakeMCPHandlerWithOptions(t, tools, &mcp.ServerOptions{PageSize: 1})) + defer hostedSrv.Close() + + regReq := &api.RegisterServiceRequest{ + Service: &api.ServiceInfo{Type: api.ServiceType_SERVICE_TYPE_MCP, Name: "code-reviewer"}, + Backend: &api.RegisterServiceRequest_TargetUrl{TargetUrl: hostedSrv.URL}, + } + if err := nodeB.RegisterService(ctx, regReq); err != nil { + t.Fatalf("RegisterService: %v", err) + } + defer func() { _ = nodeB.UnregisterService(ctx, "code-reviewer") }() + + res, _, err := nodeA.handleDescribeRemoteTool(ctx, &mcp.CallToolRequest{}, DescribeRemoteToolParams{ + PeerID: nodeB.Host.ID().String(), + ToolName: "mcp://code-reviewer/review_pr", + }) + if err != nil { + t.Fatalf("handleDescribeRemoteTool on paginated backend: %v", err) + } + tc, ok := res.Content[0].(*mcp.TextContent) + if !ok { + t.Fatalf("expected TextContent, got %T", res.Content[0]) + } + + var desc remoteToolDescription + if err := json.Unmarshal([]byte(tc.Text), &desc); err != nil { + t.Fatalf("unmarshal: %v (text: %q)", err, tc.Text) + } + if desc.ToolName != "mcp://code-reviewer/review_pr" { + t.Errorf("ToolName = %q, want %q", desc.ToolName, "mcp://code-reviewer/review_pr") + } + if desc.Description != "Run a code review" { + t.Errorf("Description = %q, want %q", desc.Description, "Run a code review") + } +} + +func TestHandleDescribeRemoteTool_BoundedAgainstEndlessCursor(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + hostileSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + var req struct { + ID any `json:"id"` + Method string `json:"method"` + } + _ = json.Unmarshal(body, &req) + if req.ID == nil { + w.WriteHeader(http.StatusAccepted) + return + } + var result any = map[string]any{} + switch req.Method { + case "initialize": + result = map[string]any{ + "protocolVersion": "2026-07-28", + "capabilities": map[string]any{"tools": map[string]any{}}, + "serverInfo": map[string]any{"name": "hostile", "version": "0"}, + } + case "tools/list": + result = map[string]any{ + "tools": []any{map[string]any{"name": "other", "inputSchema": map[string]any{"type": "object"}}}, + "nextCursor": "again", + "ttlMs": 60000, "cacheScope": "public", + } + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"jsonrpc": "2.0", "id": req.ID, "result": result}) + })) + defer hostileSrv.Close() + + nodeA, cleanupA := startBareNode(t, ctx) + defer cleanupA() + nodeB, cleanupB := startBareNode(t, ctx) + defer cleanupB() + + if err := nodeA.Host.Connect(ctx, peer.AddrInfo{ID: nodeB.Host.ID(), Addrs: nodeB.Host.Addrs()}); err != nil { + t.Fatalf("connect: %v", err) + } + + enrollUnderRoot(t, nodeA, nodeB) + + regReq := &api.RegisterServiceRequest{ + Service: &api.ServiceInfo{Type: api.ServiceType_SERVICE_TYPE_MCP, Name: "hostile-svc"}, + Backend: &api.RegisterServiceRequest_TargetUrl{TargetUrl: hostileSrv.URL}, + } + if err := nodeB.RegisterService(ctx, regReq); err != nil { + t.Fatalf("RegisterService: %v", err) + } + defer func() { _ = nodeB.UnregisterService(ctx, "hostile-svc") }() + + _, _, err := nodeA.handleDescribeRemoteTool(ctx, &mcp.CallToolRequest{}, DescribeRemoteToolParams{ + PeerID: nodeB.Host.ID().String(), + ToolName: "mcp://hostile-svc/wanted_tool", + }) + if err == nil || !strings.Contains(err.Error(), "tool not found on peer") { + t.Fatalf("expected 'tool not found on peer', got: %v", err) + } +} + func TestHandleDescribeRemoteTool_RoundTrip_UnknownTool(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() @@ -1112,7 +1314,14 @@ func TestNewMCPHandler_RegistersDescribeRemoteTool(t *testing.T) { // streamable-http with the given tools registered. func newFakeMCPHandler(t *testing.T, tools []*mcp.Tool) http.Handler { t.Helper() - srv := mcp.NewServer(&mcp.Implementation{Name: "fake", Version: "0.0.1"}, nil) + return newFakeMCPHandlerWithOptions(t, tools, nil) +} + +// newFakeMCPHandlerWithOptions returns an http.Handler serving a tiny MCP server over +// streamable-http with the given tools registered and custom server options. +func newFakeMCPHandlerWithOptions(t *testing.T, tools []*mcp.Tool, opts *mcp.ServerOptions) http.Handler { + t.Helper() + srv := mcp.NewServer(&mcp.Implementation{Name: "fake", Version: "0.0.1"}, opts) for _, tool := range tools { toolCopy := tool srv.AddTool(toolCopy, func(ctx context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) { diff --git a/internal/node/mcp_service.go b/internal/node/mcp_service.go index a7873c10..8dfa4241 100644 --- a/internal/node/mcp_service.go +++ b/internal/node/mcp_service.go @@ -160,6 +160,32 @@ func (m *MCPService) backendTransport() (mcp.Transport, error) { } } +// maxToolsPerService is the maximum number of tools collected per service +// during discovery or remote catalogue listing. It bounds the drain against +// backends with huge catalogues or endless cursor loops (#444). +const maxToolsPerService = 256 + +// listAllTools drains the SDK's Tools iterator up to maxToolsPerService. +// If the backend paginates beyond maxToolsPerService, iteration stops, +// a debug message is logged, and the collected tools are returned. +func listAllTools(ctx context.Context, session *mcp.ClientSession) ([]*mcp.Tool, error) { + var tools []*mcp.Tool + for tool, err := range session.Tools(ctx, nil) { + if err != nil { + return nil, err + } + if tool == nil { + continue + } + tools = append(tools, tool) + if len(tools) >= maxToolsPerService { + logger.Debugf("tools list reached cap of %d tools; truncating remainder", maxToolsPerService) + break + } + } + return tools, nil +} + // Tools lists the backend's tool names (sorted), cached briefly since the // discovery announcer polls it on every tick. func (m *MCPService) Tools(ctx context.Context) ([]string, error) { @@ -178,12 +204,12 @@ func (m *MCPService) Tools(ctx context.Context) ([]string, error) { return nil, fmt.Errorf("connect to backend of %q: %w", m.info.GetName(), err) } defer func() { _ = session.Close() }() - res, err := session.ListTools(ctx, nil) + tools, err := listAllTools(ctx, session) if err != nil { return nil, fmt.Errorf("list tools of %q: %w", m.info.GetName(), err) } - names := make([]string, 0, len(res.Tools)) - for _, t := range res.Tools { + names := make([]string, 0, len(tools)) + for _, t := range tools { if t != nil && t.Name != "" { names = append(names, t.Name) }