From 86af6587d72532f6eb6b96f3e6240063e8414d87 Mon Sep 17 00:00:00 2001 From: Sam Morrow Date: Sat, 3 Oct 2026 14:33:28 +0200 Subject: [PATCH 1/2] perf: cache encoded tools/list schemas Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pkg/github/encoded_schemas_test.go | 269 ++++++++++++++++++++++++++ pkg/inventory/encoded_schemas.go | 88 +++++++++ pkg/inventory/encoded_schemas_test.go | 249 ++++++++++++++++++++++++ pkg/inventory/registry.go | 8 +- pkg/inventory/server_tool.go | 5 + 5 files changed, 618 insertions(+), 1 deletion(-) create mode 100644 pkg/github/encoded_schemas_test.go create mode 100644 pkg/inventory/encoded_schemas.go create mode 100644 pkg/inventory/encoded_schemas_test.go diff --git a/pkg/github/encoded_schemas_test.go b/pkg/github/encoded_schemas_test.go new file mode 100644 index 0000000000..09e9011ae0 --- /dev/null +++ b/pkg/github/encoded_schemas_test.go @@ -0,0 +1,269 @@ +package github + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "slices" + "testing" + + ghcontext "github.com/github/github-mcp-server/pkg/context" + "github.com/github/github-mcp-server/pkg/inventory" + "github.com/github/github-mcp-server/pkg/translations" + "github.com/github/github-mcp-server/pkg/utils" + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// The baseline uses the public single-tool registration path, which does not +// install the tools/list encoder. Both paths select the same inventory first. +func schemaListServer(ctx context.Context, inv *inventory.Inventory, optimized bool, pageSize int) *mcp.Server { + server := mcp.NewServer(&mcp.Implementation{Name: "schema-list"}, &mcp.ServerOptions{PageSize: pageSize}) + if optimized { + inv.RegisterTools(ctx, server, nil) + } else { + for _, tool := range inv.ToolsForRegistration(ctx) { + tool.RegisterFunc(server, nil) + } + } + return server +} + +func schemaListContext(protocol string, ui bool) context.Context { + return ghcontext.WithUISupport(ghcontext.WithMCPMethodInfo(context.Background(), &ghcontext.MCPMethodInfo{ + Method: inventory.MCPMethodToolsList, + ProtocolVersion: protocol, + ClientCapabilities: &mcp.ClientCapabilities{}, + }), ui) +} + +type schemaListWire struct { + Result *mcp.ListToolsResult `json:"result"` + Error json.RawMessage `json:"error"` +} + +func schemaListPages(ctx context.Context, t *testing.T, inv *inventory.Inventory, optimized bool, transport, protocol string, pageSize int) ([][]byte, []string) { + t.Helper() + var request func(string) []byte + switch transport { + case "stdio": + server := schemaListServer(ctx, inv, optimized, pageSize) + serverConn, clientConn := net.Pipe() + ss, err := server.Connect(ctx, &mcp.IOTransport{Reader: serverConn, Writer: serverConn}, nil) + require.NoError(t, err) + defer func() { _ = ss.Close() }() + defer func() { _ = clientConn.Close() }() + reader := bufio.NewReader(clientConn) + exchange := func(body string) []byte { + _, err := io.WriteString(clientConn, body+"\n") + require.NoError(t, err) + response, err := reader.ReadBytes('\n') + require.NoError(t, err) + return response + } + init := exchange(fmt.Sprintf(`{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":%q,"capabilities":{},"clientInfo":{"name":"test","version":"1"}}}`, protocol)) + require.NotContains(t, string(init), `"error"`) + _, err = io.WriteString(clientConn, `{"jsonrpc":"2.0","method":"notifications/initialized"}`+"\n") + require.NoError(t, err) + request = exchange + case "http": + handler := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { + return schemaListServer(ctx, inv, optimized, pageSize) + }, &mcp.StreamableHTTPOptions{Stateless: true, JSONResponse: true}) + request = func(body string) []byte { + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + req.Header.Set("Mcp-Protocol-Version", protocol) + req.Header.Set("Mcp-Method", "tools/list") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + require.Contains(t, []int{http.StatusOK, http.StatusBadRequest}, rec.Code, rec.Body.String()) + return rec.Body.Bytes() + } + default: + t.Fatalf("unknown transport: %s", transport) + } + var pages [][]byte + names := make([]string, 0) + listRequest := func(id int, cursor string) string { + return fmt.Sprintf(`{"jsonrpc":"2.0","id":%d,"method":"tools/list","params":{"cursor":%q,"_meta":{"io.modelcontextprotocol/protocolVersion":%q,"io.modelcontextprotocol/clientInfo":{"name":"test","version":"1"},"io.modelcontextprotocol/clientCapabilities":{}}}}`, id, cursor, protocol) + } + cursor := "" + for { + response := request(listRequest(2, cursor)) + pages = append(pages, response) + var wire schemaListWire + require.NoError(t, json.Unmarshal(response, &wire)) + require.Empty(t, wire.Error) + require.NotNil(t, wire.Result) + for _, tool := range wire.Result.Tools { + names = append(names, tool.Name) + } + if wire.Result.NextCursor == "" { + break + } + require.NotEqual(t, cursor, wire.Result.NextCursor) + cursor = wire.Result.NextCursor + } + bad := request(listRequest(3, "not-a-valid-cursor")) + var wire schemaListWire + require.NoError(t, json.Unmarshal(bad, &wire)) + require.NotEmpty(t, wire.Error) + pages = append(pages, bad) + return pages, names +} + +func TestEncodedSchemasWireParity(t *testing.T) { + type selection struct { + name string + toolsets []string + features []string + exclude []string + explicit []string + scopes []string + protocol string + ui bool + } + prototype, err := NewInventory(translations.NullTranslationHelper).WithToolsets([]string{"all"}).Build() + require.NoError(t, err) + var selections []selection + for _, toolset := range append([]inventory.ToolsetID{"all", "default"}, prototype.ToolsetIDs()...) { + for _, enabled := range []bool{false, true} { + features := []string(nil) + if enabled { + features = AllowedFeatureFlags + } + selections = append(selections, selection{ + name: fmt.Sprintf("%s/all-flags=%v", toolset, enabled), toolsets: []string{string(toolset)}, features: features, + }) + } + } + for _, feature := range AllowedFeatureFlags { + selections = append(selections, selection{name: "flag/" + feature, toolsets: []string{"all"}, features: []string{feature}}) + } + selections = append(selections, + selection{name: "insiders", toolsets: []string{"all"}, features: InsidersFeatureFlags}, + selection{name: "none", toolsets: []string{}}, + selection{name: "explicit-and-excluded", toolsets: []string{"repos"}, explicit: []string{"get_me"}, exclude: []string{"get_commit"}}, + selection{name: "scopes", toolsets: []string{"all"}, scopes: []string{"read:user"}}, + selection{name: "legacy-no-ui", toolsets: []string{"all"}, features: AllowedFeatureFlags, protocol: "2025-11-25"}, + selection{name: "apps", toolsets: []string{"all"}, features: AllowedFeatureFlags, ui: true}, + ) + for _, host := range []utils.HostType{utils.HostTypeDotcom, utils.HostTypeGHEC, utils.HostTypeGHES} { + // Reuse the same static definitions across request-time variants, just as + // both the local HTTP and remote inventory factories do. + definitions := AllTools(translations.NullTranslationHelper, WithHost(host)) + for _, readOnly := range []bool{false, true} { + for _, sel := range selections { + t.Run(fmt.Sprintf("host=%d/readonly=%v/%s", host, readOnly, sel.name), func(t *testing.T) { + protocol := sel.protocol + if protocol == "" { + protocol = inventory.ProtocolVersionMultiRoundTrip + } + ctx := schemaListContext(protocol, sel.ui) + flags := ResolveFeatureFlags(sel.features, false) + builder := inventory.NewBuilder().SetTools(definitions). + WithToolsets(sel.toolsets).WithReadOnly(readOnly). + WithTools(sel.explicit).WithExcludeTools(sel.exclude). + WithFeatureChecker(func(_ context.Context, stringFlag string) (bool, error) { + return flags[stringFlag], nil + }) + if sel.scopes != nil { + builder.WithFilter(CreateToolScopeFilter(sel.scopes)) + } + inv, err := builder.Build() + require.NoError(t, err) + wantNames := make([]string, 0) + for _, tool := range inv.ToolsForRegistration(ctx) { + wantNames = append(wantNames, tool.Tool.Name) + } + slices.Sort(wantNames) + wantNames = slices.Compact(wantNames) + for _, transport := range []string{"stdio", "http"} { + t.Run(transport, func(t *testing.T) { + before, baselineNames := schemaListPages(ctx, t, inv, false, transport, protocol, 7) + after, encodedNames := schemaListPages(ctx, t, inv, true, transport, protocol, 7) + assert.Equal(t, before, after, "complete JSON-RPC wire responses, including cursors and errors") + assert.Equal(t, wantNames, baselineNames) + assert.Equal(t, wantNames, encodedNames) + }) + } + }) + } + } + } +} + +func schemaListClient(tb testing.TB, server *mcp.Server) *mcp.ClientSession { + tb.Helper() + st, ct := mcp.NewInMemoryTransports() + ss, err := server.Connect(context.Background(), st, nil) + require.NoError(tb, err) + client := mcp.NewClient(&mcp.Implementation{Name: "benchmark"}, nil) + cs, err := client.Connect(context.Background(), ct, nil) + require.NoError(tb, err) + tb.Cleanup(func() { + _ = cs.Close() + _ = ss.Close() + }) + return cs +} + +func BenchmarkEncodedSchemas(b *testing.B) { + ctx := schemaListContext(inventory.ProtocolVersionMultiRoundTrip, false) + inv, err := NewInventory(translations.NullTranslationHelper). + WithToolsets([]string{"all"}). + WithFeatureChecker(func(context.Context, string) (bool, error) { return true, nil }). + Build() + require.NoError(b, err) + for _, optimized := range []bool{false, true} { + name := map[bool]string{false: "baseline", true: "encoded"}[optimized] + b.Run(name, func(b *testing.B) { + b.Run("ListTools", func(b *testing.B) { + cs := schemaListClient(b, schemaListServer(ctx, inv, optimized, 1000)) + result, err := cs.ListTools(ctx, nil) + require.NoError(b, err) + payload, err := json.Marshal(result) + require.NoError(b, err) + b.ReportAllocs() + b.ResetTimer() + for b.Loop() { + if _, err := cs.ListTools(ctx, nil); err != nil { + b.Fatal(err) + } + } + b.ReportMetric(float64(len(payload)), "payload-B") + b.ReportMetric(float64(len(result.Tools)), "tools") + }) + b.Run("ConstructionAndList", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + st, ct := mcp.NewInMemoryTransports() + server := schemaListServer(ctx, inv, optimized, 1000) + ss, err := server.Connect(ctx, st, nil) + if err != nil { + b.Fatal(err) + } + cs, err := mcp.NewClient(&mcp.Implementation{Name: "benchmark"}, nil).Connect(ctx, ct, nil) + if err != nil { + b.Fatal(err) + } + _, err = cs.ListTools(ctx, nil) + _ = cs.Close() + _ = ss.Close() + if err != nil { + b.Fatal(err) + } + } + }) + }) + } +} diff --git a/pkg/inventory/encoded_schemas.go b/pkg/inventory/encoded_schemas.go new file mode 100644 index 0000000000..12c7961ef0 --- /dev/null +++ b/pkg/inventory/encoded_schemas.go @@ -0,0 +1,88 @@ +package inventory + +import ( + "context" + "encoding/json" + "fmt" + "sync" + + "github.com/google/jsonschema-go/jsonschema" + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +type listedToolSchemas struct { + source any + registered *mcp.Tool +} + +type encodedSchemaKey struct { + schema *jsonschema.Schema + // Input schemas receive header annotations at registration; output schemas do not. + input bool +} + +var encodedSchemas sync.Map // encodedSchemaKey -> func() (json.RawMessage, error) + +func encodeListedSchema(source, registered, listed any, input bool) (any, error) { + schema, ok := source.(*jsonschema.Schema) + if !ok || schema == nil { + return listed, nil + } + finalSchema, ok := registered.(*jsonschema.Schema) + if !ok || finalSchema == nil { + return listed, nil + } + if listedSchema, ok := listed.(*jsonschema.Schema); !ok || listedSchema != finalSchema { + // Another middleware or a replacement tool may provide a different schema. + return listed, nil + } + + key := encodedSchemaKey{schema: schema, input: input} + encode, ok := encodedSchemas.Load(key) + if !ok { + // Key by the source definition: header annotation clones the input schema + // on each HTTP registration, but its encoded representation is unchanged. + encode, _ = encodedSchemas.LoadOrStore(key, sync.OnceValues(func() (json.RawMessage, error) { + data, err := json.Marshal(finalSchema) + return json.RawMessage(data), err + })) + } + return encode.(func() (json.RawMessage, error))() +} + +func encodedToolSchemasMiddleware(schemas map[string]listedToolSchemas) mcp.Middleware { + return func(next mcp.MethodHandler) mcp.MethodHandler { + return func(ctx context.Context, method string, request mcp.Request) (mcp.Result, error) { + result, err := next(ctx, method, request) + if err != nil || method != MCPMethodToolsList { + return result, err + } + list, ok := result.(*mcp.ListToolsResult) + if !ok || list == nil { + return result, nil + } + patched := *list + if list.Tools != nil { + patched.Tools = make([]*mcp.Tool, len(list.Tools)) + } + for i, tool := range list.Tools { + definition, ok := schemas[tool.Name] + if !ok { + patched.Tools[i] = tool + continue + } + toolCopy := *tool + toolCopy.InputSchema, err = encodeListedSchema(definition.source, definition.registered.InputSchema, tool.InputSchema, true) + if err != nil { + return nil, fmt.Errorf("encode input schema for %q: %w", tool.Name, err) + } + toolCopy.OutputSchema, err = encodeListedSchema(definition.registered.OutputSchema, definition.registered.OutputSchema, tool.OutputSchema, false) + if err != nil { + return nil, fmt.Errorf("encode output schema for %q: %w", tool.Name, err) + } + patched.Tools[i] = &toolCopy + } + return &patched, nil + } + } +} diff --git a/pkg/inventory/encoded_schemas_test.go b/pkg/inventory/encoded_schemas_test.go new file mode 100644 index 0000000000..76ea7f651a --- /dev/null +++ b/pkg/inventory/encoded_schemas_test.go @@ -0,0 +1,249 @@ +package inventory + +import ( + "context" + "encoding/json" + "errors" + "sync" + "sync/atomic" + "testing" + + "github.com/google/jsonschema-go/jsonschema" + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func listThroughEncoder(t *testing.T, definitions map[string]listedToolSchemas, list *mcp.ListToolsResult) *mcp.ListToolsResult { + t.Helper() + handler := encodedToolSchemasMiddleware(definitions)(func(context.Context, string, mcp.Request) (mcp.Result, error) { + return list, nil + }) + result, err := handler(context.Background(), MCPMethodToolsList, nil) + require.NoError(t, err) + return result.(*mcp.ListToolsResult) +} + +func TestEncodedSchemasVariantIsolation(t *testing.T) { + t.Parallel() + for _, property := range []string{"dotcom", "enterprise", "feature_enabled"} { + t.Run(property, func(t *testing.T) { + t.Parallel() + schema := &jsonschema.Schema{ + Type: "object", + Properties: map[string]*jsonschema.Schema{property: {Type: "string"}, "owner": {Type: "string"}}, + } + tool := &mcp.Tool{Name: "same_name", InputSchema: schema, OutputSchema: schema} + AnnotateHeaderParams(tool) + before := &mcp.ListToolsResult{Tools: []*mcp.Tool{tool}, NextCursor: "unchanged"} + want, err := json.Marshal(before) + require.NoError(t, err) + after := listThroughEncoder(t, map[string]listedToolSchemas{ + tool.Name: {source: schema, registered: tool}, + }, before) + got, err := json.Marshal(after) + require.NoError(t, err) + assert.Equal(t, string(want), string(got)) + assert.IsType(t, json.RawMessage{}, after.Tools[0].InputSchema) + assert.IsType(t, json.RawMessage{}, after.Tools[0].OutputSchema) + assert.NotSame(t, before, after) + assert.NotSame(t, tool, after.Tools[0]) + assert.Same(t, schema, tool.OutputSchema) + assert.IsType(t, &jsonschema.Schema{}, tool.InputSchema) + assert.Contains(t, string(got), property) + assert.NotContains(t, string(after.Tools[0].OutputSchema.(json.RawMessage)), "x-mcp-header") + assert.Contains(t, string(after.Tools[0].InputSchema.(json.RawMessage)), "x-mcp-header") + }) + } +} + +type encodingProbe struct { + calls atomic.Int64 + err error +} + +func (p *encodingProbe) MarshalJSON() ([]byte, error) { + p.calls.Add(1) + return []byte(`true`), p.err +} + +func TestEncodedSchemasConcurrentRegistration(t *testing.T) { + t.Parallel() + probe := &encodingProbe{} + schema := &jsonschema.Schema{ + Type: "object", + Properties: map[string]*jsonschema.Schema{"owner": {Type: "string"}}, + Extra: map[string]any{"x-probe": probe}, + } + original, err := json.Marshal(schema) + require.NoError(t, err) + initialCalls := probe.calls.Load() + const workers = 32 + results := make(chan json.RawMessage, workers) + var wg sync.WaitGroup + for range workers { + wg.Go(func() { + tool := &mcp.Tool{Name: "concurrent", InputSchema: schema} + AnnotateHeaderParams(tool) + before := &mcp.ListToolsResult{Tools: []*mcp.Tool{tool}} + handler := encodedToolSchemasMiddleware(map[string]listedToolSchemas{ + tool.Name: {source: schema, registered: tool}, + })(func(context.Context, string, mcp.Request) (mcp.Result, error) { + return before, nil + }) + for range 10 { + result, err := handler(context.Background(), MCPMethodToolsList, nil) + if !assert.NoError(t, err) { + return + } + after := result.(*mcp.ListToolsResult) + assert.Same(t, tool, before.Tools[0]) + assert.IsType(t, &jsonschema.Schema{}, tool.InputSchema) + data := after.Tools[0].InputSchema.(json.RawMessage) + assert.True(t, json.Valid(data)) + } + result, err := handler(context.Background(), MCPMethodToolsList, nil) + if assert.NoError(t, err) { + results <- result.(*mcp.ListToolsResult).Tools[0].InputSchema.(json.RawMessage) + } + }) + } + wg.Wait() + close(results) + assert.Equal(t, int64(1), probe.calls.Load()-initialCalls, "all registration clones share one encoding") + var first json.RawMessage + for data := range results { + if first == nil { + first = data + } + assert.Equal(t, first, data) + } + after, err := json.Marshal(schema) + require.NoError(t, err) + assert.Equal(t, original, after, "source schema must remain unchanged") +} + +func TestEncodedSchemasPreserveReplacementsAndErrors(t *testing.T) { + t.Parallel() + source := &jsonschema.Schema{Type: "object"} + registered := &mcp.Tool{Name: "tool", InputSchema: source} + replacement := &mcp.Tool{Name: "tool", InputSchema: &jsonschema.Schema{Type: "object", Description: "replacement"}} + raw := &mcp.Tool{Name: "raw", InputSchema: json.RawMessage(`{"type":"object","properties":{}}`)} + before := &mcp.ListToolsResult{Tools: []*mcp.Tool{replacement, raw}} + after := listThroughEncoder(t, map[string]listedToolSchemas{ + "tool": {source: source, registered: registered}, + "raw": {source: raw.InputSchema, registered: raw}, + }, before) + assert.Same(t, replacement.InputSchema, after.Tools[0].InputSchema) + assert.Equal(t, raw.InputSchema, after.Tools[1].InputSchema) + + broken := &jsonschema.Schema{Type: "object", Extra: map[string]any{ + "x-probe": &encodingProbe{err: errors.New("encoding failed")}, + }} + tool := &mcp.Tool{Name: "broken", InputSchema: broken} + handler := encodedToolSchemasMiddleware(map[string]listedToolSchemas{ + "broken": {source: broken, registered: tool}, + })(func(context.Context, string, mcp.Request) (mcp.Result, error) { + return &mcp.ListToolsResult{Tools: []*mcp.Tool{tool}}, nil + }) + result, err := handler(context.Background(), MCPMethodToolsList, nil) + assert.Nil(t, result) + require.ErrorContains(t, err, `encode input schema for "broken"`) + + wantErr := errors.New("downstream error") + handler = encodedToolSchemasMiddleware(nil)(func(context.Context, string, mcp.Request) (mcp.Result, error) { + return nil, wantErr + }) + _, err = handler(context.Background(), MCPMethodToolsList, nil) + assert.ErrorIs(t, err, wantErr) +} + +func TestEncodedSchemasSharedResultIsImmutable(t *testing.T) { + t.Parallel() + schema := &jsonschema.Schema{Type: "object", Description: "shared"} + tool := &mcp.Tool{Name: "shared", InputSchema: schema} + list := &mcp.ListToolsResult{Tools: []*mcp.Tool{tool}, NextCursor: "cursor"} + handler := encodedToolSchemasMiddleware(map[string]listedToolSchemas{ + tool.Name: {source: schema, registered: tool}, + })(func(context.Context, string, mcp.Request) (mcp.Result, error) { + return list, nil + }) + var wg sync.WaitGroup + for range 32 { + wg.Go(func() { + for range 10 { + result, err := handler(context.Background(), MCPMethodToolsList, nil) + if !assert.NoError(t, err) { + return + } + patched := result.(*mcp.ListToolsResult) + assert.Equal(t, "cursor", patched.NextCursor) + assert.Equal(t, "shared", patched.Tools[0].Name) + assert.Equal(t, json.RawMessage(`{"type":"object","description":"shared"}`), patched.Tools[0].InputSchema) + patched.NextCursor = "caller cursor" + patched.Tools[0].Name = "caller tool" + patched.Tools[0].InputSchema = json.RawMessage(`{}`) + patched.Tools[0] = nil + } + }) + } + wg.Wait() + assert.Same(t, tool, list.Tools[0]) + assert.Same(t, schema, tool.InputSchema) + assert.Equal(t, "shared", tool.Name) + assert.Equal(t, "cursor", list.NextCursor) +} + +func TestEncodedSchemasLeaveSDKValidationIntact(t *testing.T) { + t.Parallel() + type values struct { + Count int `json:"count"` + } + input := &jsonschema.Schema{ + Type: "object", + Properties: map[string]*jsonschema.Schema{"count": {Type: "integer"}}, + Required: []string{"count"}, + } + output := &jsonschema.Schema{ + Type: "object", + Properties: map[string]*jsonschema.Schema{"count": {Type: "integer", Minimum: new(float64)}}, + Required: []string{"count"}, + } + for _, optimized := range []bool{false, true} { + t.Run(map[bool]string{false: "baseline", true: "encoded"}[optimized], func(t *testing.T) { + var calls int + server := mcp.NewServer(&mcp.Implementation{Name: "validation"}, &mcp.ServerOptions{SchemaCache: mcp.NewSchemaCache()}) + tool := &mcp.Tool{Name: "validate", InputSchema: input, OutputSchema: output} + mcp.AddTool(server, tool, func(_ context.Context, _ *mcp.CallToolRequest, args values) (*mcp.CallToolResult, values, error) { + calls++ + return nil, args, nil + }) + if optimized { + server.AddReceivingMiddleware(encodedToolSchemasMiddleware(map[string]listedToolSchemas{ + tool.Name: {source: input, registered: tool}, + })) + } + st, ct := mcp.NewInMemoryTransports() + ss, err := server.Connect(context.Background(), st, nil) + require.NoError(t, err) + defer func() { _ = ss.Close() }() + cs, err := mcp.NewClient(&mcp.Implementation{Name: "test"}, nil).Connect(context.Background(), ct, nil) + require.NoError(t, err) + defer func() { _ = cs.Close() }() + _, err = cs.ListTools(context.Background(), nil) + require.NoError(t, err) + invalid, err := cs.CallTool(context.Background(), &mcp.CallToolParams{Name: tool.Name, Arguments: map[string]any{"count": "wrong"}}) + require.NoError(t, err) + assert.True(t, invalid.IsError) + assert.Zero(t, calls, "input validation rejects before calling the handler") + _, err = cs.CallTool(context.Background(), &mcp.CallToolParams{Name: tool.Name, Arguments: values{Count: -1}}) + require.Error(t, err, "output validation must still reject a negative result") + result, err := cs.CallTool(context.Background(), &mcp.CallToolParams{Name: tool.Name, Arguments: values{Count: 1}}) + require.NoError(t, err) + assert.False(t, result.IsError) + assert.Equal(t, 2, calls) + assert.Same(t, input, tool.InputSchema) + assert.Same(t, output, tool.OutputSchema) + }) + } +} diff --git a/pkg/inventory/registry.go b/pkg/inventory/registry.go index 02dc5b0fd7..d6a0c215cb 100644 --- a/pkg/inventory/registry.go +++ b/pkg/inventory/registry.go @@ -243,9 +243,15 @@ func shouldStripMCPAppsMetadata(ctx context.Context) bool { func (r *Inventory) RegisterTools(ctx context.Context, s *mcp.Server, deps any, middleware ...ToolHandlerMiddleware) { tools := r.ToolsForRegistration(ctx) addToolAvailabilityMiddleware(s, tools) + schemas := make(map[string]listedToolSchemas, len(tools)) for _, tool := range tools { - tool.RegisterFunc(s, deps, middleware...) + registered := tool.register(s, deps, middleware...) + schemas[registered.Name] = listedToolSchemas{ + source: tool.Tool.InputSchema, + registered: registered, + } } + s.AddReceivingMiddleware(encodedToolSchemasMiddleware(schemas)) } // RegisterResourceTemplates registers all available resource templates with the server. diff --git a/pkg/inventory/server_tool.go b/pkg/inventory/server_tool.go index 9dd708685f..ec8ee44240 100644 --- a/pkg/inventory/server_tool.go +++ b/pkg/inventory/server_tool.go @@ -142,6 +142,10 @@ func (st *ServerTool) Handler(deps any) mcp.ToolHandler { // A shallow copy of the tool is made to avoid mutating the original ServerTool. // Panics if the tool has no handler - all tools should have handlers. func (st *ServerTool) RegisterFunc(s *mcp.Server, deps any, middleware ...ToolHandlerMiddleware) { + st.register(s, deps, middleware...) +} + +func (st *ServerTool) register(s *mcp.Server, deps any, middleware ...ToolHandlerMiddleware) *mcp.Tool { handler := st.Handler(deps) // This will panic if HandlerFunc is nil for _, m := range slices.Backward(middleware) { handler = m(handler) @@ -158,6 +162,7 @@ func (st *ServerTool) RegisterFunc(s *mcp.Server, deps any, middleware ...ToolHa // No-op for tools without these params. AnnotateHeaderParams(&toolCopy) s.AddTool(&toolCopy, handler) + return &toolCopy } // HeaderParams maps owner/repo input properties to the MCP-Param-* headers a From 2aa8fac2ee2bbfd0e8f4f538e5f0e5055929f5b9 Mon Sep 17 00:00:00 2001 From: Sam Morrow Date: Sat, 3 Oct 2026 14:56:04 +0200 Subject: [PATCH 2/2] test: verify encoded schema cache lifecycle Document immutable shared schema definitions and verify bounded reuse across HTTP inventories and repeated server registration. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pkg/http/handler_test.go | 36 +++++++++++++++++++ pkg/inventory/encoded_schemas_test.go | 51 +++++++++++++++++++++++++++ pkg/inventory/registry.go | 3 ++ 3 files changed, 90 insertions(+) diff --git a/pkg/http/handler_test.go b/pkg/http/handler_test.go index c9fa1de095..9edf886af5 100644 --- a/pkg/http/handler_test.go +++ b/pkg/http/handler_test.go @@ -21,6 +21,7 @@ import ( "github.com/github/github-mcp-server/pkg/translations" "github.com/github/github-mcp-server/pkg/utils" "github.com/go-chi/chi/v5" + "github.com/google/jsonschema-go/jsonschema" "github.com/modelcontextprotocol/go-sdk/jsonrpc" "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/stretchr/testify/assert" @@ -60,6 +61,41 @@ func (f allScopesFetcher) FetchTokenScopes(_ context.Context, _ string) ([]strin var _ scopes.FetcherInterface = allScopesFetcher{} +func TestDefaultInventoryFactoryReusesSchemaDefinitions(t *testing.T) { + t.Parallel() + for _, host := range []string{"https://github.com", "https://example.ghe.com", "https://github.example.com"} { + t.Run(host, func(t *testing.T) { + factory, err := NewDefaultInventoryFactory(&ServerConfig{Host: host}, translations.NullTranslationHelper, nil, allScopesFetcher{}) + require.NoError(t, err) + request := httptest.NewRequest(http.MethodPost, "/", nil) + request = request.WithContext(ghcontext.WithToolsets(request.Context(), []string{"all"})) + first, err := factory(request) + require.NoError(t, err) + definitions := first.AllTools() + require.NotEmpty(t, definitions) + for range 10 { + next, err := factory(request) + require.NoError(t, err) + tools := next.AllTools() + require.Len(t, tools, len(definitions)) + for i, tool := range tools { + assert.Equal(t, definitions[i].Tool.Name, tool.Tool.Name) + if _, ok := definitions[i].Tool.InputSchema.(*jsonschema.Schema); ok { + assert.Same(t, definitions[i].Tool.InputSchema, tool.Tool.InputSchema) + } else { + assert.Equal(t, definitions[i].Tool.InputSchema, tool.Tool.InputSchema) + } + if _, ok := definitions[i].Tool.OutputSchema.(*jsonschema.Schema); ok { + assert.Same(t, definitions[i].Tool.OutputSchema, tool.Tool.OutputSchema) + } else { + assert.Equal(t, definitions[i].Tool.OutputSchema, tool.Tool.OutputSchema) + } + } + } + }) + } +} + func mockToolWithFeatureFlag(name, toolsetID string, readOnly bool, enableFlag, disableFlag inventory.FeatureFlag) inventory.ServerTool { tool := mockTool(name, toolsetID, readOnly) features := make([]inventory.FeatureFlag, 0, 2) diff --git a/pkg/inventory/encoded_schemas_test.go b/pkg/inventory/encoded_schemas_test.go index 76ea7f651a..c3442878bc 100644 --- a/pkg/inventory/encoded_schemas_test.go +++ b/pkg/inventory/encoded_schemas_test.go @@ -123,6 +123,57 @@ func TestEncodedSchemasConcurrentRegistration(t *testing.T) { assert.Equal(t, original, after, "source schema must remain unchanged") } +func TestEncodedSchemasRepeatedInventoryRegistration(t *testing.T) { + t.Parallel() + schema := &jsonschema.Schema{ + Type: "object", + Properties: map[string]*jsonschema.Schema{"owner": {Type: "string"}}, + } + definition := ServerTool{ + Tool: mcp.Tool{Name: "repeated", InputSchema: schema, OutputSchema: schema}, + HandlerFunc: func(any) mcp.ToolHandler { + return func(context.Context, *mcp.CallToolRequest) (*mcp.CallToolResult, error) { + return &mcp.CallToolResult{}, nil + } + }, + } + original, err := json.Marshal(schema) + require.NoError(t, err) + var first []byte + for range 20 { + inv, err := NewBuilder().SetTools([]ServerTool{definition}).WithToolsets([]string{"all"}).Build() + require.NoError(t, err) + server := mcp.NewServer(&mcp.Implementation{Name: "repeated"}, nil) + inv.ForMCPRequest(MCPMethodToolsList, "").RegisterTools(context.Background(), server, nil) + st, ct := mcp.NewInMemoryTransports() + ss, err := server.Connect(context.Background(), st, nil) + require.NoError(t, err) + cs, err := mcp.NewClient(&mcp.Implementation{Name: "test"}, nil).Connect(context.Background(), ct, nil) + require.NoError(t, err) + list, err := cs.ListTools(context.Background(), nil) + _ = cs.Close() + _ = ss.Close() + require.NoError(t, err) + data, err := json.Marshal(list) + require.NoError(t, err) + if first == nil { + first = data + } + assert.Equal(t, first, data) + } + var entries int + encodedSchemas.Range(func(key, _ any) bool { + if key.(encodedSchemaKey).schema == schema { + entries++ + } + return true + }) + assert.Equal(t, 2, entries, "new inventories and annotation clones must reuse the two source-schema keys") + after, err := json.Marshal(schema) + require.NoError(t, err) + assert.Equal(t, original, after) +} + func TestEncodedSchemasPreserveReplacementsAndErrors(t *testing.T) { t.Parallel() source := &jsonschema.Schema{Type: "object"} diff --git a/pkg/inventory/registry.go b/pkg/inventory/registry.go index d6a0c215cb..c613664936 100644 --- a/pkg/inventory/registry.go +++ b/pkg/inventory/registry.go @@ -240,6 +240,9 @@ func shouldStripMCPAppsMetadata(ctx context.Context) bool { // the client did not advertise the io.modelcontextprotocol/ui extension. The // strip happens here (rather than at Build() time) so the per-request // context, which carries the client capability, is in scope. +// +// Schema definitions must remain immutable and be reused across registrations. +// Their encodings are retained process-wide, keyed by source schema identity. func (r *Inventory) RegisterTools(ctx context.Context, s *mcp.Server, deps any, middleware ...ToolHandlerMiddleware) { tools := r.ToolsForRegistration(ctx) addToolAvailabilityMiddleware(s, tools)