Skip to content
Draft
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
40 changes: 38 additions & 2 deletions pkg/context/mcp_info.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package context
import (
"context"
"encoding/json"
"sync"
)

type mcpMethodInfoCtx string
Expand All @@ -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.
Expand Down
117 changes: 117 additions & 0 deletions pkg/context/mcp_info_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
19 changes: 19 additions & 0 deletions pkg/http/middleware/mcp_parse.go
Original file line number Diff line number Diff line change
Expand Up @@ -106,10 +106,29 @@ 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":
methodInfo.ItemName = mcpReq.Params.URI
}
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)
}
}
82 changes: 82 additions & 0 deletions pkg/http/middleware/mcp_parse_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}{
Expand Down Expand Up @@ -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"},
},
{
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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")
Expand Down Expand Up @@ -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) {
Expand Down
Loading