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
6 changes: 3 additions & 3 deletions internal/node/mcp_discovery_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,11 +27,11 @@ import (
"github.com/modelcontextprotocol/go-sdk/mcp"
)

func TestMCPService_Tools(t *testing.T) {
backend := httptest.NewServer(newFakeMCPHandler(t, []*mcp.Tool{
func TestMCPService_ToolsPagination(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"}},
}))
}, &mcp.ServerOptions{PageSize: 1}))
defer backend.Close()

svc := &MCPService{baseService: baseService{
Expand Down
14 changes: 4 additions & 10 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,12 @@ func (n *SamNode) fetchRemoteToolDescription(ctx context.Context, pid peer.ID, t
}
defer cleanup()

listRes, err := session.ListTools(ctx, nil)
tools, err := listAllTools(ctx, session)
if err != nil {
return nil, err
}
if listRes == nil {
return nil, fmt.Errorf("list tools response was nil")
}

for _, tool := range listRes.Tools {
for _, tool := range tools {
if tool == nil {
continue
}
Expand Down
15 changes: 10 additions & 5 deletions internal/node/mcp_handlers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -155,15 +155,15 @@ func contextWithShortTimeout() (context.Context, context.CancelFunc) {
return context.WithTimeout(context.Background(), 3*time.Second)
}

func TestHandleFindRemoteTools_SinglePeer(t *testing.T) {
func TestRemoteToolCataloguePagination(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(newFakeMCPHandler(t, tools))
hostedSrv := httptest.NewServer(newFakeMCPHandlerWithOptions(t, tools, &mcp.ServerOptions{PageSize: 1}))
defer hostedSrv.Close()

nodeA, cleanupA := startBareNode(t, ctx)
Expand Down Expand Up @@ -914,7 +914,7 @@ func TestHandleDescribeRemoteTool_InvalidPeerID(t *testing.T) {
}
}

func TestHandleDescribeRemoteTool_RoundTrip(t *testing.T) {
func TestRemoteToolDescriptionPagination(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()

Expand All @@ -930,6 +930,7 @@ func TestHandleDescribeRemoteTool_RoundTrip(t *testing.T) {
enrollUnderRoot(t, nodeA, nodeB)

tools := []*mcp.Tool{
{Name: "alpha", Description: "First page", InputSchema: map[string]any{"type": "object"}},
{
Name: "review_pr",
Description: "Run a code review",
Expand All @@ -948,7 +949,7 @@ func TestHandleDescribeRemoteTool_RoundTrip(t *testing.T) {
},
},
}
hostedSrv := httptest.NewServer(newFakeMCPHandler(t, tools))
hostedSrv := httptest.NewServer(newFakeMCPHandlerWithOptions(t, tools, &mcp.ServerOptions{PageSize: 1}))
defer hostedSrv.Close()

regReq := &api.RegisterServiceRequest{
Expand Down Expand Up @@ -1082,8 +1083,12 @@ func TestNewMCPHandler_RegistersDescribeRemoteTool(t *testing.T) {
// newFakeMCPHandler returns an http.Handler serving a tiny MCP server over
// streamable-http with the given tools registered.
func newFakeMCPHandler(t *testing.T, tools []*mcp.Tool) http.Handler {
return newFakeMCPHandlerWithOptions(t, tools, nil)
}

func newFakeMCPHandlerWithOptions(t *testing.T, tools []*mcp.Tool, options *mcp.ServerOptions) http.Handler {
t.Helper()
srv := mcp.NewServer(&mcp.Implementation{Name: "fake", Version: "0.0.1"}, nil)
srv := mcp.NewServer(&mcp.Implementation{Name: "fake", Version: "0.0.1"}, options)
for _, tool := range tools {
toolCopy := tool
srv.AddTool(toolCopy, func(ctx context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) {
Expand Down
17 changes: 14 additions & 3 deletions internal/node/mcp_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -178,12 +178,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)
}
Expand All @@ -194,6 +194,17 @@ func (m *MCPService) Tools(ctx context.Context) ([]string, error) {
return names, nil
}

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
}
tools = append(tools, tool)
}
return tools, nil
}
Comment on lines +197 to +206

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

To ensure defensive programming and prevent potential nil pointer dereference panics, we should add a guard check to verify that the session parameter is not nil before calling session.Tools(ctx, nil).

Suggested change
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
}
tools = append(tools, tool)
}
return tools, nil
}
func listAllTools(ctx context.Context, session *mcp.ClientSession) ([]*mcp.Tool, error) {
if session == nil {
return nil, fmt.Errorf("mcp session is nil")
}
var tools []*mcp.Tool
for tool, err := range session.Tools(ctx, nil) {
if err != nil {
return nil, err
}
tools = append(tools, tool)
}
return tools, nil
}


// preflightMethodsUnsupportedByPassThrough lists stateless MCP capability
// probes that HandleStreamPassThrough answers locally instead of forwarding.
// It opens a fresh, sessionless backend connection per stream, so it can
Expand Down
Loading