Skip to content

Commit e493a9b

Browse files
fix(context): reject explicit null team user arguments
Preserve the legacy optional user type error before typed decoding, without changing omitted or empty user fallback. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent f1c61b3 commit e493a9b

2 files changed

Lines changed: 73 additions & 0 deletions

File tree

‎pkg/github/context_tools.go‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
package github
22

33
import (
4+
"bytes"
45
"context"
6+
"encoding/json"
57
"fmt"
68
"time"
79

@@ -51,6 +53,17 @@ func contextToolValidationSchema(schema *jsonschema.Schema) *jsonschema.Schema {
5153
return validationSchema
5254
}
5355

56+
func normalizeGetTeamsInput(arguments json.RawMessage) (json.RawMessage, error) {
57+
var fields map[string]json.RawMessage
58+
if err := json.Unmarshal(arguments, &fields); err != nil {
59+
return nil, err
60+
}
61+
if user, exists := fields["user"]; exists && bytes.Equal(bytes.TrimSpace(user), []byte("null")) {
62+
return nil, &inventory.ToolInputError{Message: "parameter user is not of type string, is <nil>"}
63+
}
64+
return arguments, nil
65+
}
66+
5467
// UserDetails contains additional fields about a GitHub user not already
5568
// present in MinimalUser. Used by get_me context tool but omitted from search_users.
5669
type UserDetails struct {
@@ -253,6 +266,7 @@ func GetTeams(t translations.TranslationHelperFunc) inventory.ServerTool {
253266
}
254267
return result, organizations, nil
255268
},
269+
normalizeGetTeamsInput,
256270
)
257271
}
258272

‎pkg/github/context_tools_test.go‎

Lines changed: 59 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -677,6 +677,17 @@ func Test_GetTeams(t *testing.T) {
677677
expectToolError: false,
678678
expectedTeamsCount: 2,
679679
},
680+
{
681+
name: "empty user uses authenticated user",
682+
makeDeps: func() ToolDependencies {
683+
return BaseDeps{
684+
Client: mustNewGHClient(t, httpClientWithUser()),
685+
GQLClient: gqlClientForTestuser(),
686+
}
687+
},
688+
requestArgs: map[string]any{"user": ""},
689+
expectedTeamsCount: 2,
690+
},
680691
{
681692
name: "no teams found",
682693
makeDeps: func() ToolDependencies {
@@ -904,6 +915,54 @@ func Test_GetTeamMembers(t *testing.T) {
904915
}
905916
}
906917

918+
func Test_GetTeams_NullUserForDirectHandler(t *testing.T) {
919+
t.Parallel()
920+
921+
clientCalls := 0
922+
gqlClientCalls := 0
923+
deps := stubDeps{
924+
clientFn: func(context.Context) (*github.Client, error) {
925+
clientCalls++
926+
return nil, nil
927+
},
928+
gqlClientFn: func(context.Context) (*githubv4.Client, error) {
929+
gqlClientCalls++
930+
return nil, nil
931+
},
932+
obsv: stubExporters(),
933+
}
934+
serverTool := GetTeams(translations.NullTranslationHelper)
935+
request := createMCPRequest(map[string]any{"user": nil})
936+
result, err := serverTool.Handler(deps)(ContextWithDeps(context.Background(), deps), &request)
937+
938+
require.NoError(t, err)
939+
require.True(t, result.IsError)
940+
assert.Equal(t, "parameter user is not of type string, is <nil>", getErrorResult(t, result).Text)
941+
assert.Nil(t, result.StructuredContent)
942+
assert.Zero(t, clientCalls)
943+
assert.Zero(t, gqlClientCalls)
944+
}
945+
946+
func TestNormalizeGetTeamsInput(t *testing.T) {
947+
t.Parallel()
948+
949+
for _, arguments := range []string{`{}`, `{"user":""}`, `{"user":"octocat"}`, `{"legacy_ignored_argument":true}`} {
950+
t.Run(arguments, func(t *testing.T) {
951+
normalized, err := normalizeGetTeamsInput(json.RawMessage(arguments))
952+
require.NoError(t, err)
953+
assert.Equal(t, arguments, string(normalized))
954+
})
955+
}
956+
for _, arguments := range []string{`{"user":null}`, `{"user": null }`} {
957+
t.Run(arguments, func(t *testing.T) {
958+
_, err := normalizeGetTeamsInput(json.RawMessage(arguments))
959+
var inputError *inventory.ToolInputError
960+
require.ErrorAs(t, err, &inputError)
961+
assert.Equal(t, "parameter user is not of type string, is <nil>", inputError.Message)
962+
})
963+
}
964+
}
965+
907966
func Test_GetTeamMembers_RequiredIdentifiersForDirectHandler(t *testing.T) {
908967
t.Parallel()
909968

0 commit comments

Comments
 (0)