Skip to content

Commit 6baa805

Browse files
fix(inventory): unwrap pointer inputs before schema inference
Match SDK input inference by unwrapping one pointer level before the cached input key and inference. Preserve nullable output pointer schemas. Cover both public constructors with actual modern and legacy discovery and calls, including structured null output. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 17595e2 commit 6baa805

4 files changed

Lines changed: 147 additions & 0 deletions

File tree

‎docs/typed-tool-schemas.md‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,9 @@ does not provide them, validates arguments before the handler runs, and
66
validates typed output. Keep business rules that JSON Schema cannot express
77
in the handler or a preflight callback.
88

9+
Input inference unwraps one pointer level, matching the SDK's object input
10+
contract. Output pointers retain nullable schemas and may return JSON null.
11+
912
```go
1013
tool := github.NewToolWithSchemaOptions[workflowInput, workflowOutput](
1114
toolset,

‎pkg/github/dependencies_test.go‎

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,9 @@ package github_test
22

33
import (
44
"context"
5+
"encoding/json"
56
"errors"
7+
"fmt"
68
"log/slog"
79
"net/http"
810
"net/http/httptest"
@@ -14,9 +16,11 @@ import (
1416
ghcontext "github.com/github/github-mcp-server/pkg/context"
1517
"github.com/github/github-mcp-server/pkg/github"
1618
"github.com/github/github-mcp-server/pkg/http/headers"
19+
"github.com/github/github-mcp-server/pkg/inventory"
1720
"github.com/github/github-mcp-server/pkg/observability"
1821
"github.com/github/github-mcp-server/pkg/observability/metrics"
1922
"github.com/github/github-mcp-server/pkg/translations"
23+
"github.com/modelcontextprotocol/go-sdk/mcp"
2024
"github.com/shurcooL/githubv4"
2125
"github.com/stretchr/testify/assert"
2226
"github.com/stretchr/testify/require"
@@ -27,6 +31,94 @@ func testExporters() observability.Exporters {
2731
return obs
2832
}
2933

34+
func TestNewToolWithSchemaOptionsPointerInputAndNullableOutput(t *testing.T) {
35+
type input struct {
36+
Query string `json:"query"`
37+
}
38+
type output struct {
39+
Query string `json:"query"`
40+
}
41+
for _, legacy := range []bool{false, true} {
42+
t.Run(fmt.Sprintf("legacy=%t", legacy), func(t *testing.T) {
43+
options := &mcp.ServerOptions{}
44+
if legacy {
45+
options.SupportedProtocolVersions = []string{"2025-11-25"}
46+
}
47+
server := mcp.NewServer(&mcp.Implementation{Name: "test", Version: "1"}, options)
48+
if legacy {
49+
server.AddReceivingMiddleware(func(next mcp.MethodHandler) mcp.MethodHandler {
50+
return func(ctx context.Context, method string, req mcp.Request) (mcp.Result, error) {
51+
if method == "server/discover" {
52+
return nil, errors.New("legacy initialize required")
53+
}
54+
return next(ctx, method, req)
55+
}
56+
})
57+
}
58+
deps := github.NewRequestDeps(newRequestDepsAPIHostResolver(t, "https://api.github.com"),
59+
"test", false, nil, translations.NullTranslationHelper, 0, nil, testExporters())
60+
server.AddReceivingMiddleware(github.InjectDepsMiddleware(deps))
61+
tool := github.NewToolWithSchemaOptions(
62+
inventory.ToolsetMetadata{ID: "test"},
63+
mcp.Tool{Name: "pointer_tool"},
64+
inventory.ScopeAccess{},
65+
inventory.TypedSchemaOptions{},
66+
func(_ context.Context, gotDeps github.ToolDependencies, _ *mcp.CallToolRequest, args *input) (*mcp.CallToolResult, *output, error) {
67+
require.Same(t, deps, gotDeps)
68+
require.NotNil(t, args)
69+
result := &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: args.Query}}}
70+
if args.Query == "null" {
71+
return result, nil, nil
72+
}
73+
return result, &output{Query: args.Query}, nil
74+
},
75+
)
76+
require.NotPanics(t, func() { tool.RegisterFunc(server, deps) })
77+
serverTransport, clientTransport := mcp.NewInMemoryTransports()
78+
serverSession, err := server.Connect(context.Background(), serverTransport, nil)
79+
require.NoError(t, err)
80+
defer serverSession.Close()
81+
client := mcp.NewClient(&mcp.Implementation{Name: "test-client", Version: "1"}, nil)
82+
session, err := client.Connect(context.Background(), clientTransport, nil)
83+
require.NoError(t, err)
84+
defer session.Close()
85+
if legacy {
86+
assert.Equal(t, "2025-11-25", session.InitializeResult().ProtocolVersion)
87+
}
88+
list, err := session.ListTools(context.Background(), nil)
89+
require.NoError(t, err)
90+
require.Len(t, list.Tools, 1)
91+
inputJSON, err := json.Marshal(list.Tools[0].InputSchema)
92+
require.NoError(t, err)
93+
assert.Contains(t, string(inputJSON), `"type":"object"`)
94+
for _, query := range []string{"value", "null"} {
95+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{
96+
Name: tool.Tool.Name, Arguments: map[string]any{"query": query},
97+
})
98+
require.NoError(t, err)
99+
require.False(t, result.IsError)
100+
if legacy {
101+
assert.Nil(t, list.Tools[0].OutputSchema)
102+
assert.Nil(t, result.StructuredContent)
103+
require.Len(t, result.Content, 1)
104+
assert.Equal(t, query, result.Content[0].(*mcp.TextContent).Text)
105+
} else {
106+
schemaJSON, err := json.Marshal(list.Tools[0].OutputSchema)
107+
require.NoError(t, err)
108+
assert.Contains(t, string(schemaJSON), `"null"`)
109+
valueJSON, err := json.Marshal(result.StructuredContent)
110+
require.NoError(t, err)
111+
if query == "null" {
112+
assert.JSONEq(t, `null`, string(valueJSON))
113+
} else {
114+
assert.JSONEq(t, `{"query":"value"}`, string(valueJSON))
115+
}
116+
}
117+
}
118+
})
119+
}
120+
}
121+
30122
type requestDepsAPIHostResolver struct {
31123
endpoint *url.URL
32124
}

‎pkg/inventory/typed_output_test.go‎

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -205,6 +205,54 @@ func TestProtocolEraForSupportedVersions(t *testing.T) {
205205
}
206206
}
207207

208+
func TestTypedPointerInputUsesObjectSchemaInBothEras(t *testing.T) {
209+
type input struct {
210+
Query string `json:"query"`
211+
}
212+
valueSchema, err := CachedInputSchemaFor[input](nil)
213+
require.NoError(t, err)
214+
pointerSchema, err := CachedInputSchemaFor[*input](nil)
215+
require.NoError(t, err)
216+
assert.Same(t, valueSchema, pointerSchema, "SDK inference unwraps one input pointer before caching")
217+
assert.Equal(t, "object", pointerSchema.Type)
218+
assert.Empty(t, pointerSchema.Types)
219+
220+
outputSchema, err := CachedSchemaFor[*input](nil)
221+
require.NoError(t, err)
222+
assert.Contains(t, outputSchema.Types, "null", "output pointers must retain nullability")
223+
224+
tool := NewServerToolWithContextHandler(
225+
mcp.Tool{Name: "pointer_input"},
226+
testToolsetMetadata("test"),
227+
func(_ context.Context, _ *mcp.CallToolRequest, args *input) (*mcp.CallToolResult, typedTestOutput, error) {
228+
require.NotNil(t, args)
229+
return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: args.Query}}},
230+
typedTestOutput{Query: args.Query}, nil
231+
},
232+
)
233+
server := mcp.NewServer(&mcp.Implementation{Name: "test", Version: "1"}, nil)
234+
require.NotPanics(t, func() { tool.RegisterFunc(server, nil) })
235+
for _, version := range []string{"", "2025-11-25"} {
236+
session := connectTypedTestClient(t, server, version)
237+
list, err := session.ListTools(context.Background(), nil)
238+
require.NoError(t, err)
239+
require.Len(t, list.Tools, 1)
240+
assert.Contains(t, mustMarshalJSON(t, list.Tools[0].InputSchema), `"type":"object"`)
241+
result, err := session.CallTool(context.Background(), &mcp.CallToolParams{
242+
Name: tool.Tool.Name, Arguments: map[string]any{"query": "pointer input decoded"},
243+
})
244+
require.NoError(t, err)
245+
require.False(t, result.IsError)
246+
if version == "" {
247+
assert.JSONEq(t, `{"query":"pointer input decoded"}`, mustMarshalJSON(t, result.StructuredContent))
248+
} else {
249+
assert.Nil(t, result.StructuredContent)
250+
require.Len(t, result.Content, 1)
251+
assert.Equal(t, "pointer input decoded", result.Content[0].(*mcp.TextContent).Text)
252+
}
253+
}
254+
}
255+
208256
func TestCachedSchemaDeepCopiesMutableMetadata(t *testing.T) {
209257
constant := any(map[string]any{"nested": []string{"const"}})
210258
nullConstant := any(nil)

‎pkg/inventory/typed_schema.go‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -357,6 +357,10 @@ func CachedInputSchemaFor[T any](options *jsonschema.ForOptions, enums ...Schema
357357
}
358358

359359
func cachedSchemaFor(goType reflect.Type, options *jsonschema.ForOptions, enums []SchemaEnum, inputSchema bool) (*jsonschema.Schema, error) {
360+
// Match SDK input inference without removing nullable output semantics.
361+
if inputSchema && goType.Kind() == reflect.Pointer {
362+
goType = goType.Elem()
363+
}
360364
optionsKey, err := schemaOptionsKey(options)
361365
if err != nil {
362366
return nil, err

0 commit comments

Comments
 (0)