From 9df1e3b0e2001138a32b62bb47176add50e86e8e Mon Sep 17 00:00:00 2001 From: Sam Morrow Date: Wed, 19 Aug 2026 11:45:08 +0200 Subject: [PATCH] fix(tools): allow omitted tool arguments Normalize missing or zero-length tool arguments to an empty object while preserving invalid JSON and required-parameter validation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pkg/github/context_tools_test.go | 30 ++++++ pkg/inventory/server_tool.go | 28 ++++-- pkg/inventory/server_tool_test.go | 157 ++++++++++++++++++------------ 3 files changed, 147 insertions(+), 68 deletions(-) diff --git a/pkg/github/context_tools_test.go b/pkg/github/context_tools_test.go index 65c4741a4b..0825158abb 100644 --- a/pkg/github/context_tools_test.go +++ b/pkg/github/context_tools_test.go @@ -11,6 +11,7 @@ import ( "github.com/github/github-mcp-server/internal/toolsnaps" "github.com/github/github-mcp-server/pkg/translations" "github.com/google/go-github/v89/github" + "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/shurcooL/githubv4" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -139,6 +140,35 @@ func Test_GetMe(t *testing.T) { } } +func Test_GetMe_OmittedArguments(t *testing.T) { + t.Parallel() + + mockUser := &github.User{ + Login: github.Ptr("testuser"), + HTMLURL: github.Ptr("https://github.com/testuser"), + CreatedAt: &github.Timestamp{Time: time.Now()}, + } + mockedClient := MockHTTPClientWithHandlers(map[string]http.HandlerFunc{ + GetUser: mockResponse(t, http.StatusOK, mockUser), + }) + deps := BaseDeps{Client: mustNewGHClient(t, mockedClient), Obsv: stubExporters()} + serverTool := GetMe(translations.NullTranslationHelper) + handler := serverTool.Handler(deps) + request := mcp.CallToolRequest{ + Params: &mcp.CallToolParamsRaw{Name: "get_me"}, + } + + result, err := handler(ContextWithDeps(context.Background(), deps), &request) + + require.NoError(t, err) + require.False(t, result.IsError) + textContent := getTextResult(t, result) + var returnedUser MinimalUser + require.NoError(t, json.Unmarshal([]byte(textContent.Text), &returnedUser)) + assert.Equal(t, mockUser.GetLogin(), returnedUser.Login) + assert.Equal(t, mockUser.GetHTMLURL(), returnedUser.ProfileURL) +} + func Test_GetMe_IFC_FeatureFlag(t *testing.T) { t.Parallel() diff --git a/pkg/inventory/server_tool.go b/pkg/inventory/server_tool.go index d25458253f..1d8cbcf885 100644 --- a/pkg/inventory/server_tool.go +++ b/pkg/inventory/server_tool.go @@ -1,6 +1,7 @@ package inventory import ( + "bytes" "context" "encoding/json" "fmt" @@ -209,14 +210,18 @@ func NewServerToolWithContextHandler[In any, Out any](tool mcp.Tool, toolset Too // HandlerFunc ignores deps - deps are retrieved from context at call time HandlerFunc: func(_ any) mcp.ToolHandler { return func(ctx context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) { + rawArguments := req.Params.Arguments + if len(rawArguments) == 0 { + rawArguments = json.RawMessage(`{}`) + } + + if bytes.Equal(bytes.TrimSpace(rawArguments), []byte("null")) { + return invalidArgumentsResult(fmt.Errorf("arguments must be a JSON object")), nil + } + var arguments In - if err := json.Unmarshal(req.Params.Arguments, &arguments); err != nil { - return &mcp.CallToolResult{ - Content: []mcp.Content{ - &mcp.TextContent{Text: fmt.Sprintf("invalid arguments: %s", err)}, - }, - IsError: true, - }, nil + if err := json.Unmarshal(rawArguments, &arguments); err != nil { + return invalidArgumentsResult(err), nil } resp, _, err := handler(ctx, req, arguments) return resp, err @@ -225,6 +230,15 @@ func NewServerToolWithContextHandler[In any, Out any](tool mcp.Tool, toolset Too } } +func invalidArgumentsResult(err error) *mcp.CallToolResult { + return &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: fmt.Sprintf("invalid arguments: %s", err)}, + }, + IsError: true, + } +} + // NewServerTool creates a ServerTool with a raw handler that receives deps via context. // This is the preferred constructor for tools that use mcp.ToolHandler directly because // it doesn't create closures at registration time, which is critical for performance in diff --git a/pkg/inventory/server_tool_test.go b/pkg/inventory/server_tool_test.go index c6d2a6fdd8..9e32b5f30c 100644 --- a/pkg/inventory/server_tool_test.go +++ b/pkg/inventory/server_tool_test.go @@ -12,73 +12,108 @@ import ( "github.com/stretchr/testify/require" ) -func TestNewServerToolWithContextHandler_InvalidArguments_ReturnsIsError(t *testing.T) { - type expectedArgs struct { - Query string `json:"query"` - Limit int `json:"limit"` - } - - tool := NewServerToolWithContextHandler( - mcp.Tool{Name: "test_context_tool"}, - testToolsetMetadata("test"), - func(_ context.Context, _ *mcp.CallToolRequest, _ expectedArgs) (*mcp.CallToolResult, any, error) { - t.Fatal("handler should not be called with invalid arguments") - return nil, nil, nil +func TestNewServerToolWithContextHandler_Arguments(t *testing.T) { + tests := []struct { + name string + arguments json.RawMessage + requireQuery bool + wantHandlerCalled bool + wantIsError bool + wantText string + }{ + { + name: "omitted arguments", + arguments: nil, + wantHandlerCalled: true, + wantText: "success", }, - ) - - handler := tool.HandlerFunc(nil) - - result, err := handler(context.Background(), &mcp.CallToolRequest{ - Params: &mcp.CallToolParamsRaw{ - Name: "test_context_tool", - Arguments: json.RawMessage(`{not valid json`), + { + name: "empty argument bytes", + arguments: json.RawMessage{}, + wantHandlerCalled: true, + wantText: "success", + }, + { + name: "explicit empty object", + arguments: json.RawMessage(`{}`), + wantHandlerCalled: true, + wantText: "success", + }, + { + name: "explicit null", + arguments: json.RawMessage(`null`), + wantIsError: true, + wantText: "arguments must be a JSON object", + }, + { + name: "malformed JSON", + arguments: json.RawMessage(`{not valid json`), + wantIsError: true, + wantText: "invalid arguments", + }, + { + name: "omitted arguments reach required parameter validation", + arguments: nil, + requireQuery: true, + wantHandlerCalled: true, + wantIsError: true, + wantText: "missing required parameter: query", + }, + { + name: "required parameter is decoded", + arguments: json.RawMessage(`{"query":"is:open"}`), + requireQuery: true, + wantHandlerCalled: true, + wantText: "success: is:open", }, - }) - - require.NoError(t, err) - require.NotNil(t, result) - assert.True(t, result.IsError) - assert.Len(t, result.Content, 1) - textContent, ok := result.Content[0].(*mcp.TextContent) - require.True(t, ok) - assert.Contains(t, textContent.Text, "invalid arguments") -} - -func TestNewServerToolWithContextHandler_ValidArguments_Succeeds(t *testing.T) { - type expectedArgs struct { - Owner string `json:"owner"` - Repo string `json:"repo"` } - tool := NewServerToolWithContextHandler( - mcp.Tool{Name: "test_tool"}, - testToolsetMetadata("test"), - func(_ context.Context, _ *mcp.CallToolRequest, args expectedArgs) (*mcp.CallToolResult, any, error) { - return &mcp.CallToolResult{ - Content: []mcp.Content{ - &mcp.TextContent{Text: "success: " + args.Owner + "/" + args.Repo}, + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + handlerCalled := false + tool := NewServerToolWithContextHandler( + mcp.Tool{Name: "test_context_tool"}, + testToolsetMetadata("test"), + func(_ context.Context, _ *mcp.CallToolRequest, args map[string]any) (*mcp.CallToolResult, any, error) { + handlerCalled = true + query, _ := args["query"].(string) + if tc.requireQuery && query == "" { + return &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: "missing required parameter: query"}, + }, + IsError: true, + }, nil, nil + } + text := "success" + if query != "" { + text += ": " + query + } + return &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: text}, + }, + }, nil, nil }, - }, nil, nil - }, - ) - - handler := tool.HandlerFunc(nil) + ) - goodArgs, _ := json.Marshal(map[string]any{"owner": "octocat", "repo": "hello-world"}) - result, err := handler(context.Background(), &mcp.CallToolRequest{ - Params: &mcp.CallToolParamsRaw{ - Name: "test_tool", - Arguments: goodArgs, - }, - }) - - require.NoError(t, err) - require.NotNil(t, result) - assert.False(t, result.IsError) - textContent, ok := result.Content[0].(*mcp.TextContent) - require.True(t, ok) - assert.Equal(t, "success: octocat/hello-world", textContent.Text) + result, err := tool.HandlerFunc(nil)(context.Background(), &mcp.CallToolRequest{ + Params: &mcp.CallToolParamsRaw{ + Name: "test_context_tool", + Arguments: tc.arguments, + }, + }) + + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, tc.wantHandlerCalled, handlerCalled) + assert.Equal(t, tc.wantIsError, result.IsError) + require.Len(t, result.Content, 1) + textContent, ok := result.Content[0].(*mcp.TextContent) + require.True(t, ok) + assert.Contains(t, textContent.Text, tc.wantText) + }) + } } func TestServerToolRegisterFuncAppliesMiddleware(t *testing.T) {