Skip to content
Merged
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
30 changes: 30 additions & 0 deletions pkg/github/context_tools_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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()

Expand Down
28 changes: 21 additions & 7 deletions pkg/inventory/server_tool.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package inventory

import (
"bytes"
"context"
"encoding/json"
"fmt"
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
157 changes: 96 additions & 61 deletions pkg/inventory/server_tool_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
Loading