diff --git a/pkg/context/mcp_info.go b/pkg/context/mcp_info.go index af474b13a0..bfdfd169ae 100644 --- a/pkg/context/mcp_info.go +++ b/pkg/context/mcp_info.go @@ -3,6 +3,7 @@ package context import ( "context" "encoding/json" + "sync" ) type mcpMethodInfoCtx string @@ -20,23 +21,58 @@ type MCPMethodInfo struct { // ItemName is the name of the specific item being accessed (tool name, resource URI, prompt name) // Only populated for call/get methods (tools/call, prompts/get, resources/read) ItemName string + // Owner is the repository owner parsed from tools/call arguments, when available. + // + // Prefer RawArguments and DecodeArguments in new code. + Owner string + // Repo is the repository name parsed from tools/call arguments, when available. + // + // Prefer RawArguments and DecodeArguments in new code. + Repo string + // Arguments contains the decoded tools/call arguments after DecodeArguments is called. + // + // Prefer RawArguments and DecodeArguments in new code. + Arguments map[string]any // RawArguments contains the unmaterialized tool arguments for tools/call requests. RawArguments json.RawMessage + + decodeMu sync.Mutex + decodeDone bool + decodeErr error + decodedArguments map[string]any } // DecodeArguments materializes tool arguments when request middleware needs // call-specific values. Invalid argument shapes are returned to the caller so // the request can continue to the tool handler's normal validation path. func (info *MCPMethodInfo) DecodeArguments() (map[string]any, error) { + info.decodeMu.Lock() + defer info.decodeMu.Unlock() + + if info.decodeDone { + return info.decodedArguments, info.decodeErr + } + if info.Arguments != nil { + info.decodedArguments = info.Arguments + info.decodeDone = true + return info.decodedArguments, nil + } if len(info.RawArguments) == 0 { - return nil, nil + info.decodeDone = true + return info.decodedArguments, info.decodeErr } var arguments map[string]any if err := json.Unmarshal(info.RawArguments, &arguments); err != nil { + info.decodeErr = err + info.decodeDone = true return nil, err } - return arguments, nil + info.decodedArguments = arguments + info.Arguments = arguments + info.decodeDone = true + + return info.decodedArguments, info.decodeErr } // WithMCPMethodInfo stores the MCPMethodInfo in the context. diff --git a/pkg/context/mcp_info_test.go b/pkg/context/mcp_info_test.go new file mode 100644 index 0000000000..d471d8dcb4 --- /dev/null +++ b/pkg/context/mcp_info_test.go @@ -0,0 +1,117 @@ +package context + +import ( + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMCPMethodInfoDecodeArguments(t *testing.T) { + t.Run("caches decoded arguments for compatibility", func(t *testing.T) { + info := &MCPMethodInfo{ + RawArguments: []byte(`{"owner":"github","repo":"github-mcp-server","path":"README.md"}`), + } + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + require.NotNil(t, decoded) + assert.Equal(t, map[string]any{ + "owner": "github", + "repo": "github-mcp-server", + "path": "README.md", + }, decoded) + + cached, err := info.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, decoded, cached) + assert.Equal(t, decoded, info.Arguments) + }) + + t.Run("returns predecoded arguments when present", func(t *testing.T) { + arguments := map[string]any{"owner": "github"} + info := &MCPMethodInfo{Arguments: arguments} + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, arguments, decoded) + }) + + t.Run("null arguments decode as nil", func(t *testing.T) { + info := &MCPMethodInfo{RawArguments: []byte(`null`)} + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + assert.Nil(t, decoded) + assert.Nil(t, info.Arguments) + + decoded, err = info.DecodeArguments() + require.NoError(t, err) + assert.Nil(t, decoded) + }) + + t.Run("non object arguments return an error", func(t *testing.T) { + info := &MCPMethodInfo{RawArguments: []byte(`"not an object"`)} + + decoded, err := info.DecodeArguments() + require.Error(t, err) + assert.Nil(t, decoded) + assert.Nil(t, info.Arguments) + }) + + t.Run("duplicate keys keep the last value", func(t *testing.T) { + info := &MCPMethodInfo{RawArguments: []byte(`{"owner":"first","owner":"second"}`)} + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, map[string]any{"owner": "second"}, decoded) + }) + + t.Run("key casing is preserved", func(t *testing.T) { + info := &MCPMethodInfo{RawArguments: []byte(`{"owner":"lower","Owner":"upper"}`)} + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, map[string]any{"owner": "lower", "Owner": "upper"}, decoded) + }) + + t.Run("cached decode survives compatibility field mutation", func(t *testing.T) { + info := &MCPMethodInfo{RawArguments: []byte(`{"owner":"github"}`)} + + decoded, err := info.DecodeArguments() + require.NoError(t, err) + info.Arguments = nil + + cached, err := info.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, decoded, cached) + }) + + t.Run("concurrent decode is safe", func(t *testing.T) { + info := &MCPMethodInfo{ + RawArguments: []byte(`{"owner":"github","repo":"github-mcp-server","nested":{"path":"README.md"}}`), + } + + var wg sync.WaitGroup + errors := make(chan error, 32) + for range 32 { + wg.Go(func() { + decoded, err := info.DecodeArguments() + if err != nil { + errors <- err + return + } + if decoded["owner"] != "github" { + errors <- assert.AnError + } + }) + } + wg.Wait() + close(errors) + + for err := range errors { + require.NoError(t, err) + } + }) +} diff --git a/pkg/http/middleware/mcp_parse.go b/pkg/http/middleware/mcp_parse.go index 0b56902261..90605ffee9 100644 --- a/pkg/http/middleware/mcp_parse.go +++ b/pkg/http/middleware/mcp_parse.go @@ -106,6 +106,7 @@ func parseMCPMethodInfo(body []byte) (*ghcontext.MCPMethodInfo, error) { case "tools/call": methodInfo.ItemName = mcpReq.Params.Name methodInfo.RawArguments = mcpReq.Params.Arguments + populateToolArgumentCompatibilityFields(methodInfo) case "prompts/get": methodInfo.ItemName = mcpReq.Params.Name case "resources/read": @@ -113,3 +114,21 @@ func parseMCPMethodInfo(body []byte) (*ghcontext.MCPMethodInfo, error) { } return methodInfo, nil } + +func populateToolArgumentCompatibilityFields(methodInfo *ghcontext.MCPMethodInfo) { + if len(methodInfo.RawArguments) == 0 { + return + } + + var arguments map[string]json.RawMessage + if err := json.Unmarshal(methodInfo.RawArguments, &arguments); err != nil { + return + } + + if owner, ok := arguments["owner"]; ok { + _ = json.Unmarshal(owner, &methodInfo.Owner) + } + if repo, ok := arguments["repo"]; ok { + _ = json.Unmarshal(repo, &methodInfo.Repo) + } +} diff --git a/pkg/http/middleware/mcp_parse_test.go b/pkg/http/middleware/mcp_parse_test.go index e067f7808a..6f123d7b35 100644 --- a/pkg/http/middleware/mcp_parse_test.go +++ b/pkg/http/middleware/mcp_parse_test.go @@ -23,6 +23,8 @@ func TestWithMCPParse(t *testing.T) { expectedMethod string expectedItem string expectedRaw string + expectedOwner string + expectedRepo string expectedArgs map[string]any expectArgsError bool }{ @@ -94,6 +96,8 @@ func TestWithMCPParse(t *testing.T) { expectedMethod: "tools/call", expectedItem: "get_file_contents", expectedRaw: `{"owner":"github","repo":"github-mcp-server","path":"README.md"}`, + expectedOwner: "github", + expectedRepo: "github-mcp-server", expectedArgs: map[string]any{"owner": "github", "repo": "github-mcp-server", "path": "README.md"}, }, { @@ -107,6 +111,40 @@ func TestWithMCPParse(t *testing.T) { expectedRaw: `"not an object"`, expectArgsError: true, }, + { + name: "tools/call keeps valid owner when repo has invalid type", + method: http.MethodPost, + path: "/mcp", + body: `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"get_file_contents","arguments":{"owner":"github","repo":123,"path":"README.md"}}}`, + expectInfo: true, + expectedMethod: "tools/call", + expectedItem: "get_file_contents", + expectedRaw: `{"owner":"github","repo":123,"path":"README.md"}`, + expectedOwner: "github", + expectedArgs: map[string]any{"owner": "github", "repo": 123.0, "path": "README.md"}, + }, + { + name: "tools/call ignores non exact owner key casing", + method: http.MethodPost, + path: "/mcp", + body: `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"get_file_contents","arguments":{"Owner":"github","repo":"github-mcp-server"}}}`, + expectInfo: true, + expectedMethod: "tools/call", + expectedItem: "get_file_contents", + expectedRaw: `{"Owner":"github","repo":"github-mcp-server"}`, + expectedRepo: "github-mcp-server", + expectedArgs: map[string]any{"Owner": "github", "repo": "github-mcp-server"}, + }, + { + name: "tools/call with null arguments keeps empty compatibility fields", + method: http.MethodPost, + path: "/mcp", + body: `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"get_file_contents","arguments":null}}`, + expectInfo: true, + expectedMethod: "tools/call", + expectedItem: "get_file_contents", + expectedRaw: `null`, + }, { name: "prompts/get parses name", method: http.MethodPost, @@ -161,6 +199,8 @@ func TestWithMCPParse(t *testing.T) { if tt.expectedRaw != "" { assert.JSONEq(t, tt.expectedRaw, string(capturedInfo.RawArguments)) } + assert.Equal(t, tt.expectedOwner, capturedInfo.Owner) + assert.Equal(t, tt.expectedRepo, capturedInfo.Repo) decodedArgs, err := capturedInfo.DecodeArguments() if tt.expectArgsError { assert.Error(t, err) @@ -169,6 +209,7 @@ func TestWithMCPParse(t *testing.T) { } if tt.expectedArgs != nil { assert.Equal(t, tt.expectedArgs, decodedArgs) + assert.Equal(t, decodedArgs, capturedInfo.Arguments) } } else { assert.False(t, infoCaptured, "MCPMethodInfo should not be present in context") @@ -197,6 +238,47 @@ func TestWithMCPParseRetainsLargeArgumentsWithoutMaterializingThem(t *testing.T) require.NotNil(t, capturedInfo) assert.Equal(t, json.RawMessage(rawArguments), capturedInfo.RawArguments) + assert.Empty(t, capturedInfo.Owner) + assert.Empty(t, capturedInfo.Repo) +} + +func TestWithMCPParseCompatibilityFieldsUseExactKeysAndLastDuplicates(t *testing.T) { + body := `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"get_file_contents","arguments":{"owner":"first","owner":"second","repo":"one","repo":"two","nested":{"keep":["all",{"payload":true}]}}}}` + + var capturedInfo *ghcontext.MCPMethodInfo + next := http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + capturedInfo, _ = ghcontext.MCPMethod(r.Context()) + }) + + request := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body)) + WithMCPParse()(next).ServeHTTP(httptest.NewRecorder(), request) + + require.NotNil(t, capturedInfo) + assert.Equal(t, `{"owner":"first","owner":"second","repo":"one","repo":"two","nested":{"keep":["all",{"payload":true}]}}`, string(capturedInfo.RawArguments)) + assert.Equal(t, "second", capturedInfo.Owner) + assert.Equal(t, "two", capturedInfo.Repo) + + decodedArgs, err := capturedInfo.DecodeArguments() + require.NoError(t, err) + assert.Equal(t, map[string]any{ + "owner": "second", + "repo": "two", + "nested": map[string]any{ + "keep": []any{"all", map[string]any{"payload": true}}, + }, + }, decodedArgs) +} + +func BenchmarkParseMCPMethodInfoToolsCall(b *testing.B) { + body := []byte(`{"jsonrpc":"2.0","method":"tools/call","params":{"name":"get_file_contents","arguments":{"owner":"github","repo":"github-mcp-server","nested":{"items":[{"payload":"` + strings.Repeat("x", 4096) + `"}]}}}}`) + + b.ReportAllocs() + for b.Loop() { + info, err := parseMCPMethodInfo(body) + if err != nil || info == nil { + b.Fatalf("parseMCPMethodInfo() error = %v, info = %v", err, info) + } + } } func TestWithMCPParse_BodyRestoration(t *testing.T) {