Skip to content

Commit 1c30ec5

Browse files
fix(inventory): finalize OAuth short-circuit tool results
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 70c15f2 commit 1c30ec5

7 files changed

Lines changed: 514 additions & 16 deletions

File tree

‎docs/typed-tool-schemas.md‎

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
# Typed tool schemas
22

33
Typed tool registrations use concrete Go input and output types with the MCP
4-
Go SDK's `mcp.AddTool` path. The SDK infers schemas when the tool definition
5-
does not provide them, validates arguments before the handler runs, and
6-
validates typed output. Keep business rules that JSON Schema cannot express
4+
Go SDK's `mcp.AddTool` path. Inventory infers missing schemas and validates
5+
arguments against the cached input schema before the handler runs; the SDK
6+
adapts and validates typed output. Keep business rules that JSON Schema cannot express
77
in the handler or a preflight callback.
88

99
Input inference unwraps one pointer level, matching the SDK's object input
@@ -31,13 +31,13 @@ schemas; do not mutate the returned pointers.
3131

3232
When compatibility requires a broader runtime input contract than the one
3333
advertised to clients, provide `ValidationInputSchema`. The tool's declared
34-
`InputSchema` remains visible while the SDK validates calls against the
34+
`InputSchema` remains visible while inventory validates calls against the
3535
runtime-only schema. Use `Preflight` for checks that need raw arguments or
3636
request dependencies before typed decoding; it may return a derived context
3737
for the handler. Input normalizers are only for compatibility transformations,
3838
not a replacement for schema validation.
3939

40-
The SDK applies defaults from the runtime input schema before decoding. If an
40+
Inventory applies defaults from the runtime input schema before decoding. If an
4141
omitted field must remain omitted, build the runtime schema with
4242
`inventory.CloneSchemaWithoutDefaults` before applying validation-only
4343
changes. Use `inventory.CloneSchema` when deriving other runtime-only schema
@@ -46,6 +46,14 @@ subschemas. Default removal visits each child once per parent. The advertised
4646
schema can retain its defaults; the constructor caches the runtime schema
4747
without mutating either caller-owned schema.
4848

49+
Availability and authorization guards run before preflight, normalization,
50+
and input validation. A guard or preflight result is passed through the
51+
registered SDK handler without decoding the original arguments or invoking
52+
user code, so the SDK still finalizes multi-round-trip results with
53+
`resultType: "input_required"` (or `"complete"`). The permissive SDK input
54+
envelope used for this handoff is never advertised and does not replace the
55+
original validation schema for real calls.
56+
4957
The output schema and `structuredContent` are exposed only for a negotiated,
5058
SDK-supported protocol version `2026-07-28` or later. Unknown versions are
5159
treated as legacy, including unsupported future dates; older, absent, and

‎go.mod‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ require (
1212
github.com/microcosm-cc/bluemonday v1.0.27
1313
github.com/modelcontextprotocol/go-sdk v1.8.0
1414
github.com/muesli/cache2go v0.0.0-20221011235721-518229cd8021
15+
github.com/segmentio/encoding v0.5.4
1516
github.com/shurcooL/githubv4 v0.0.0-20260209031235-2402fdf4a9ed
1617
github.com/shurcooL/graphql v0.0.0-20240915155400-7ee5256398cf
1718
github.com/spf13/cobra v1.10.2
@@ -32,7 +33,6 @@ require (
3233
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
3334
github.com/sagikazarmark/locafero v0.11.0 // indirect
3435
github.com/segmentio/asm v1.1.3 // indirect
35-
github.com/segmentio/encoding v0.5.4 // indirect
3636
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
3737
github.com/spf13/afero v1.15.0 // indirect
3838
github.com/spf13/cast v1.10.0 // indirect

‎internal/ghmcp/oauth_typed_test.go‎

Lines changed: 308 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,308 @@
1+
package ghmcp
2+
3+
import (
4+
"context"
5+
"encoding/json"
6+
"testing"
7+
8+
"github.com/github/github-mcp-server/internal/oauth"
9+
"github.com/github/github-mcp-server/pkg/inventory"
10+
"github.com/google/jsonschema-go/jsonschema"
11+
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
12+
"github.com/modelcontextprotocol/go-sdk/mcp"
13+
"github.com/stretchr/testify/assert"
14+
"github.com/stretchr/testify/require"
15+
)
16+
17+
type oauthProbeOutput struct {
18+
Status string `json:"status"`
19+
}
20+
21+
func oauthRegisteredSession(
22+
t *testing.T,
23+
typed bool,
24+
fake oauthAuthenticator,
25+
protocol string,
26+
options *mcp.ClientOptions,
27+
toolCalls *int,
28+
captures ...chan json.RawMessage,
29+
) *mcp.ClientSession {
30+
t.Helper()
31+
server := mcp.NewServer(&mcp.Implementation{Name: "oauth-test", Version: "test"}, nil)
32+
middleware := createOAuthToolMiddleware(fake, discardLogger())
33+
if typed {
34+
tool := inventory.NewServerToolWithContextHandler(
35+
mcp.Tool{Name: probeToolName},
36+
inventory.ToolsetMetadata{ID: "test"},
37+
func(context.Context, *mcp.CallToolRequest, struct{}) (*mcp.CallToolResult, oauthProbeOutput, error) {
38+
(*toolCalls)++
39+
return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: "tool-ran"}}},
40+
oauthProbeOutput{Status: "tool-ran"}, nil
41+
},
42+
)
43+
tool.RegisterFunc(server, nil, middleware)
44+
} else {
45+
server.AddTool(&mcp.Tool{Name: probeToolName, InputSchema: &jsonschema.Schema{Type: "object"}},
46+
middleware(func(context.Context, *mcp.CallToolRequest) (*mcp.CallToolResult, error) {
47+
(*toolCalls)++
48+
return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: "tool-ran"}}}, nil
49+
}))
50+
}
51+
return connectOAuthRegisteredServer(t, server, protocol, options, captures...)
52+
}
53+
54+
func connectOAuthRegisteredServer(t *testing.T, server *mcp.Server, protocol string, options *mcp.ClientOptions, captures ...chan json.RawMessage) *mcp.ClientSession {
55+
t.Helper()
56+
st, ct := mcp.NewInMemoryTransports()
57+
ss, err := server.Connect(context.Background(), st, nil)
58+
require.NoError(t, err)
59+
t.Cleanup(func() { _ = ss.Close() })
60+
client := mcp.NewClient(&mcp.Implementation{Name: "oauth-client", Version: "test"}, options)
61+
var transport mcp.Transport = ct
62+
if len(captures) > 0 {
63+
transport = oauthCaptureTransport{Transport: ct, responses: captures[0]}
64+
}
65+
cs, err := client.Connect(context.Background(), transport, &mcp.ClientSessionOptions{ProtocolVersion: protocol})
66+
require.NoError(t, err)
67+
t.Cleanup(func() { _ = cs.Close() })
68+
return cs
69+
}
70+
71+
type oauthCaptureTransport struct {
72+
mcp.Transport
73+
responses chan json.RawMessage
74+
}
75+
76+
func (t oauthCaptureTransport) Connect(ctx context.Context) (mcp.Connection, error) {
77+
connection, err := t.Transport.Connect(ctx)
78+
if err != nil {
79+
return nil, err
80+
}
81+
return oauthCaptureConnection{Connection: connection, responses: t.responses}, nil
82+
}
83+
84+
type oauthCaptureConnection struct {
85+
mcp.Connection
86+
responses chan json.RawMessage
87+
}
88+
89+
func (c oauthCaptureConnection) Read(ctx context.Context) (jsonrpc.Message, error) {
90+
message, err := c.Connection.Read(ctx)
91+
if err == nil {
92+
if response, ok := message.(*jsonrpc.Response); ok && response.Result != nil {
93+
c.responses <- append(json.RawMessage(nil), response.Result...)
94+
}
95+
}
96+
return message, err
97+
}
98+
99+
func oauthPendingAuthenticator() *fakeAuthenticator {
100+
return &fakeAuthenticator{
101+
outcome: &oauth.Outcome{
102+
UserAction: &oauth.UserAction{URL: "https://example.com/auth", Message: "Authorize, then retry."},
103+
FlowID: "flow-1",
104+
},
105+
tokenAfterAwait: true,
106+
cancelResult: true,
107+
}
108+
}
109+
110+
func TestOAuthTypedRegistration(t *testing.T) {
111+
for _, typed := range []bool{false, true} {
112+
name := "untyped-control"
113+
if typed {
114+
name = "typed"
115+
}
116+
t.Run(name, func(t *testing.T) {
117+
urlCaps := &mcp.ClientCapabilities{Elicitation: &mcp.ElicitationCapabilities{URL: &mcp.URLElicitationCapabilities{}}}
118+
t.Run("wire-input-required", func(t *testing.T) {
119+
fake := oauthPendingAuthenticator()
120+
calls := 0
121+
captures := make(chan json.RawMessage, 8)
122+
session := oauthRegisteredSession(t, typed, fake, "", &mcp.ClientOptions{
123+
Capabilities: urlCaps, MultiRoundTrip: &mcp.MultiRoundTripOptions{Disabled: true},
124+
}, &calls, captures)
125+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: probeToolName})
126+
require.NoError(t, err)
127+
<-captures // initialize
128+
wire := <-captures
129+
var fields map[string]json.RawMessage
130+
require.NoError(t, json.Unmarshal(wire, &fields))
131+
assert.JSONEq(t, `"input_required"`, string(fields["resultType"]), "wire: %s", wire)
132+
assert.True(t, result.NeedsInput())
133+
assert.Contains(t, result.InputRequests, oauthElicitIDPrefix+"flow-1")
134+
assert.Nil(t, result.StructuredContent)
135+
assert.Zero(t, calls)
136+
t.Logf("raw tools/call response: %s", wire)
137+
})
138+
for _, action := range []string{"accept", "decline"} {
139+
t.Run(action, func(t *testing.T) {
140+
fake := oauthPendingAuthenticator()
141+
calls, prompts := 0, 0
142+
session := oauthRegisteredSession(t, typed, fake, "", &mcp.ClientOptions{
143+
Capabilities: urlCaps,
144+
ElicitationHandler: func(_ context.Context, req *mcp.ElicitRequest) (*mcp.ElicitResult, error) {
145+
prompts++
146+
assert.Equal(t, "url", req.Params.Mode)
147+
assert.Equal(t, "https://example.com/auth", req.Params.URL)
148+
return &mcp.ElicitResult{Action: action}, nil
149+
},
150+
}, &calls)
151+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: probeToolName})
152+
require.NoError(t, err)
153+
assert.Equal(t, 1, prompts)
154+
require.Len(t, result.Content, 1)
155+
assert.False(t, result.NeedsInput())
156+
wire, err := json.Marshal(result)
157+
require.NoError(t, err)
158+
assert.Contains(t, string(wire), `"resultType":"complete"`)
159+
if action == "accept" {
160+
assert.Equal(t, 1, calls)
161+
assert.Equal(t, 1, fake.awaitCalls)
162+
assert.Equal(t, "flow-1", fake.lastAwaitFlowID)
163+
assert.Zero(t, fake.cancelCalls)
164+
if typed {
165+
output, err := json.Marshal(result.StructuredContent)
166+
require.NoError(t, err)
167+
assert.JSONEq(t, `{"status":"tool-ran"}`, string(output))
168+
}
169+
} else {
170+
assert.Zero(t, calls)
171+
assert.Zero(t, fake.awaitCalls)
172+
assert.Equal(t, 1, fake.cancelCalls)
173+
assert.Contains(t, result.Content[0].(*mcp.TextContent).Text, "declined")
174+
assert.Nil(t, result.StructuredContent)
175+
}
176+
})
177+
}
178+
for _, protocol := range []string{"", "2025-11-25"} {
179+
t.Run("manual-fallback/"+protocol, func(t *testing.T) {
180+
fake := oauthPendingAuthenticator()
181+
calls := 0
182+
session := oauthRegisteredSession(t, typed, fake, protocol, &mcp.ClientOptions{
183+
Capabilities: &mcp.ClientCapabilities{},
184+
}, &calls)
185+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: probeToolName})
186+
require.NoError(t, err)
187+
require.Len(t, result.Content, 1)
188+
assert.Equal(t, "Authorize, then retry.", result.Content[0].(*mcp.TextContent).Text)
189+
assert.False(t, result.NeedsInput())
190+
assert.Nil(t, result.StructuredContent)
191+
assert.Zero(t, calls)
192+
assert.Zero(t, fake.awaitCalls)
193+
assert.Zero(t, fake.cancelCalls)
194+
assert.Equal(t, protocol != "", fake.lastPrompter != nil)
195+
})
196+
}
197+
})
198+
}
199+
}
200+
201+
type legacyPromptAuthenticator struct {
202+
fakeAuthenticator
203+
}
204+
205+
func (f *legacyPromptAuthenticator) Authenticate(ctx context.Context, prompter oauth.Prompter) (*oauth.Outcome, error) {
206+
f.authCalls++
207+
f.lastPrompter = prompter
208+
if err := prompter.PromptURL(ctx, oauth.Prompt{URL: "https://example.com/auth", Message: "Authorize"}); err != nil {
209+
return nil, err
210+
}
211+
f.hasToken = true
212+
return nil, nil
213+
}
214+
215+
func TestOAuthTypedRegistrationLegacyElicitation(t *testing.T) {
216+
for _, typed := range []bool{false, true} {
217+
t.Run(map[bool]string{false: "untyped-control", true: "typed"}[typed], func(t *testing.T) {
218+
fake := &legacyPromptAuthenticator{}
219+
calls, prompts := 0, 0
220+
session := oauthRegisteredSession(t, typed, fake, "2025-11-25", &mcp.ClientOptions{
221+
Capabilities: &mcp.ClientCapabilities{Elicitation: &mcp.ElicitationCapabilities{URL: &mcp.URLElicitationCapabilities{}}},
222+
ElicitationHandler: func(context.Context, *mcp.ElicitRequest) (*mcp.ElicitResult, error) {
223+
prompts++
224+
return &mcp.ElicitResult{Action: "accept"}, nil
225+
},
226+
}, &calls)
227+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: probeToolName})
228+
require.NoError(t, err)
229+
assert.Equal(t, 1, prompts)
230+
assert.Equal(t, 1, calls)
231+
assert.Equal(t, 1, fake.authCalls)
232+
assert.False(t, result.NeedsInput())
233+
assert.Nil(t, result.StructuredContent)
234+
require.Len(t, result.Content, 1)
235+
assert.Equal(t, "tool-ran", result.Content[0].(*mcp.TextContent).Text)
236+
})
237+
}
238+
}
239+
240+
func TestOAuthTypedRegistrationGuardsInvalidArguments(t *testing.T) {
241+
type input struct {
242+
Mode string `json:"mode"`
243+
}
244+
for _, raw := range []string{
245+
`{}`, `{"mode":"invalid"}`, `{"mode":42}`, `[]`, `null`, `"not an object"`,
246+
} {
247+
t.Run(raw, func(t *testing.T) {
248+
arguments := json.RawMessage(raw)
249+
fake := oauthPendingAuthenticator()
250+
handlerCalls, preflightCalls, normalizerCalls := 0, 0, 0
251+
tool := inventory.NewServerToolWithContextHandlerAndSchemaOptions(
252+
mcp.Tool{
253+
Name: probeToolName,
254+
InputSchema: &jsonschema.Schema{
255+
Type: "object",
256+
Properties: map[string]*jsonschema.Schema{"mode": {Type: "string", Enum: []any{"valid"}}},
257+
Required: []string{"mode"},
258+
},
259+
},
260+
inventory.ToolsetMetadata{ID: "test"},
261+
func(context.Context, *mcp.CallToolRequest, input) (*mcp.CallToolResult, oauthProbeOutput, error) {
262+
handlerCalls++
263+
return nil, oauthProbeOutput{Status: "unexpected"}, nil
264+
},
265+
inventory.TypedSchemaOptions{
266+
Preflight: func(ctx context.Context, _ *mcp.CallToolRequest) (context.Context, *mcp.CallToolResult, error) {
267+
preflightCalls++
268+
return ctx, nil, nil
269+
},
270+
},
271+
func(raw json.RawMessage) (json.RawMessage, error) {
272+
normalizerCalls++
273+
return raw, nil
274+
},
275+
)
276+
server := mcp.NewServer(&mcp.Implementation{Name: "test", Version: "test"}, nil)
277+
tool.RegisterFunc(server, nil, createOAuthToolMiddleware(fake, discardLogger()))
278+
session := connectOAuthRegisteredServer(t, server, "", &mcp.ClientOptions{
279+
Capabilities: &mcp.ClientCapabilities{Elicitation: &mcp.ElicitationCapabilities{URL: &mcp.URLElicitationCapabilities{}}},
280+
MultiRoundTrip: &mcp.MultiRoundTripOptions{Disabled: true},
281+
})
282+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: probeToolName, Arguments: arguments})
283+
require.NoError(t, err)
284+
assert.True(t, result.NeedsInput(), "auth must win over invalid arguments")
285+
assert.Zero(t, handlerCalls)
286+
assert.Zero(t, preflightCalls)
287+
assert.Zero(t, normalizerCalls)
288+
result, err = session.CallTool(context.Background(), &mcp.CallToolParams{
289+
Name: probeToolName, Arguments: arguments,
290+
InputResponses: mcp.InputResponseMap{oauthElicitIDPrefix + "flow-1": &mcp.ElicitResult{Action: "accept"}},
291+
})
292+
require.NoError(t, err)
293+
assert.True(t, result.IsError, "authorized retries must enforce original input schema")
294+
assert.False(t, result.NeedsInput())
295+
assert.Zero(t, handlerCalls)
296+
assert.Equal(t, 1, preflightCalls)
297+
assert.Equal(t, 1, normalizerCalls)
298+
assert.Equal(t, 1, fake.awaitCalls)
299+
list, err := session.ListTools(context.Background(), nil)
300+
require.NoError(t, err)
301+
require.Len(t, list.Tools, 1)
302+
encoded, err := json.Marshal(list.Tools[0].InputSchema)
303+
require.NoError(t, err)
304+
assert.Contains(t, string(encoded), `"required":["mode"]`)
305+
assert.Contains(t, string(encoded), `"enum":["valid"]`)
306+
})
307+
}
308+
}

‎pkg/inventory/server_tool.go‎

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -465,13 +465,19 @@ func NewServerToolWithContextHandlerAndSchemaOptions[In any, Out any](
465465
preserveContent: schemaOptions.PreserveHandlerContent,
466466
}
467467
}
468-
serverTool.registerTyped = func(server *mcp.Server, registration *typedToolRegistration, era ProtocolEra, middleware ...ToolHandlerMiddleware) {
469-
switch era {
470-
case ProtocolEraLegacy:
471-
mcp.AddTool[In, any](server, registration.legacyRuntimeTool, wrapTypedHandler(handler, middleware...))
472-
default:
473-
mcp.AddTool[In, any](server, registration.modernRuntimeTool, wrapTypedHandler(handler, middleware...))
468+
serverTool.registerTyped = func(server *mcp.Server, registration *typedToolRegistration, era ProtocolEra, _ ...ToolHandlerMiddleware) {
469+
tool := *registration.runtimeTool(era)
470+
inputSchema, err := cachedResolvedInputSchema(tool.InputSchema)
471+
if err != nil {
472+
panic(fmt.Sprintf("failed to resolve input schema for tool %q: %v", tool.Name, err))
474473
}
474+
// Input validation is deferred until guards have allowed the call.
475+
// The advertised and validation schemas remain the original schemas.
476+
tool.InputSchema, err = cachedObjectInputSchema()
477+
if err != nil {
478+
panic(fmt.Sprintf("failed to prepare input schema for tool %q: %v", tool.Name, err))
479+
}
480+
mcp.AddTool[any, any](server, &tool, wrapTypedHandler(handler, inputSchema))
475481
}
476482
}
477483
return serverTool

0 commit comments

Comments
 (0)