Skip to content
Open
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
82 changes: 82 additions & 0 deletions internal/node/mcp_discovery_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,16 @@

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"
Expand Down Expand Up @@ -58,6 +63,83 @@
}
}

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()

Check failure on line 132 in internal/node/mcp_discovery_test.go

View workflow job for this annotation

GitHub Actions / lint

Error return value of `session.Close` is not checked (errcheck)

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
Expand Down
25 changes: 11 additions & 14 deletions internal/node/mcp_handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{{
Expand All @@ -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
}
Expand Down Expand Up @@ -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
}
Expand All @@ -664,6 +657,10 @@ func (n *SamNode) fetchRemoteToolDescription(ctx context.Context, pid peer.ID, t
OutputSchema: tool.OutputSchema,
}, nil
}
count++
if count >= maxToolsPerService {
break
}
Comment on lines +661 to +663

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

For consistency with listAllTools and to aid in troubleshooting, consider logging a debug message when the tool search is truncated due to reaching the maxToolsPerService cap. This helps operators understand if a tool lookup failed because the tool list was truncated.

		if count >= maxToolsPerService {
			logger.Debugf("reached cap of %d tools while searching for %q; truncating", maxToolsPerService, toolName)
			break
		}

}

return nil, fmt.Errorf("tool not found on peer")
Expand Down
211 changes: 210 additions & 1 deletion internal/node/mcp_handlers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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) {
Expand Down
Loading
Loading