Skip to content
Closed
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
42 changes: 42 additions & 0 deletions pkg/github/context_tools_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,48 @@ func Test_GetMe(t *testing.T) {
}
}

func Test_GetMe_OmittedArguments(t *testing.T) {
t.Parallel()

serverTool := GetMe(translations.NullTranslationHelper)

mockUser := &github.User{
Login: github.Ptr("testuser"),
Name: github.Ptr("Test User"),
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()}
handler := serverTool.Handler(deps)

tests := []struct {
name string
arguments json.RawMessage
}{
{name: "nil arguments", arguments: nil},
{name: "empty byte slice", arguments: json.RawMessage{}},
{name: "explicit empty object", arguments: json.RawMessage(`{}`)},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
request := createMCPRequestWithRawArguments(tc.arguments)
result, err := handler(ContextWithDeps(context.Background(), deps), &request)
require.NoError(t, err)
require.False(t, result.IsError)

textContent := getTextResult(t, result)
var returnedUser MinimalUser
err = json.Unmarshal([]byte(textContent.Text), &returnedUser)
require.NoError(t, err)
assert.Equal(t, "testuser", returnedUser.Login)
})
}
}

func Test_GetMe_IFC_FeatureFlag(t *testing.T) {
t.Parallel()

Expand Down
10 changes: 10 additions & 0 deletions pkg/github/helper_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -334,6 +334,16 @@ func createMCPRequest(args any) mcp.CallToolRequest {
}
}

// createMCPRequestWithRawArguments creates a CallToolRequest with the given raw JSON arguments.
// Use nil or an empty slice to simulate clients that omit the arguments field.
func createMCPRequestWithRawArguments(args json.RawMessage) mcp.CallToolRequest {
return mcp.CallToolRequest{
Params: &mcp.CallToolParamsRaw{
Arguments: args,
},
}
}

// Well-known MCP client names used in tests.
const (
ClientNameVSCodeInsiders = "Visual Studio Code - Insiders"
Expand Down
26 changes: 19 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 @@ -210,13 +211,15 @@ func NewServerToolWithContextHandler[In any, Out any](tool mcp.Tool, toolset Too
HandlerFunc: func(_ any) mcp.ToolHandler {
return func(ctx context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) {
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
args := req.Params.Arguments
if len(args) == 0 {
args = json.RawMessage(`{}`)
}
if bytes.Equal(bytes.TrimSpace(args), []byte("null")) {
return invalidArgumentsResult("arguments must be a JSON object"), nil
}
if err := json.Unmarshal(args, &arguments); err != nil {
return invalidArgumentsResult(err.Error()), nil
}
resp, _, err := handler(ctx, req, arguments)
return resp, err
Expand All @@ -225,6 +228,15 @@ func NewServerToolWithContextHandler[In any, Out any](tool mcp.Tool, toolset Too
}
}

func invalidArgumentsResult(message string) *mcp.CallToolResult {
return &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: fmt.Sprintf("invalid arguments: %s", message)},
},
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
107 changes: 107 additions & 0 deletions pkg/inventory/server_tool_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,113 @@ func TestNewServerToolWithContextHandler_ValidArguments_Succeeds(t *testing.T) {
assert.Equal(t, "success: octocat/hello-world", textContent.Text)
}

func TestNewServerToolWithContextHandler_ArgumentObjects(t *testing.T) {
type emptyArgs struct{}

handlerCalled := false
tool := NewServerToolWithContextHandler(
mcp.Tool{Name: "zero_arg_tool"},
testToolsetMetadata("test"),
func(_ context.Context, _ *mcp.CallToolRequest, _ emptyArgs) (*mcp.CallToolResult, any, error) {
handlerCalled = true
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "ok"}},
}, nil, nil
},
)
handler := tool.HandlerFunc(nil)

tests := []struct {
name string
arguments json.RawMessage
wantError bool
wantMessage string
wantHandler bool
}{
{
name: "nil arguments",
arguments: nil,
wantHandler: true,
},
{
name: "zero-length arguments",
arguments: json.RawMessage{},
wantHandler: true,
},
{
name: "empty object",
arguments: json.RawMessage(`{}`),
wantHandler: true,
},
{
name: "null",
arguments: json.RawMessage(`null`),
wantError: true,
wantMessage: "arguments must be a JSON object",
},
{
name: "malformed JSON",
arguments: json.RawMessage(`{not valid json`),
wantError: true,
wantMessage: "invalid character",
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
handlerCalled = false
result, err := handler(context.Background(), &mcp.CallToolRequest{
Params: &mcp.CallToolParamsRaw{
Name: "zero_arg_tool",
Arguments: tt.arguments,
},
})

require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, tt.wantError, result.IsError)
assert.Equal(t, tt.wantHandler, handlerCalled)
if tt.wantMessage != "" {
textContent, ok := result.Content[0].(*mcp.TextContent)
require.True(t, ok)
assert.Contains(t, textContent.Text, tt.wantMessage)
}
})
}
}

func TestNewServerToolWithContextHandler_OmittedArgumentsReachRequiredParameterValidation(t *testing.T) {
tool := NewServerToolWithContextHandler(
mcp.Tool{Name: "parameterized_tool"},
testToolsetMetadata("test"),
func(_ context.Context, _ *mcp.CallToolRequest, args map[string]any) (*mcp.CallToolResult, any, error) {
if _, ok := args["owner"]; !ok {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "missing required parameter: owner"}},
IsError: true,
}, nil, nil
}
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "ok"}},
}, nil, nil
},
)

result, err := tool.HandlerFunc(nil)(context.Background(), &mcp.CallToolRequest{
Params: &mcp.CallToolParamsRaw{
Name: "parameterized_tool",
Arguments: nil,
},
})

require.NoError(t, err)
require.NotNil(t, result)
require.True(t, result.IsError)
textContent, ok := result.Content[0].(*mcp.TextContent)
require.True(t, ok)
assert.Contains(t, textContent.Text, "missing required parameter: owner")
}

func TestServerToolRegisterFuncAppliesMiddleware(t *testing.T) {
tool := NewServerTool(
mcp.Tool{
Expand Down
Loading