Skip to content

Commit f027e5d

Browse files
fix(context): preserve exact routing argument keys
Ignore case-variant routing aliases before typed decoding while retaining canonical fields and unrelated unknown arguments. Cover direct handlers and legacy/modern registered calls, authenticated fallback, canonical null rejection, and missing identifiers. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent bae4e60 commit f027e5d

2 files changed

Lines changed: 184 additions & 1 deletion

File tree

‎pkg/github/context_tools.go‎

Lines changed: 28 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ import (
55
"context"
66
"encoding/json"
77
"fmt"
8+
"strings"
89
"time"
910

1011
ghErrors "github.com/github/github-mcp-server/pkg/errors"
@@ -61,7 +62,32 @@ func normalizeGetTeamsInput(arguments json.RawMessage) (json.RawMessage, error)
6162
if user, exists := fields["user"]; exists && bytes.Equal(bytes.TrimSpace(user), []byte("null")) {
6263
return nil, &inventory.ToolInputError{Message: "parameter user is not of type string, is <nil>"}
6364
}
64-
return arguments, nil
65+
return normalizeContextRoutingKeys(arguments, fields, "user")
66+
}
67+
68+
func normalizeGetTeamMembersInput(arguments json.RawMessage) (json.RawMessage, error) {
69+
var fields map[string]json.RawMessage
70+
if err := json.Unmarshal(arguments, &fields); err != nil {
71+
return nil, err
72+
}
73+
return normalizeContextRoutingKeys(arguments, fields, "org", "team_slug")
74+
}
75+
76+
func normalizeContextRoutingKeys(arguments json.RawMessage, fields map[string]json.RawMessage, routingKeys ...string) (json.RawMessage, error) {
77+
changed := false
78+
for name := range fields {
79+
for _, key := range routingKeys {
80+
if name != key && strings.EqualFold(name, key) {
81+
delete(fields, name)
82+
changed = true
83+
break
84+
}
85+
}
86+
}
87+
if !changed {
88+
return arguments, nil
89+
}
90+
return json.Marshal(fields)
6591
}
6692

6793
// UserDetails contains additional fields about a GitHub user not already
@@ -338,5 +364,6 @@ func GetTeamMembers(t translations.TranslationHelperFunc) inventory.ServerTool {
338364
}
339365
return result, members, nil
340366
},
367+
normalizeGetTeamMembersInput,
341368
)
342369
}

‎pkg/github/context_tools_test.go‎

Lines changed: 156 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -963,6 +963,162 @@ func TestNormalizeGetTeamsInput(t *testing.T) {
963963
}
964964
}
965965

966+
func TestContextToolsExactRoutingKeys(t *testing.T) {
967+
for _, mode := range []string{"direct", "legacy", "modern", "unknown"} {
968+
t.Run(mode, func(t *testing.T) {
969+
var clientCalls, gqlCalls int
970+
deps := stubDeps{
971+
clientFn: func(ctx context.Context) (*github.Client, error) {
972+
clientCalls++
973+
return stubClientFnFromHTTP(t, MockHTTPClientWithHandlers(map[string]http.HandlerFunc{
974+
GetUser: mockResponse(t, http.StatusOK, &github.User{Login: new("authenticated")}),
975+
}))(ctx)
976+
},
977+
gqlClientFn: func(context.Context) (*githubv4.Client, error) {
978+
gqlCalls++
979+
return githubv4.NewClient(githubv4mock.NewMockedHTTPClient(
980+
githubv4mock.NewQueryMatcher(
981+
"query($login:String!){user(login: $login){organizations(first: 100){nodes{login,teams(first: 100, userLogins: [$login]){nodes{name,slug,description}}}}}}",
982+
map[string]any{"login": "authenticated"},
983+
githubv4mock.DataResponse(map[string]any{
984+
"user": map[string]any{"organizations": map[string]any{"nodes": []any{}}},
985+
}),
986+
),
987+
githubv4mock.NewQueryMatcher(
988+
"query($org:String!$teamSlug:String!){organization(login: $org){team(slug: $teamSlug){members(first: 100){nodes{login}}}}}",
989+
map[string]any{"org": "testorg", "teamSlug": "testteam"},
990+
githubv4mock.DataResponse(map[string]any{
991+
"organization": map[string]any{"team": map[string]any{"members": map[string]any{"nodes": []any{}}}},
992+
}),
993+
),
994+
)), nil
995+
},
996+
obsv: stubExporters(),
997+
}
998+
tools := []inventory.ServerTool{
999+
GetTeams(translations.NullTranslationHelper),
1000+
GetTeamMembers(translations.NullTranslationHelper),
1001+
}
1002+
call := func(name string, args map[string]any) *mcp.CallToolResult {
1003+
for _, tool := range tools {
1004+
if tool.Tool.Name == name {
1005+
request := createMCPRequest(args)
1006+
if mode == "unknown" {
1007+
request.Params.Meta = mcp.Meta{mcp.MetaKeyProtocolVersion: "2099-01-01"}
1008+
}
1009+
result, err := tool.Handler(deps)(ContextWithDeps(context.Background(), deps), &request)
1010+
require.NoError(t, err)
1011+
return result
1012+
}
1013+
}
1014+
t.Fatalf("unknown tool %s", name)
1015+
return nil
1016+
}
1017+
if mode == "legacy" || mode == "modern" {
1018+
server := mcp.NewServer(&mcp.Implementation{Name: "test-server", Version: "v0.0.1"}, nil)
1019+
server.AddReceivingMiddleware(InjectDepsMiddleware(deps))
1020+
inv, err := inventory.NewBuilder().SetTools(tools).WithToolsets([]string{"all"}).Build()
1021+
require.NoError(t, err)
1022+
inv.RegisterTools(context.Background(), server, deps)
1023+
serverTransport, clientTransport := mcp.NewInMemoryTransports()
1024+
serverSession, err := server.Connect(context.Background(), serverTransport, nil)
1025+
require.NoError(t, err)
1026+
t.Cleanup(func() { _ = serverSession.Close() })
1027+
client := mcp.NewClient(&mcp.Implementation{Name: "test-client", Version: "v0.0.1"}, nil)
1028+
version := "2025-03-26"
1029+
if mode == "modern" {
1030+
version = inventory.ProtocolVersionMultiRoundTrip
1031+
}
1032+
session, err := client.Connect(context.Background(), clientTransport, &mcp.ClientSessionOptions{ProtocolVersion: version})
1033+
require.NoError(t, err)
1034+
t.Cleanup(func() { _ = session.Close() })
1035+
meta := mcp.Meta{}
1036+
if mode == "modern" {
1037+
meta[mcp.MetaKeyProtocolVersion] = inventory.ProtocolVersionMultiRoundTrip
1038+
}
1039+
call = func(name string, args map[string]any) *mcp.CallToolResult {
1040+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{
1041+
Name: name, Arguments: args, Meta: meta,
1042+
})
1043+
require.NoError(t, err)
1044+
return result
1045+
}
1046+
}
1047+
for _, args := range []map[string]any{
1048+
{"USER": "wrong-user", "unrelated": true},
1049+
{"UsEr": nil},
1050+
{"user": "", "USER": 42},
1051+
{"user": "authenticated", "USER": "wrong-user"},
1052+
} {
1053+
beforeREST := clientCalls
1054+
beforeGQL := gqlCalls
1055+
result := call("get_teams", args)
1056+
require.False(t, result.IsError)
1057+
assert.Equal(t, beforeGQL+1, gqlCalls)
1058+
if args["user"] == "authenticated" {
1059+
assert.Equal(t, beforeREST, clientCalls)
1060+
} else {
1061+
assert.Equal(t, beforeREST+1, clientCalls)
1062+
}
1063+
if mode == "modern" {
1064+
assert.Equal(t, "[]", getTextResult(t, result).Text)
1065+
assert.NotNil(t, result.StructuredContent)
1066+
} else {
1067+
assert.Equal(t, "null", getTextResult(t, result).Text)
1068+
assert.Nil(t, result.StructuredContent)
1069+
}
1070+
}
1071+
beforeREST, beforeGQL := clientCalls, gqlCalls
1072+
result := call("get_teams", map[string]any{"user": nil, "USER": "authenticated"})
1073+
require.True(t, result.IsError)
1074+
assert.Equal(t, "parameter user is not of type string, is <nil>", getErrorResult(t, result).Text)
1075+
assert.Equal(t, beforeREST, clientCalls)
1076+
assert.Equal(t, beforeGQL, gqlCalls)
1077+
for _, args := range []map[string]any{
1078+
{"ORG": "testorg", "team_slug": "testteam"},
1079+
{"org": "testorg", "TEAM_SLUG": "testteam"},
1080+
{"ORG": "testorg", "TEAM_SLUG": "testteam"},
1081+
} {
1082+
result := call("get_team_members", args)
1083+
require.True(t, result.IsError)
1084+
assert.Nil(t, result.StructuredContent)
1085+
assert.Equal(t, beforeGQL, gqlCalls)
1086+
}
1087+
result = call("get_team_members", map[string]any{
1088+
"org": "testorg", "team_slug": "testteam", "ORG": 42, "TEAM_SLUG": nil, "unrelated": true,
1089+
})
1090+
require.False(t, result.IsError)
1091+
assert.Equal(t, beforeGQL+1, gqlCalls)
1092+
if mode == "modern" {
1093+
assert.Equal(t, "[]", getTextResult(t, result).Text)
1094+
} else {
1095+
assert.Equal(t, "null", getTextResult(t, result).Text)
1096+
}
1097+
})
1098+
}
1099+
}
1100+
1101+
func TestNormalizeContextRoutingKeys(t *testing.T) {
1102+
t.Parallel()
1103+
1104+
for _, tc := range []struct {
1105+
name string
1106+
normalize inventory.InputNormalizer
1107+
args string
1108+
want string
1109+
}{
1110+
{"teams", normalizeGetTeamsInput, `{"user":"exact","USER":42,"UsEr":null,"unknown":{"USER":"retained"}}`, `{"user":"exact","unknown":{"USER":"retained"}}`},
1111+
{"members", normalizeGetTeamMembersInput, `{"org":"exact","ORG":false,"team_slug":"slug","Team_Slug":null,"unknown":true}`, `{"org":"exact","team_slug":"slug","unknown":true}`},
1112+
{"canonical null retained", normalizeGetTeamMembersInput, `{"org":null,"team_slug":null,"ORG":"ignored"}`, `{"org":null,"team_slug":null}`},
1113+
} {
1114+
t.Run(tc.name, func(t *testing.T) {
1115+
result, err := tc.normalize(json.RawMessage(tc.args))
1116+
require.NoError(t, err)
1117+
assert.JSONEq(t, tc.want, string(result))
1118+
})
1119+
}
1120+
}
1121+
9661122
func Test_GetTeamMembers_RequiredIdentifiersForDirectHandler(t *testing.T) {
9671123
t.Parallel()
9681124

0 commit comments

Comments
 (0)