From dedc33f1bfae0d0e51070aa77644c7e75944dcde Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 29 Sep 2026 21:26:35 -0700 Subject: [PATCH 1/3] feat(mcp): add opt-in tool search discovery --- .env.template | 4 + config/config.example.yaml | 4 + config/mcp.go | 20 ++ config/mcp_test.go | 25 ++ docs/advanced/configuration.mdx | 1 + docs/features/mcp-gateway.mdx | 48 +++ internal/mcpgateway/discovery.go | 439 ++++++++++++++++++++++++++ internal/mcpgateway/discovery_test.go | 251 +++++++++++++++ internal/mcpgateway/factory.go | 1 + internal/mcpgateway/service.go | 58 +++- internal/server/mcp_service.go | 10 +- internal/server/mcp_service_test.go | 10 + tests/e2e/mcp_test.go | 77 ++++- 13 files changed, 929 insertions(+), 19 deletions(-) create mode 100644 internal/mcpgateway/discovery.go create mode 100644 internal/mcpgateway/discovery_test.go diff --git a/.env.template b/.env.template index 71f7f446b..a66a9df7d 100644 --- a/.env.template +++ b/.env.template @@ -171,6 +171,10 @@ # other, so comparing them proves nothing. Add an origin only if you serve an MCP web # client from it. "*" trusts every origin and disables the check. # MCP_ALLOWED_ORIGINS=https://console.example.com +# MCP_TOOL_DISCOVERY: "off" (default) lists every tool; "search" lists only search_tools +# and call_tool, so large catalogs stop costing model context on every turn. Use it for +# clients without their own tool search. Clients override it with X-MCP-Tool-Discovery. +# MCP_TOOL_DISCOVERY=off # HTTP Client Configuration (for upstream API requests) # The standard HTTP_PROXY / HTTPS_PROXY / NO_PROXY variables route every upstream diff --git a/config/config.example.yaml b/config/config.example.yaml index 1a177aaa6..8f703f856 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -137,6 +137,10 @@ models: # # page's Origin and Host agree with each other, so comparing them proves # # nothing). Add an origin only if you serve an MCP web client from it. # allowed_origins: [] # env: MCP_ALLOWED_ORIGINS; "*" disables the check +# # off (default) lists every tool; search lists only search_tools and call_tool, +# # so large catalogs stop costing model context. Clients override it per session +# # with the X-MCP-Tool-Discovery header. +# tool_discovery: off # env: MCP_TOOL_DISCOVERY # servers: # github: # url: https://api.githubcopilot.com/mcp diff --git a/config/mcp.go b/config/mcp.go index d4d1f0592..9cf1c572a 100644 --- a/config/mcp.go +++ b/config/mcp.go @@ -42,11 +42,23 @@ type MCPConfig struct { // ("*") turns the check off and is logged as a warning at startup. AllowedOrigins []string `yaml:"allowed_origins" env:"MCP_ALLOWED_ORIGINS"` + // ToolDiscovery selects how sessions see tools by default. "off" (the + // default) lists every visible tool. "search" lists only search_tools and + // call_tool, so large catalogs stop costing context on every model turn. + // Clients override it per session with the X-MCP-Tool-Discovery header. + ToolDiscovery string `yaml:"tool_discovery" env:"MCP_TOOL_DISCOVERY"` + // Servers maps stable server slugs to upstream definitions. Slugs become // tool namespaces and URL segments, so they are restricted to [a-z0-9_-]. Servers map[string]MCPServerConfig `yaml:"servers"` } +// MCP tool discovery modes accepted in MCPConfig.ToolDiscovery. +const ( + MCPToolDiscoveryOff = "off" + MCPToolDiscoverySearch = "search" +) + // TrustAnyOrigin is the mcp.allowed_origins entry that trusts every browser // origin. It exists for deployments that enforce their own origin checks in // front of the gateway, and disables the gateway's DNS-rebinding defense. @@ -224,6 +236,14 @@ func expandMCPServerEnv(server *MCPServerConfig) { // invalid entries. It runs at load time so a bad declaration fails startup // loudly instead of silently dropping the server. func normalizeMCPConfig(cfg *MCPConfig) error { + switch mode := strings.ToLower(strings.TrimSpace(cfg.ToolDiscovery)); mode { + case "": + cfg.ToolDiscovery = MCPToolDiscoveryOff + case MCPToolDiscoveryOff, MCPToolDiscoverySearch: + cfg.ToolDiscovery = mode + default: + return fmt.Errorf("mcp.tool_discovery must be %q or %q, got %q", MCPToolDiscoveryOff, MCPToolDiscoverySearch, cfg.ToolDiscovery) + } if len(cfg.AllowedOrigins) > 0 { normalized := make([]string, 0, len(cfg.AllowedOrigins)) for _, raw := range cfg.AllowedOrigins { diff --git a/config/mcp_test.go b/config/mcp_test.go index 2be771e81..e3da3df3e 100644 --- a/config/mcp_test.go +++ b/config/mcp_test.go @@ -250,3 +250,28 @@ func TestNormalizeMCPAllowedOrigins(t *testing.T) { }) } } + +func TestNormalizeMCPToolDiscovery(t *testing.T) { + tests := []struct { + input string + want string + wantErr bool + }{ + {input: "", want: MCPToolDiscoveryOff}, + {input: "off", want: MCPToolDiscoveryOff}, + {input: " Search ", want: MCPToolDiscoverySearch}, + {input: "semantic", wantErr: true}, + } + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + cfg := MCPConfig{ToolDiscovery: tt.input} + err := normalizeMCPConfig(&cfg) + if tt.wantErr { + require.ErrorContains(t, err, "mcp.tool_discovery") + return + } + require.NoError(t, err) + assert.Equal(t, tt.want, cfg.ToolDiscovery) + }) + } +} diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 8d7e6e7d6..524f54afa 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -165,6 +165,7 @@ See [MCP Gateway](/features/mcp-gateway) for the full feature guide. | `MCP_ENABLED` | Expose the MCP endpoints `/mcp` and `/mcp/{server}` | `true` | | `MCP_SERVERS` | JSON object of upstream MCP servers, merged over the `mcp.servers` YAML map per name (entries are read-only in the dashboard) | _(none)_ | | `MCP_ALLOWED_ORIGINS` | Comma-separated browser origins (`scheme://host[:port]`) allowed to call `/mcp`; the default trusts none, which is what blocks DNS rebinding. `*` disables the check | _(none)_ | +| `MCP_TOOL_DISCOVERY` | `off` lists every tool; `search` serves `search_tools` and `call_tool` instead, so large catalogs stop filling model context. Clients override it per session with `X-MCP-Tool-Discovery` | `off` | #### Logging diff --git a/docs/features/mcp-gateway.mdx b/docs/features/mcp-gateway.mdx index 0d6cca447..9876488f5 100644 --- a/docs/features/mcp-gateway.mdx +++ b/docs/features/mcp-gateway.mdx @@ -60,6 +60,54 @@ Three ways to narrow what a client sees: - `/mcp/{slug}` — a single server with its **original** tool names. - `X-MCP-Servers: github,jira` header — a comma-separated subset on `/mcp`. +## Search tools instead of listing them + +Every tool in `tools/list` is sent to the model on every turn. With a few +large servers that can be tens of thousands of tokens. Search discovery +replaces the list with two tools: + +- `search_tools` takes keywords (or an exact tool name) and returns the + matching tools with their descriptions and input schemas. +- `call_tool` runs a found tool by `name` with its `arguments`. + +It is off by default. Turn it on for every client: + +```yaml +mcp: + tool_discovery: search # env: MCP_TOOL_DISCOVERY; default: off +``` + +Or per client, with a header that overrides the default for that session: + +```json +{ + "mcpServers": { + "gomodel": { + "type": "http", + "url": "https://your-gateway/mcp", + "headers": { + "Authorization": "Bearer sk-your-gomodel-key", + "X-MCP-Tool-Discovery": "search" + } + } + } +} +``` + +Use it for clients that send every tool to the model, such as custom agents +and SDK loops. Leave it off for clients with their own tool search, such as +Claude Code: they already defer tools, and they get per-tool permission +prompts and read-only/destructive hints only when tools are listed directly. +Through `call_tool`, every call looks like the same tool to the client. + +Search results and calls respect user paths and tool filters. Failures, such +as an unknown name or an unreachable server, come back as tool errors the +model can read and recover from. Usage entries and the request log record the +real tool name, not `call_tool`. If the catalog is only +large because of tools nobody uses, trimming it with +[tool filters](#choose-which-tools-each-server-exposes) or `X-MCP-Servers` +is simpler and keeps tools listed directly. + ## Declare servers Servers can be managed in the dashboard (**MCP Servers** page) or declared as diff --git a/internal/mcpgateway/discovery.go b/internal/mcpgateway/discovery.go new file mode 100644 index 000000000..a9b3d725b --- /dev/null +++ b/internal/mcpgateway/discovery.go @@ -0,0 +1,439 @@ +package mcpgateway + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "math" + "net/http" + "sort" + "strings" + "unicode" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/enterpilot/gomodel/config" +) + +// ToolDiscoveryHeader overrides mcp.tool_discovery for one session ("off" or +// "search"), so a client without its own tool search can opt in while other +// clients of the same gateway keep the full tool list. Read at initialize; +// unknown values fall back to the configured default. +const ToolDiscoveryHeader = "X-MCP-Tool-Discovery" + +// Meta-tool names served instead of the catalog in search discovery mode. +// CallToolName is exported so the request log can label a relayed call with +// the tool it runs. +const ( + searchToolsName = "search_tools" + CallToolName = "call_tool" +) + +const ( + defaultSearchLimit = 5 + maxSearchLimit = 20 +) + +// discoveryMode resolves the session's discovery mode from the header, falling +// back to the configured default. +func (s *Service) discoveryMode(r *http.Request) bool { + switch strings.ToLower(strings.TrimSpace(r.Header.Get(ToolDiscoveryHeader))) { + case config.MCPToolDiscoverySearch: + return true + case config.MCPToolDiscoveryOff: + return false + } + return s.searchDiscovery +} + +// indexedTool is one searchable tool with its pre-tokenized fields. +type indexedTool struct { + exposed string + upstream string + tool *mcp.Tool + + nameTerms []string + detailTerms []string // title and parameter names + descTerms []string +} + +// toolIndex is a discovery session's tool snapshot, mirroring what tools/list +// would have returned. Like registered tools, it is fixed at initialize; +// visibility and tool filters are re-checked on every search and call. +type toolIndex struct { + tools []indexedTool + aliases map[string]string // unambiguous bare name -> exposed name +} + +func (idx *toolIndex) add(exposed, upstream string, tool *mcp.Tool) { + idx.tools = append(idx.tools, indexedTool{ + exposed: exposed, + upstream: upstream, + tool: tool, + nameTerms: searchTerms(exposed), + detailTerms: searchTerms(tool.Title + " " + strings.Join(schemaPropertyNames(tool.InputSchema), " ")), + descTerms: searchTerms(tool.Description), + }) +} + +func (idx *toolIndex) lookup(name string) (indexedTool, bool) { + if exposed, ok := idx.aliases[name]; ok { + name = exposed + } + for _, tool := range idx.tools { + if tool.exposed == name { + return tool, true + } + } + return indexedTool{}, false +} + +// rankTools ranks tools by weighted keyword matches: a term found in the tool +// name counts most, then title and parameter names, then the description. +// Each term is weighted by how rare it is, so a server prefix shared by every +// tool does not drown out the words that tell tools apart. +func rankTools(query string, candidates []indexedTool, limit int) []indexedTool { + terms := dedupe(searchTerms(query)) + if len(terms) == 0 || len(candidates) == 0 { + return nil + } + exact := strings.ToLower(strings.TrimSpace(query)) + + type scored struct { + tool indexedTool + score float64 + } + weights := make([][]float64, len(candidates)) + docFreq := make([]int, len(terms)) + for i, tool := range candidates { + weights[i] = make([]float64, len(terms)) + for j, term := range terms { + switch { + case termIn(term, tool.nameTerms): + weights[i][j] = 3 + case termIn(term, tool.detailTerms): + weights[i][j] = 2 + case termIn(term, tool.descTerms): + weights[i][j] = 1 + } + if weights[i][j] > 0 { + docFreq[j]++ + } + } + } + + results := make([]scored, 0, len(candidates)) + for i, tool := range candidates { + score := 0.0 + for j := range terms { + if weights[i][j] > 0 { + score += weights[i][j] * math.Log(1+float64(len(candidates))/float64(docFreq[j])) + } + } + if strings.EqualFold(tool.exposed, exact) || strings.EqualFold(tool.tool.Name, exact) { + score += 100 + } + if score > 0 { + results = append(results, scored{tool: tool, score: score}) + } + } + sort.SliceStable(results, func(a, b int) bool { return results[a].score > results[b].score }) + if len(results) > limit { + results = results[:limit] + } + out := make([]indexedTool, len(results)) + for i, result := range results { + out[i] = result.tool + } + return out +} + +// termIn reports whether a query term matches a tool term. A query term of +// three or more letters matches as a prefix ("repo" finds "repository"), and +// a tool term matches a query term that only adds a short suffix ("issues" +// finds "issue"), covering plurals and simple stems without a stemmer. +func termIn(term string, toolTerms []string) bool { + for _, candidate := range toolTerms { + if candidate == term || + (len(term) >= 3 && strings.HasPrefix(candidate, term)) || + (len(candidate) >= 3 && len(term)-len(candidate) <= 2 && strings.HasPrefix(term, candidate)) { + return true + } + } + return false +} + +// stopWords carry no signal for telling tools apart. +var stopWords = map[string]struct{}{ + "an": {}, "and": {}, "are": {}, "by": {}, "for": {}, "from": {}, "in": {}, + "into": {}, "is": {}, "it": {}, "of": {}, "on": {}, "or": {}, "the": {}, + "this": {}, "to": {}, "with": {}, +} + +// searchTerms lowercases text and splits it into words on punctuation, +// underscores, and camelCase boundaries. Single characters and stop words +// are dropped. +func searchTerms(text string) []string { + var terms []string + var current []rune + flush := func() { + if len(current) > 1 { + term := strings.ToLower(string(current)) + if _, stop := stopWords[term]; !stop { + terms = append(terms, term) + } + } + current = current[:0] + } + var prev rune + for _, r := range text { + switch { + case !unicode.IsLetter(r) && !unicode.IsDigit(r): + flush() + case unicode.IsUpper(r) && unicode.IsLower(prev): + flush() + current = append(current, r) + default: + current = append(current, r) + } + prev = r + } + flush() + return terms +} + +func dedupe(terms []string) []string { + seen := make(map[string]struct{}, len(terms)) + out := terms[:0] + for _, term := range terms { + if _, ok := seen[term]; ok { + continue + } + seen[term] = struct{}{} + out = append(out, term) + } + return out +} + +// schemaPropertyNames returns the top-level property names of a tool's input +// schema, whatever Go type the SDK decoded it into. +func schemaPropertyNames(schema any) []string { + raw, err := json.Marshal(schema) + if err != nil { + return nil + } + var parsed struct { + Properties map[string]json.RawMessage `json:"properties"` + } + if err := json.Unmarshal(raw, &parsed); err != nil { + return nil + } + names := make([]string, 0, len(parsed.Properties)) + for name := range parsed.Properties { + names = append(names, name) + } + sort.Strings(names) + return names +} + +// registerDiscoveryTools serves search_tools and call_tool in place of the +// session's tools. The list the client sees never changes, which keeps +// provider prompt caches warm; each call is logged under the real tool name. +func (s *Service) registerDiscoveryTools(server *mcp.Server, idx *toolIndex, endpoint string) { + server.AddTool(&mcp.Tool{ + Name: searchToolsName, + Description: searchToolsDescription(idx), + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "query": map[string]any{ + "type": "string", + "description": "Keywords describing the task, such as \"create github issue\", or an exact tool name.", + }, + "limit": map[string]any{ + "type": "integer", + "minimum": 1, + "maximum": maxSearchLimit, + "description": fmt.Sprintf("Maximum number of tools to return. Default %d.", defaultSearchLimit), + }, + }, + "required": []string{"query"}, + }, + Annotations: &mcp.ToolAnnotations{Title: "Search tools", ReadOnlyHint: true}, + }, s.searchToolsHandler(idx)) + + server.AddTool(&mcp.Tool{ + Name: CallToolName, + Description: "Call a tool found with " + searchToolsName + ". Pass its exact name and arguments matching its input schema.", + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "name": map[string]any{ + "type": "string", + "description": "Tool name as returned by " + searchToolsName + ".", + }, + "arguments": map[string]any{ + "type": "object", + "description": "Arguments for the tool, matching its input schema.", + }, + }, + "required": []string{"name"}, + }, + }, s.callToolHandler(idx, endpoint)) +} + +func searchToolsDescription(idx *toolIndex) string { + counts := make(map[string]int) + var servers []string + for _, tool := range idx.tools { + if counts[tool.upstream] == 0 { + servers = append(servers, tool.upstream) + } + counts[tool.upstream]++ + } + parts := make([]string, len(servers)) + for i, server := range servers { + parts[i] = fmt.Sprintf("%s (%d)", server, counts[server]) + } + desc := "Search the tools available through this gateway by keyword. Returns matching tool names, descriptions, and input schemas; run one with " + CallToolName + "." + if len(parts) > 0 { + desc += " Servers: " + strings.Join(parts, ", ") + "." + } + return desc +} + +// searchResult is one tool in a search_tools response. +type searchResult struct { + Name string `json:"name"` + Title string `json:"title,omitempty"` + Description string `json:"description,omitempty"` + InputSchema any `json:"inputSchema"` + Annotations *mcp.ToolAnnotations `json:"annotations,omitempty"` +} + +func (s *Service) searchToolsHandler(idx *toolIndex) mcp.ToolHandler { + return func(_ context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) { + var args struct { + Query string `json:"query"` + Limit int `json:"limit"` + } + if err := unmarshalArguments(req, &args); err != nil { + return toolError("invalid " + searchToolsName + " arguments: " + err.Error()), nil + } + if strings.TrimSpace(args.Query) == "" { + return toolError("query is required"), nil + } + limit := args.Limit + if limit <= 0 { + limit = defaultSearchLimit + } + limit = min(limit, maxSearchLimit) + + matches := rankTools(args.Query, s.callableTools(req.Session, idx), limit) + if len(matches) == 0 { + return textResult(fmt.Sprintf("No tools matched %q. Try fewer or broader keywords.", args.Query)), nil + } + results := make([]searchResult, len(matches)) + for i, match := range matches { + results[i] = searchResult{ + Name: match.exposed, + Title: match.tool.Title, + Description: match.tool.Description, + InputSchema: match.tool.InputSchema, + Annotations: match.tool.Annotations, + } + } + encoded, err := json.Marshal(map[string]any{"tools": results}) + if err != nil { + return nil, fmt.Errorf("encode search results: %w", err) + } + return textResult(string(encoded)), nil + } +} + +func (s *Service) callToolHandler(idx *toolIndex, endpoint string) mcp.ToolHandler { + return func(ctx context.Context, req *mcp.CallToolRequest) (*mcp.CallToolResult, error) { + var args struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } + if err := unmarshalArguments(req, &args); err != nil { + return toolError("invalid " + CallToolName + " arguments: " + err.Error()), nil + } + target, ok := idx.lookup(strings.TrimSpace(args.Name)) + if !ok { + return toolError(fmt.Sprintf("unknown tool %q; use %s to find tool names", args.Name, searchToolsName)), nil + } + // Dispatch through the regular handler so session authorization, + // tool filters, and usage accounting match a direct tools/call. + inner := &mcp.CallToolRequest{ + Session: req.Session, + Extra: req.Extra, + Params: &mcp.CallToolParamsRaw{Name: target.exposed, Arguments: toolArguments(args.Arguments)}, + } + result, err := s.toolHandler(target.upstream, target.tool.Name, target.exposed, endpoint)(ctx, inner) + if err != nil { + // Clients relay a tool error to the model but may only surface a + // JSON-RPC error to the user, so the model could not recover. + return toolError(err.Error()), nil + } + return result, nil + } +} + +// callableTools drops index entries the session may no longer use: servers +// hidden from its user path since initialize, and tools excluded by an edited +// filter. Search must not advertise what call_tool would reject. +func (s *Service) callableTools(session *mcp.ServerSession, idx *toolIndex) []indexedTool { + allowed := make(map[string]bool) + out := make([]indexedTool, 0, len(idx.tools)) + for _, tool := range idx.tools { + visible, checked := allowed[tool.upstream] + if !checked { + visible = s.authorizeSession(session, tool.upstream) == nil + allowed[tool.upstream] = visible + } + if !visible { + continue + } + if u, ok := s.manager.get(tool.upstream); !ok || !u.toolExposed(tool.tool.Name) { + continue + } + out = append(out, tool) + } + return out +} + +// toolArguments normalizes call_tool's arguments: null means none, and a JSON +// object sent as a string — a common model mistake — is unwrapped. +func toolArguments(raw json.RawMessage) json.RawMessage { + trimmed := bytes.TrimSpace(raw) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil + } + if trimmed[0] == '"' { + var encoded string + if err := json.Unmarshal(trimmed, &encoded); err == nil && json.Valid([]byte(encoded)) { + return json.RawMessage(encoded) + } + } + return trimmed +} + +func unmarshalArguments(req *mcp.CallToolRequest, dst any) error { + if req.Params == nil || len(req.Params.Arguments) == 0 { + return nil + } + return json.Unmarshal(req.Params.Arguments, dst) +} + +func textResult(text string) *mcp.CallToolResult { + return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: text}}} +} + +func toolError(text string) *mcp.CallToolResult { + result := textResult(text) + result.IsError = true + return result +} diff --git a/internal/mcpgateway/discovery_test.go b/internal/mcpgateway/discovery_test.go new file mode 100644 index 000000000..4768d53a6 --- /dev/null +++ b/internal/mcpgateway/discovery_test.go @@ -0,0 +1,251 @@ +package mcpgateway + +import ( + "context" + "encoding/json" + "strings" + "testing" + "time" + + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSearchTerms(t *testing.T) { + tests := []struct { + input string + want []string + }{ + {input: "github_create_issue", want: []string{"github", "create", "issue"}}, + {input: "listPullRequests", want: []string{"list", "pull", "requests"}}, + {input: "Search the web, fast!", want: []string{"search", "web", "fast"}}, + {input: "a b", want: nil}, + } + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + assert.Equal(t, tt.want, searchTerms(tt.input)) + }) + } +} + +func testIndex() *toolIndex { + idx := &toolIndex{} + schema := func(props ...string) map[string]any { + properties := make(map[string]any, len(props)) + for _, prop := range props { + properties[prop] = map[string]any{"type": "string"} + } + return map[string]any{"type": "object", "properties": properties} + } + idx.add("github_create_issue", "github", &mcp.Tool{Name: "create_issue", Description: "Create a new issue in a repository", InputSchema: schema("owner", "repo", "title")}) + idx.add("github_list_issues", "github", &mcp.Tool{Name: "list_issues", Description: "List issues in a repository", InputSchema: schema("owner", "repo")}) + idx.add("github_merge_pull_request", "github", &mcp.Tool{Name: "merge_pull_request", Description: "Merge a pull request", InputSchema: schema("owner", "repo", "pullNumber")}) + idx.add("jira_create_ticket", "jira", &mcp.Tool{Name: "create_ticket", Description: "Create a Jira ticket for an issue", InputSchema: schema("project", "summary")}) + return idx +} + +func exposedNames(tools []indexedTool) []string { + names := make([]string, len(tools)) + for i, tool := range tools { + names[i] = tool.exposed + } + return names +} + +func TestRankTools(t *testing.T) { + idx := testIndex() + tests := []struct { + name string + query string + limit int + want []string + }{ + {name: "more matched terms rank higher", query: "create issue", limit: 5, want: []string{"github_create_issue", "jira_create_ticket", "github_list_issues"}}, + {name: "plural matches singular", query: "issues", limit: 1, want: []string{"github_create_issue"}}, + {name: "parameter names are searchable", query: "pull number", limit: 5, want: []string{"github_merge_pull_request"}}, + {name: "exact bare name wins", query: "create_ticket", limit: 1, want: []string{"jira_create_ticket"}}, + {name: "limit truncates", query: "github", limit: 2, want: []string{"github_create_issue", "github_list_issues"}}, + {name: "short tool terms do not match long query terms", query: "weather forecast", limit: 5, want: []string{}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, exposedNames(rankTools(tt.query, idx.tools, tt.limit))) + }) + } +} + +func TestToolArguments(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {name: "absent", input: "", want: ""}, + {name: "null", input: "null", want: ""}, + {name: "object", input: ` {"a":1} `, want: `{"a":1}`}, + {name: "object encoded as string", input: `"{\"a\":1}"`, want: `{"a":1}`}, + {name: "plain string passes through", input: `"hello"`, want: `"hello"`}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, string(toolArguments(json.RawMessage(tt.input)))) + }) + } +} + +func searchTools(t *testing.T, session *mcp.ClientSession, query string) []searchResult { + t.Helper() + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: searchToolsName, + Arguments: map[string]any{"query": query}, + }) + require.NoError(t, err) + require.False(t, result.IsError, "search_tools failed: %#v", result.Content) + require.Len(t, result.Content, 1) + text, ok := result.Content[0].(*mcp.TextContent) + require.True(t, ok) + if !strings.HasPrefix(text.Text, "{") { + return nil + } + var payload struct { + Tools []searchResult `json:"tools"` + } + require.NoError(t, json.Unmarshal([]byte(text.Text), &payload)) + return payload.Tools +} + +func TestSearchDiscoveryServesMetaToolsAndRelaysCalls(t *testing.T) { + alphaURL := newTestUpstream(t, "alpha", addEchoTool("echo")) + betaURL := newTestUpstream(t, "beta", addEchoTool("search")) + usageLog := &recordingUsageLogger{} + service, gatewayURL := newTestService(t, usageLog, + testSpec("alpha", alphaURL, nil), + testSpec("beta", betaURL, nil), + ) + service.searchDiscovery = true + + session := connectClient(t, gatewayURL+"/mcp", map[string]string{"X-Request-ID": "req-1"}) + assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, session)) + assert.Contains(t, session.InitializeResult().Instructions, searchToolsName) + + results := searchTools(t, session, "echo") + require.Len(t, results, 2, "beta_search matches through its description") + assert.Equal(t, "alpha_echo", results[0].Name, "a name match ranks first") + assert.Equal(t, "echoes its input", results[0].Description) + assert.NotNil(t, results[0].InputSchema) + + for _, name := range []string{"alpha_echo", "echo"} { + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: CallToolName, + Arguments: map[string]any{"name": name, "arguments": map[string]any{"value": 1}}, + }) + require.NoError(t, err) + require.False(t, result.IsError) + text, ok := result.Content[0].(*mcp.TextContent) + require.True(t, ok) + assert.Equal(t, `echo:{"value":1}`, text.Text) + } + + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) && len(usageLog.all()) < 2 { + time.Sleep(10 * time.Millisecond) + } + entries := usageLog.all() + require.Len(t, entries, 2, "only relayed calls are usage entries, not searches") + for _, entry := range entries { + assert.Equal(t, "alpha_echo", entry.Model) + assert.Equal(t, "alpha", entry.ProviderName) + assert.Equal(t, "req-1", entry.RequestID) + } +} + +func TestSearchDiscoveryRejectsUnknownToolAsToolError(t *testing.T) { + url := newTestUpstream(t, "alpha", addEchoTool("echo")) + service, gatewayURL := newTestService(t, nil, testSpec("alpha", url, nil)) + service.searchDiscovery = true + + session := connectClient(t, gatewayURL+"/mcp", nil) + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: CallToolName, + Arguments: map[string]any{"name": "alpha_missing"}, + }) + require.NoError(t, err) + assert.True(t, result.IsError) + + result, err = session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: searchToolsName, + Arguments: map[string]any{"query": " "}, + }) + require.NoError(t, err) + assert.True(t, result.IsError) +} + +func TestToolDiscoveryHeaderOverridesDefault(t *testing.T) { + url := newTestUpstream(t, "alpha", addEchoTool("echo")) + service, gatewayURL := newTestService(t, nil, testSpec("alpha", url, nil)) + + optIn := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "Search"}) + assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, optIn)) + + unknown := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "semantic"}) + assert.Equal(t, []string{"alpha_echo"}, listToolNames(t, unknown)) + + service.searchDiscovery = true + optOut := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "off"}) + assert.Equal(t, []string{"alpha_echo"}, listToolNames(t, optOut)) + + pinned := connectClient(t, gatewayURL+"/mcp/alpha", nil) + assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, pinned)) + results := searchTools(t, pinned, "echo") + require.NotEmpty(t, results) + assert.Equal(t, "echo", results[0].Name, "a pinned endpoint keeps original names") +} + +func TestSearchDiscoveryHidesToolsExcludedAfterInitialize(t *testing.T) { + url := newTestUpstream(t, "alpha", func(server *mcp.Server) { + addEchoTool("read")(server) + addEchoTool("write")(server) + }) + spec := testSpec("alpha", url, nil) + service, gatewayURL := newTestService(t, nil, spec) + service.searchDiscovery = true + + session := connectClient(t, gatewayURL+"/mcp", nil) + require.Len(t, searchTools(t, session, "echoes"), 2) + + spec.DisallowedTools = []string{"write"} + service.manager.Apply([]ServerSpec{spec}) + + results := searchTools(t, session, "echoes") + require.Len(t, results, 1) + assert.Equal(t, "alpha_read", results[0].Name) + + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: CallToolName, + Arguments: map[string]any{"name": "alpha_write"}, + }) + require.NoError(t, err, "failures reach the model as tool errors, not JSON-RPC errors") + assert.True(t, result.IsError) + require.Len(t, result.Content, 1) + text, ok := result.Content[0].(*mcp.TextContent) + require.True(t, ok) + assert.Contains(t, text.Text, "excluded") +} + +func TestSearchDiscoveryReportsUpstreamFailureAsToolError(t *testing.T) { + url := newTestUpstream(t, "alpha", addEchoTool("echo")) + usageLog := &recordingUsageLogger{} + service, gatewayURL := newTestService(t, usageLog, testSpec("alpha", url, nil)) + service.searchDiscovery = true + + session := connectClient(t, gatewayURL+"/mcp", nil) + service.manager.Apply(nil) // the upstream disappears after initialize + + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ + Name: CallToolName, + Arguments: map[string]any{"name": "alpha_echo"}, + }) + require.NoError(t, err) + assert.True(t, result.IsError) +} diff --git a/internal/mcpgateway/factory.go b/internal/mcpgateway/factory.go index 45b139e83..801a6c213 100644 --- a/internal/mcpgateway/factory.go +++ b/internal/mcpgateway/factory.go @@ -80,6 +80,7 @@ func newResult(ctx context.Context, cfg *config.Config, storeConn storage.Storag UsageLogger: usageLogger, UserPathHeader: cfg.Server.UserPathHeader, AllowedOrigins: cfg.MCP.AllowedOrigins, + ToolDiscovery: cfg.MCP.ToolDiscovery, }) if err != nil { return nil, err diff --git a/internal/mcpgateway/service.go b/internal/mcpgateway/service.go index 205157366..8da30b31c 100644 --- a/internal/mcpgateway/service.go +++ b/internal/mcpgateway/service.go @@ -14,6 +14,7 @@ import ( "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/usage" "github.com/enterpilot/gomodel/internal/version" @@ -46,6 +47,9 @@ type Service struct { usageLogger usage.LoggerInterface userPathHeader string configSpecs map[string]ServerSpec + // searchDiscovery is the default for sessions that do not send + // ToolDiscoveryHeader: serve search_tools/call_tool instead of the catalog. + searchDiscovery bool handler http.Handler origins *originGuard @@ -92,20 +96,24 @@ type Options struct { // AllowedOrigins are the browser origins permitted to reach the endpoint. // Empty trusts none, which is the default; see originGuard. AllowedOrigins []string + // ToolDiscovery is the default discovery mode, config.MCPToolDiscoveryOff + // or config.MCPToolDiscoverySearch. + ToolDiscovery string } // NewService builds the gateway service and starts connecting to the merged // server set. Upstream connects are asynchronous; construction never blocks. func NewService(ctx context.Context, opts Options) (*Service, error) { s := &Service{ - manager: NewManager(opts.HTTPClient), - store: opts.Store, - usageLogger: opts.UsageLogger, - userPathHeader: core.UserPathHeaderName(opts.UserPathHeader), - configSpecs: opts.ConfigServers, - bindings: make(map[string]sessionBinding), - requestCancels: make(map[uint64]context.CancelFunc), - stop: make(chan struct{}), + manager: NewManager(opts.HTTPClient), + store: opts.Store, + usageLogger: opts.UsageLogger, + userPathHeader: core.UserPathHeaderName(opts.UserPathHeader), + configSpecs: opts.ConfigServers, + searchDiscovery: opts.ToolDiscovery == config.MCPToolDiscoverySearch, + bindings: make(map[string]sessionBinding), + requestCancels: make(map[uint64]context.CancelFunc), + stop: make(chan struct{}), } guard, err := newOriginGuard(opts.AllowedOrigins) if err != nil { @@ -313,12 +321,14 @@ type requestScope struct { userPath string pinned string include map[string]struct{} + discovery bool } func (s *Service) scopeFromRequest(r *http.Request) requestScope { scope := requestScope{ authKeyID: core.GetAuthKeyID(r.Context()), userPath: core.UserPathFromContext(r.Context()), + discovery: s.discoveryMode(r), } if pinned, ok := r.Context().Value(pinnedServerKey{}).(string); ok { scope.pinned = pinned @@ -369,6 +379,10 @@ func (s *Service) getServer(r *http.Request) *mcp.Server { endpoint = "/mcp/" + scope.pinned } + var index *toolIndex + if scope.discovery { + index = &toolIndex{} + } toolOwners := make(map[string]string) promptOwners := make(map[string]string) resourceOwners := make(map[string]string) @@ -377,14 +391,19 @@ func (s *Service) getServer(r *http.Request) *mcp.Server { if snapshot == nil { continue } - s.registerTools(server, view.Spec.Name, snapshot, prefixNames, endpoint, toolOwners) + s.registerTools(server, view.Spec.Name, snapshot, prefixNames, endpoint, toolOwners, index) s.registerPrompts(server, view.Spec.Name, snapshot, prefixNames, promptOwners) s.registerResources(server, view.Spec.Name, snapshot, resourceOwners) } + var aliases map[string]string if prefixNames { - if aliases := bareToolAliases(toolOwners); len(aliases) > 0 { - server.AddReceivingMiddleware(bareToolCallMiddleware(aliases)) - } + aliases = bareToolAliases(toolOwners) + } + if index != nil { + index.aliases = aliases + s.registerDiscoveryTools(server, index, endpoint) + } else if len(aliases) > 0 { + server.AddReceivingMiddleware(bareToolCallMiddleware(aliases)) } return server } @@ -491,6 +510,12 @@ func (s *Service) composeInstructions(scope requestScope, views []ServerView) st fmt.Fprintf(&b, "GoModel MCP gateway aggregating %d server(s): %s. Tools and prompts are namespaced as {server}%s{name}.", len(names), strings.Join(names, ", "), namespaceSeparator) } + if scope.discovery { + if b.Len() > 0 { + b.WriteString(" ") + } + fmt.Fprintf(&b, "Tools are not listed directly: find them with %s and run them with %s.", searchToolsName, CallToolName) + } for _, view := range views { snapshot, _ := s.upstreamCatalog(view.Spec.Name) if snapshot == nil || strings.TrimSpace(snapshot.instructions) == "" { @@ -512,8 +537,9 @@ func (s *Service) composeInstructions(scope requestScope, views []ServerView) st // registerTools adds one upstream's tools to a session server. Tool metadata // and valid schemas relay verbatim; only the name is prefixed on the // aggregated endpoint. Arguments relay as raw JSON — validation belongs to -// the upstream. -func (s *Service) registerTools(server *mcp.Server, upstreamName string, snapshot *catalog, prefix bool, endpoint string, owners map[string]string) { +// the upstream. A non-nil index collects the tools for search discovery +// instead of listing them. +func (s *Service) registerTools(server *mcp.Server, upstreamName string, snapshot *catalog, prefix bool, endpoint string, owners map[string]string, index *toolIndex) { for _, tool := range snapshot.tools { exposed := tool.Name if prefix { @@ -525,6 +551,10 @@ func (s *Service) registerTools(server *mcp.Server, upstreamName string, snapsho continue } owners[exposed] = upstreamName + if index != nil { + index.add(exposed, upstreamName, tool) + continue + } clone := *tool clone.Name = exposed server.AddTool(&clone, s.toolHandler(upstreamName, tool.Name, exposed, endpoint)) diff --git a/internal/server/mcp_service.go b/internal/server/mcp_service.go index 79edd9e63..1566d99a0 100644 --- a/internal/server/mcp_service.go +++ b/internal/server/mcp_service.go @@ -95,8 +95,9 @@ func enrichMCPAuditEntry(c *echo.Context, logBodies bool) { } // mcpAuditLabel derives the request-log label from one JSON-RPC frame: the -// tool/prompt name for calls, otherwise the method. Empty means unlabelable -// (a bare response or malformed frame). +// tool/prompt name for calls, otherwise the method. A search-discovery +// call_tool is labelled with the tool it runs, matching its usage entry. +// Empty means unlabelable (a bare response or malformed frame). func mcpAuditLabel(body []byte) string { method := strings.TrimSpace(gjson.GetBytes(body, "method").String()) if method == "" { @@ -104,6 +105,11 @@ func mcpAuditLabel(body []byte) string { } if name := strings.TrimSpace(gjson.GetBytes(body, "params.name").String()); name != "" && (method == "tools/call" || method == "prompts/get") { + if method == "tools/call" && name == mcpgateway.CallToolName { + if inner := strings.TrimSpace(gjson.GetBytes(body, "params.arguments.name").String()); inner != "" { + return inner + } + } return name } return method diff --git a/internal/server/mcp_service_test.go b/internal/server/mcp_service_test.go index 4ced022d6..15dca2222 100644 --- a/internal/server/mcp_service_test.go +++ b/internal/server/mcp_service_test.go @@ -31,6 +31,16 @@ func TestMCPAuditLabel(t *testing.T) { body: `{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"github_create_issue","arguments":{}}}`, want: "github_create_issue", }, + { + name: "discovery call_tool labels with the tool it runs", + body: `{"jsonrpc":"2.0","id":5,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"github_create_issue","arguments":{}}}}`, + want: "github_create_issue", + }, + { + name: "call_tool without a target keeps its own name", + body: `{"jsonrpc":"2.0","id":6,"method":"tools/call","params":{"name":"call_tool","arguments":{}}}`, + want: "call_tool", + }, { name: "prompts/get labels with the prompt name", body: `{"jsonrpc":"2.0","id":4,"method":"prompts/get","params":{"name":"github_triage"}}`, diff --git a/tests/e2e/mcp_test.go b/tests/e2e/mcp_test.go index 95bdd116d..7c914080e 100644 --- a/tests/e2e/mcp_test.go +++ b/tests/e2e/mcp_test.go @@ -41,7 +41,12 @@ func startMockMCPServer(t *testing.T, name string, tools ...string) *httptest.Se func newE2EMCPGateway(t *testing.T, specs map[string]mcpgateway.ServerSpec) *mcpgateway.Service { t.Helper() - service, err := mcpgateway.NewService(context.Background(), mcpgateway.Options{ConfigServers: specs}) + return newE2EMCPGatewayWithOptions(t, mcpgateway.Options{ConfigServers: specs}) +} + +func newE2EMCPGatewayWithOptions(t *testing.T, opts mcpgateway.Options) *mcpgateway.Service { + t.Helper() + service, err := mcpgateway.NewService(context.Background(), opts) require.NoError(t, err) t.Cleanup(service.Close) @@ -70,7 +75,8 @@ func e2eMCPSpec(name, url string) mcpgateway.ServerSpec { // bearerTransport injects the gateway API key the way MCP clients configure // custom headers. type bearerTransport struct { - token string + token string + headers map[string]string } func (t *bearerTransport) RoundTrip(req *http.Request) (*http.Response, error) { @@ -78,15 +84,23 @@ func (t *bearerTransport) RoundTrip(req *http.Request) (*http.Response, error) { if t.token != "" { clone.Header.Set("Authorization", "Bearer "+t.token) } + for name, value := range t.headers { + clone.Header.Set(name, value) + } return http.DefaultTransport.RoundTrip(clone) } func connectMCPClient(t *testing.T, endpoint, token string) *sdk.ClientSession { + t.Helper() + return connectMCPClientWithHeaders(t, endpoint, token, nil) +} + +func connectMCPClientWithHeaders(t *testing.T, endpoint, token string, headers map[string]string) *sdk.ClientSession { t.Helper() client := sdk.NewClient(&sdk.Implementation{Name: "e2e-client", Version: "1"}, nil) session, err := client.Connect(context.Background(), &sdk.StreamableClientTransport{ Endpoint: endpoint, - HTTPClient: &http.Client{Transport: &bearerTransport{token: token}}, + HTTPClient: &http.Client{Transport: &bearerTransport{token: token, headers: headers}}, }, nil) require.NoError(t, err, "MCP client failed to connect through the gateway") t.Cleanup(func() { _ = session.Close() }) @@ -208,3 +222,60 @@ func TestMCPGatewayDisabled(t *testing.T) { defer resp.Body.Close() assert.Equal(t, http.StatusNotFound, resp.StatusCode) } + +func e2eToolNames(t *testing.T, session *sdk.ClientSession) []string { + t.Helper() + tools, err := session.ListTools(context.Background(), nil) + require.NoError(t, err) + names := make([]string, 0, len(tools.Tools)) + for _, tool := range tools.Tools { + names = append(names, tool.Name) + } + return names +} + +// TestMCPGatewaySearchDiscovery drives search discovery through the fully +// wired server: the configured default, the per-session header override, and +// a search-then-call round trip to the upstream. +func TestMCPGatewaySearchDiscovery(t *testing.T) { + alpha := startMockMCPServer(t, "alpha", "echo") + beta := startMockMCPServer(t, "beta", "search", "fetch") + gateway := newE2EMCPGatewayWithOptions(t, mcpgateway.Options{ + ConfigServers: map[string]mcpgateway.ServerSpec{ + "alpha": e2eMCPSpec("alpha", alpha.URL), + "beta": e2eMCPSpec("beta", beta.URL), + }, + ToolDiscovery: "search", + }) + + srv := setupE2EServer(t, e2eServerOptions{masterKey: "sk-e2e-master", mcpGateway: gateway}) + ts := httptest.NewServer(srv) + t.Cleanup(ts.Close) + + session := connectMCPClient(t, ts.URL+"/mcp", "sk-e2e-master") + assert.Equal(t, []string{"call_tool", "search_tools"}, e2eToolNames(t, session)) + + found, err := session.CallTool(context.Background(), &sdk.CallToolParams{ + Name: "search_tools", + Arguments: map[string]any{"query": "fetch"}, + }) + require.NoError(t, err) + require.False(t, found.IsError) + text, ok := found.Content[0].(*sdk.TextContent) + require.True(t, ok) + assert.Contains(t, text.Text, `"name":"beta_fetch"`) + + result, err := session.CallTool(context.Background(), &sdk.CallToolParams{ + Name: "call_tool", + Arguments: map[string]any{"name": "beta_fetch", "arguments": map[string]any{"url": "x"}}, + }) + require.NoError(t, err) + require.False(t, result.IsError) + text, ok = result.Content[0].(*sdk.TextContent) + require.True(t, ok) + assert.Equal(t, `fetch:{"url":"x"}`, text.Text) + + optOut := connectMCPClientWithHeaders(t, ts.URL+"/mcp", "sk-e2e-master", + map[string]string{mcpgateway.ToolDiscoveryHeader: "off"}) + assert.Equal(t, []string{"alpha_echo", "beta_fetch", "beta_search"}, e2eToolNames(t, optOut)) +} From 26eaa91303e28ff5ef9da539e13823a8081ac064 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 1 Oct 2026 08:17:39 -0700 Subject: [PATCH 2/3] fix(mcp): find short exact tool names and log resolved bare names --- internal/mcpgateway/discovery.go | 4 ++- internal/mcpgateway/discovery_test.go | 10 +++++++ internal/mcpgateway/service.go | 32 +++++++++++++++----- internal/mcpgateway/service_test.go | 28 +++++++++++++++-- internal/server/mcp_service.go | 24 ++++++++++----- internal/server/mcp_service_test.go | 43 ++++++++++++++++++++++++--- 6 files changed, 120 insertions(+), 21 deletions(-) diff --git a/internal/mcpgateway/discovery.go b/internal/mcpgateway/discovery.go index a9b3d725b..c1220312b 100644 --- a/internal/mcpgateway/discovery.go +++ b/internal/mcpgateway/discovery.go @@ -94,8 +94,10 @@ func (idx *toolIndex) lookup(name string) (indexedTool, bool) { // Each term is weighted by how rare it is, so a server prefix shared by every // tool does not drown out the words that tell tools apart. func rankTools(query string, candidates []indexedTool, limit int) []indexedTool { + // No early return on an empty term list: a stop word or single letter + // yields no terms but can still be an exact tool name. terms := dedupe(searchTerms(query)) - if len(terms) == 0 || len(candidates) == 0 { + if len(candidates) == 0 { return nil } exact := strings.ToLower(strings.TrimSpace(query)) diff --git a/internal/mcpgateway/discovery_test.go b/internal/mcpgateway/discovery_test.go index 4768d53a6..e9d32d901 100644 --- a/internal/mcpgateway/discovery_test.go +++ b/internal/mcpgateway/discovery_test.go @@ -75,6 +75,16 @@ func TestRankTools(t *testing.T) { } } +func TestRankToolsFindsExactNamesWithoutSearchTerms(t *testing.T) { + idx := &toolIndex{} + idx.add("in", "alpha", &mcp.Tool{Name: "in", InputSchema: map[string]any{"type": "object"}}) + idx.add("x", "alpha", &mcp.Tool{Name: "x", InputSchema: map[string]any{"type": "object"}}) + + assert.Equal(t, []string{"in"}, exposedNames(rankTools("in", idx.tools, 5)), "stop word") + assert.Equal(t, []string{"x"}, exposedNames(rankTools("X", idx.tools, 5)), "single letter") + assert.Empty(t, rankTools("the", idx.tools, 5)) +} + func TestToolArguments(t *testing.T) { tests := []struct { name string diff --git a/internal/mcpgateway/service.go b/internal/mcpgateway/service.go index 8da30b31c..e78654e3b 100644 --- a/internal/mcpgateway/service.go +++ b/internal/mcpgateway/service.go @@ -79,6 +79,9 @@ type sessionBinding struct { userPath string pinned string lastSeen time.Time + // toolAliases maps the session's unambiguous bare tool names to their + // namespaced names, so request logs can name the tool a call resolved to. + toolAliases map[string]string } // Options configures NewService. @@ -351,6 +354,9 @@ func (s *Service) scopeFromRequest(r *http.Request) requestScope { func (s *Service) getServer(r *http.Request) *mcp.Server { scope := s.scopeFromRequest(r) views := s.visibleServers(scope) + // Assigned below; the SDK asks for the session ID only after getServer + // returns, so the binding sees the final map. + var aliases map[string]string server := mcp.NewServer(&mcp.Implementation{ Name: "gomodel", @@ -368,7 +374,7 @@ func (s *Service) getServer(r *http.Request) *mcp.Server { }, GetSessionID: func() string { id := rand.Text() - s.bindSession(id, scope.authKeyID, scope.userPath, scope.pinned) + s.bindSession(id, scope.authKeyID, scope.userPath, scope.pinned, aliases) return id }, }) @@ -395,7 +401,6 @@ func (s *Service) getServer(r *http.Request) *mcp.Server { s.registerPrompts(server, view.Spec.Name, snapshot, prefixNames, promptOwners) s.registerResources(server, view.Spec.Name, snapshot, resourceOwners) } - var aliases map[string]string if prefixNames { aliases = bareToolAliases(toolOwners) } @@ -671,17 +676,30 @@ func (s *Service) authorizeSessionID(sessionID, upstreamName string) error { } // bindSession records the principal a new session was initialized under. -func (s *Service) bindSession(sessionID, authKeyID, userPath, pinned string) { +func (s *Service) bindSession(sessionID, authKeyID, userPath, pinned string, toolAliases map[string]string) { s.bindMu.Lock() s.bindings[sessionID] = sessionBinding{ - authKeyID: authKeyID, - userPath: userPath, - pinned: pinned, - lastSeen: time.Now(), + authKeyID: authKeyID, + userPath: userPath, + pinned: pinned, + lastSeen: time.Now(), + toolAliases: toolAliases, } s.bindMu.Unlock() } +// CanonicalToolName resolves a bare tool name the session accepts on the +// aggregated endpoint to the namespaced name it runs, matching usage +// entries. Any other name is returned unchanged. +func (s *Service) CanonicalToolName(sessionID, name string) string { + s.bindMu.Lock() + defer s.bindMu.Unlock() + if exposed, ok := s.bindings[sessionID].toolAliases[name]; ok { + return exposed + } + return name +} + // touchBinding refreshes a known session binding and reports whether the // caller's authenticated identity, user path, and endpoint pin match it. // Unknown session IDs pass through: the SDK rejects them itself, and bindings diff --git a/internal/mcpgateway/service_test.go b/internal/mcpgateway/service_test.go index 87a43f500..1d68af9f4 100644 --- a/internal/mcpgateway/service_test.go +++ b/internal/mcpgateway/service_test.go @@ -300,6 +300,30 @@ func TestAggregatedEndpointAcceptsUniqueBareToolName(t *testing.T) { require.Equal(t, "alpha", entries[0].ProviderName) } +func TestCanonicalToolNameResolvesSessionAliases(t *testing.T) { + alphaURL := newTestUpstream(t, "alpha", addEchoTool("echo")) + betaURL := newTestUpstream(t, "beta", func(server *mcp.Server) { + addEchoTool("echo")(server) + addEchoTool("search")(server) + }) + gammaURL := newTestUpstream(t, "gamma", addEchoTool("fetch")) + service, gatewayURL := newTestService(t, nil, + testSpec("alpha", alphaURL, nil), + testSpec("beta", betaURL, nil), + testSpec("gamma", gammaURL, nil), + ) + + aggregated := connectClient(t, gatewayURL+"/mcp", nil).ID() + assert.Equal(t, "beta_search", service.CanonicalToolName(aggregated, "search")) + assert.Equal(t, "gamma_fetch", service.CanonicalToolName(aggregated, "fetch")) + assert.Equal(t, "echo", service.CanonicalToolName(aggregated, "echo"), "ambiguous names stay as sent") + assert.Equal(t, "alpha_echo", service.CanonicalToolName(aggregated, "alpha_echo")) + + pinned := connectClient(t, gatewayURL+"/mcp/beta", nil).ID() + assert.Equal(t, "search", service.CanonicalToolName(pinned, "search"), "pinned endpoints use original names") + assert.Equal(t, "search", service.CanonicalToolName("unknown-session", "search")) +} + func TestAggregatedEndpointRejectsAmbiguousBareToolName(t *testing.T) { alphaURL := newTestUpstream(t, "alpha", addEchoTool("echo")) betaURL := newTestUpstream(t, "beta", addEchoTool("echo")) @@ -527,9 +551,9 @@ func TestAuthorizeSessionFailsClosedWithoutBinding(t *testing.T) { err := service.authorizeSessionID("deleted-session", "alpha") require.ErrorIs(t, err, ErrServerNotVisible) - service.bindSession("live", "", "/staff", "") + service.bindSession("live", "", "/staff", "", nil) require.NoError(t, service.authorizeSessionID("live", "alpha")) - service.bindSession("contractor", "", "/contractors/acme", "") + service.bindSession("contractor", "", "/contractors/acme", "", nil) require.ErrorIs(t, service.authorizeSessionID("contractor", "alpha"), ErrServerNotVisible) } diff --git a/internal/server/mcp_service.go b/internal/server/mcp_service.go index 1566d99a0..410cb7566 100644 --- a/internal/server/mcp_service.go +++ b/internal/server/mcp_service.go @@ -41,7 +41,10 @@ func (s *mcpService) handle(c *echo.Context, pinnedServer string) error { // so they count against user-path rate limits and budget gates. GET (the // notification stream) and DELETE (session teardown) stay free. if c.Request().Method == http.MethodPost { - enrichMCPAuditEntry(c, s.logBodies) + sessionID := strings.TrimSpace(c.Request().Header.Get("Mcp-Session-Id")) + enrichMCPAuditEntry(c, s.logBodies, func(name string) string { + return s.gateway.CanonicalToolName(sessionID, name) + }) release, err := enforceRateLimit(c, s.rateLimiter, rateLimitRoute{}) if err != nil { return handleError(c, err) @@ -74,7 +77,7 @@ func (s *mcpService) handle(c *echo.Context, pinnedServer string) error { // the JSON-RPC frame as the request body when body logging is on. The body is // restored for the gateway handler; the body-limit middleware has already // bounded its size. -func enrichMCPAuditEntry(c *echo.Context, logBodies bool) { +func enrichMCPAuditEntry(c *echo.Context, logBodies bool, resolveTool func(string) string) { req := c.Request() if req.Body == nil { return @@ -86,7 +89,7 @@ func enrichMCPAuditEntry(c *echo.Context, logBodies bool) { return } - if label := mcpAuditLabel(body); label != "" { + if label := mcpAuditLabel(body, resolveTool); label != "" { auditlog.EnrichEntry(c, label, "mcp") } if logBodies { @@ -96,20 +99,27 @@ func enrichMCPAuditEntry(c *echo.Context, logBodies bool) { // mcpAuditLabel derives the request-log label from one JSON-RPC frame: the // tool/prompt name for calls, otherwise the method. A search-discovery -// call_tool is labelled with the tool it runs, matching its usage entry. +// call_tool is labelled with the tool it runs, and resolveTool (optional) +// maps a bare tool name to the namespaced one, both matching usage entries. // Empty means unlabelable (a bare response or malformed frame). -func mcpAuditLabel(body []byte) string { +func mcpAuditLabel(body []byte, resolveTool func(string) string) string { method := strings.TrimSpace(gjson.GetBytes(body, "method").String()) if method == "" { return "" } if name := strings.TrimSpace(gjson.GetBytes(body, "params.name").String()); name != "" && (method == "tools/call" || method == "prompts/get") { - if method == "tools/call" && name == mcpgateway.CallToolName { + if method != "tools/call" { + return name + } + if name == mcpgateway.CallToolName { if inner := strings.TrimSpace(gjson.GetBytes(body, "params.arguments.name").String()); inner != "" { - return inner + name = inner } } + if resolveTool != nil { + name = resolveTool(name) + } return name } return method diff --git a/internal/server/mcp_service_test.go b/internal/server/mcp_service_test.go index 15dca2222..7e95d7f8c 100644 --- a/internal/server/mcp_service_test.go +++ b/internal/server/mcp_service_test.go @@ -74,12 +74,47 @@ func TestMCPAuditLabel(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got := mcpAuditLabel([]byte(tt.body)) + got := mcpAuditLabel([]byte(tt.body), nil) require.Equal(t, tt.want, got, "mcpAuditLabel(%s) = %q, want %q", tt.body, got, tt.want) }) } } +func TestMCPAuditLabelResolvesToolNames(t *testing.T) { + resolve := func(name string) string { + if name == "echo" { + return "alpha_echo" + } + return name + } + tests := []struct { + name string + body string + want string + }{ + { + name: "bare tools/call name", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"echo"}}`, + want: "alpha_echo", + }, + { + name: "bare call_tool target", + body: `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"echo"}}}`, + want: "alpha_echo", + }, + { + name: "prompts are not tool names", + body: `{"jsonrpc":"2.0","id":3,"method":"prompts/get","params":{"name":"echo"}}`, + want: "echo", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, mcpAuditLabel([]byte(tt.body), resolve)) + }) + } +} + // rejectingRateLimiter breaches every acquisition with a requests-window rule. type rejectingRateLimiter struct{} @@ -382,7 +417,7 @@ func TestEnrichMCPAuditEntryRestoresBody(t *testing.T) { body := `{"jsonrpc":"2.0","id":1,"method":"tools/list"}` c, _ := echotest.Post(t, "/mcp", body) - enrichMCPAuditEntry(c, false) + enrichMCPAuditEntry(c, false, nil) restored, err := io.ReadAll(c.Request().Body) require.NoError(t, err) @@ -394,7 +429,7 @@ func TestEnrichMCPAuditEntryCapturesRequestBody(t *testing.T) { entry := &auditlog.LogEntry{} c, _ := echotest.Post(t, "/mcp", body, echotest.WithValue(string(auditlog.LogEntryKey), entry)) - enrichMCPAuditEntry(c, true) + enrichMCPAuditEntry(c, true, nil) require.NotNil(t, entry.Data) require.NotNil(t, entry.Data.RequestBody) @@ -410,7 +445,7 @@ func TestEnrichMCPAuditEntryBodyLoggingOff(t *testing.T) { entry := &auditlog.LogEntry{} c, _ := echotest.Post(t, "/mcp", `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, echotest.WithValue(string(auditlog.LogEntryKey), entry)) - enrichMCPAuditEntry(c, false) + enrichMCPAuditEntry(c, false, nil) if entry.Data != nil { require.Nil(t, entry.Data.RequestBody, "body logging is off") From 2e20d78fc5a9fe12197f50144aa1ddf97207319e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Fri, 2 Oct 2026 14:38:58 -0700 Subject: [PATCH 3/3] fix(mcp): resolve call_tool log labels only in discovery sessions --- internal/mcpgateway/discovery.go | 10 ++++----- internal/mcpgateway/discovery_test.go | 14 ++++++------- internal/mcpgateway/service.go | 30 ++++++++++++++++++--------- internal/mcpgateway/service_test.go | 24 +++++++++++++-------- internal/server/mcp_service.go | 28 +++++++++---------------- internal/server/mcp_service_test.go | 20 +++++++++++------- 6 files changed, 68 insertions(+), 58 deletions(-) diff --git a/internal/mcpgateway/discovery.go b/internal/mcpgateway/discovery.go index c1220312b..dba685156 100644 --- a/internal/mcpgateway/discovery.go +++ b/internal/mcpgateway/discovery.go @@ -23,11 +23,9 @@ import ( const ToolDiscoveryHeader = "X-MCP-Tool-Discovery" // Meta-tool names served instead of the catalog in search discovery mode. -// CallToolName is exported so the request log can label a relayed call with -// the tool it runs. const ( searchToolsName = "search_tools" - CallToolName = "call_tool" + callToolName = "call_tool" ) const ( @@ -266,7 +264,7 @@ func (s *Service) registerDiscoveryTools(server *mcp.Server, idx *toolIndex, end }, s.searchToolsHandler(idx)) server.AddTool(&mcp.Tool{ - Name: CallToolName, + Name: callToolName, Description: "Call a tool found with " + searchToolsName + ". Pass its exact name and arguments matching its input schema.", InputSchema: map[string]any{ "type": "object", @@ -298,7 +296,7 @@ func searchToolsDescription(idx *toolIndex) string { for i, server := range servers { parts[i] = fmt.Sprintf("%s (%d)", server, counts[server]) } - desc := "Search the tools available through this gateway by keyword. Returns matching tool names, descriptions, and input schemas; run one with " + CallToolName + "." + desc := "Search the tools available through this gateway by keyword. Returns matching tool names, descriptions, and input schemas; run one with " + callToolName + "." if len(parts) > 0 { desc += " Servers: " + strings.Join(parts, ", ") + "." } @@ -361,7 +359,7 @@ func (s *Service) callToolHandler(idx *toolIndex, endpoint string) mcp.ToolHandl Arguments json.RawMessage `json:"arguments"` } if err := unmarshalArguments(req, &args); err != nil { - return toolError("invalid " + CallToolName + " arguments: " + err.Error()), nil + return toolError("invalid " + callToolName + " arguments: " + err.Error()), nil } target, ok := idx.lookup(strings.TrimSpace(args.Name)) if !ok { diff --git a/internal/mcpgateway/discovery_test.go b/internal/mcpgateway/discovery_test.go index e9d32d901..b1abeac58 100644 --- a/internal/mcpgateway/discovery_test.go +++ b/internal/mcpgateway/discovery_test.go @@ -136,7 +136,7 @@ func TestSearchDiscoveryServesMetaToolsAndRelaysCalls(t *testing.T) { service.searchDiscovery = true session := connectClient(t, gatewayURL+"/mcp", map[string]string{"X-Request-ID": "req-1"}) - assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, session)) + assert.Equal(t, []string{callToolName, searchToolsName}, listToolNames(t, session)) assert.Contains(t, session.InitializeResult().Instructions, searchToolsName) results := searchTools(t, session, "echo") @@ -147,7 +147,7 @@ func TestSearchDiscoveryServesMetaToolsAndRelaysCalls(t *testing.T) { for _, name := range []string{"alpha_echo", "echo"} { result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ - Name: CallToolName, + Name: callToolName, Arguments: map[string]any{"name": name, "arguments": map[string]any{"value": 1}}, }) require.NoError(t, err) @@ -177,7 +177,7 @@ func TestSearchDiscoveryRejectsUnknownToolAsToolError(t *testing.T) { session := connectClient(t, gatewayURL+"/mcp", nil) result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ - Name: CallToolName, + Name: callToolName, Arguments: map[string]any{"name": "alpha_missing"}, }) require.NoError(t, err) @@ -196,7 +196,7 @@ func TestToolDiscoveryHeaderOverridesDefault(t *testing.T) { service, gatewayURL := newTestService(t, nil, testSpec("alpha", url, nil)) optIn := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "Search"}) - assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, optIn)) + assert.Equal(t, []string{callToolName, searchToolsName}, listToolNames(t, optIn)) unknown := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "semantic"}) assert.Equal(t, []string{"alpha_echo"}, listToolNames(t, unknown)) @@ -206,7 +206,7 @@ func TestToolDiscoveryHeaderOverridesDefault(t *testing.T) { assert.Equal(t, []string{"alpha_echo"}, listToolNames(t, optOut)) pinned := connectClient(t, gatewayURL+"/mcp/alpha", nil) - assert.Equal(t, []string{CallToolName, searchToolsName}, listToolNames(t, pinned)) + assert.Equal(t, []string{callToolName, searchToolsName}, listToolNames(t, pinned)) results := searchTools(t, pinned, "echo") require.NotEmpty(t, results) assert.Equal(t, "echo", results[0].Name, "a pinned endpoint keeps original names") @@ -232,7 +232,7 @@ func TestSearchDiscoveryHidesToolsExcludedAfterInitialize(t *testing.T) { assert.Equal(t, "alpha_read", results[0].Name) result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ - Name: CallToolName, + Name: callToolName, Arguments: map[string]any{"name": "alpha_write"}, }) require.NoError(t, err, "failures reach the model as tool errors, not JSON-RPC errors") @@ -253,7 +253,7 @@ func TestSearchDiscoveryReportsUpstreamFailureAsToolError(t *testing.T) { service.manager.Apply(nil) // the upstream disappears after initialize result, err := session.CallTool(context.Background(), &mcp.CallToolParams{ - Name: CallToolName, + Name: callToolName, Arguments: map[string]any{"name": "alpha_echo"}, }) require.NoError(t, err) diff --git a/internal/mcpgateway/service.go b/internal/mcpgateway/service.go index e78654e3b..ec350929b 100644 --- a/internal/mcpgateway/service.go +++ b/internal/mcpgateway/service.go @@ -79,8 +79,10 @@ type sessionBinding struct { userPath string pinned string lastSeen time.Time - // toolAliases maps the session's unambiguous bare tool names to their - // namespaced names, so request logs can name the tool a call resolved to. + // discovery and toolAliases let request logs name the tool a call runs: + // a discovery session's call_tool target, and the namespaced name an + // unambiguous bare name resolves to. + discovery bool toolAliases map[string]string } @@ -374,7 +376,7 @@ func (s *Service) getServer(r *http.Request) *mcp.Server { }, GetSessionID: func() string { id := rand.Text() - s.bindSession(id, scope.authKeyID, scope.userPath, scope.pinned, aliases) + s.bindSession(id, scope.authKeyID, scope.userPath, scope.pinned, scope.discovery, aliases) return id }, }) @@ -519,7 +521,7 @@ func (s *Service) composeInstructions(scope requestScope, views []ServerView) st if b.Len() > 0 { b.WriteString(" ") } - fmt.Fprintf(&b, "Tools are not listed directly: find them with %s and run them with %s.", searchToolsName, CallToolName) + fmt.Fprintf(&b, "Tools are not listed directly: find them with %s and run them with %s.", searchToolsName, callToolName) } for _, view := range views { snapshot, _ := s.upstreamCatalog(view.Spec.Name) @@ -676,25 +678,33 @@ func (s *Service) authorizeSessionID(sessionID, upstreamName string) error { } // bindSession records the principal a new session was initialized under. -func (s *Service) bindSession(sessionID, authKeyID, userPath, pinned string, toolAliases map[string]string) { +func (s *Service) bindSession(sessionID, authKeyID, userPath, pinned string, discovery bool, toolAliases map[string]string) { s.bindMu.Lock() s.bindings[sessionID] = sessionBinding{ authKeyID: authKeyID, userPath: userPath, pinned: pinned, lastSeen: time.Now(), + discovery: discovery, toolAliases: toolAliases, } s.bindMu.Unlock() } -// CanonicalToolName resolves a bare tool name the session accepts on the -// aggregated endpoint to the namespaced name it runs, matching usage -// entries. Any other name is returned unchanged. -func (s *Service) CanonicalToolName(sessionID, name string) string { +// ToolCallLabel names the tool a session's tools/call runs, matching its +// usage entry. name is the called tool and target the call's +// arguments.name. In a discovery session call_tool resolves to its target, +// and a bare name the session accepts resolves to its namespaced name. Only +// the session knows whether call_tool is the meta-tool: on a pinned endpoint +// without discovery it can be an upstream tool of that name. +func (s *Service) ToolCallLabel(sessionID, name, target string) string { s.bindMu.Lock() defer s.bindMu.Unlock() - if exposed, ok := s.bindings[sessionID].toolAliases[name]; ok { + binding := s.bindings[sessionID] + if binding.discovery && name == callToolName && target != "" { + name = target + } + if exposed, ok := binding.toolAliases[name]; ok { return exposed } return name diff --git a/internal/mcpgateway/service_test.go b/internal/mcpgateway/service_test.go index 1d68af9f4..63ed458cd 100644 --- a/internal/mcpgateway/service_test.go +++ b/internal/mcpgateway/service_test.go @@ -300,7 +300,7 @@ func TestAggregatedEndpointAcceptsUniqueBareToolName(t *testing.T) { require.Equal(t, "alpha", entries[0].ProviderName) } -func TestCanonicalToolNameResolvesSessionAliases(t *testing.T) { +func TestToolCallLabelResolvesSessionAliases(t *testing.T) { alphaURL := newTestUpstream(t, "alpha", addEchoTool("echo")) betaURL := newTestUpstream(t, "beta", func(server *mcp.Server) { addEchoTool("echo")(server) @@ -314,14 +314,20 @@ func TestCanonicalToolNameResolvesSessionAliases(t *testing.T) { ) aggregated := connectClient(t, gatewayURL+"/mcp", nil).ID() - assert.Equal(t, "beta_search", service.CanonicalToolName(aggregated, "search")) - assert.Equal(t, "gamma_fetch", service.CanonicalToolName(aggregated, "fetch")) - assert.Equal(t, "echo", service.CanonicalToolName(aggregated, "echo"), "ambiguous names stay as sent") - assert.Equal(t, "alpha_echo", service.CanonicalToolName(aggregated, "alpha_echo")) + assert.Equal(t, "beta_search", service.ToolCallLabel(aggregated, "search", "")) + assert.Equal(t, "gamma_fetch", service.ToolCallLabel(aggregated, "fetch", "")) + assert.Equal(t, "echo", service.ToolCallLabel(aggregated, "echo", ""), "ambiguous names stay as sent") + assert.Equal(t, "alpha_echo", service.ToolCallLabel(aggregated, "alpha_echo", "")) + + discovery := connectClient(t, gatewayURL+"/mcp", map[string]string{ToolDiscoveryHeader: "search"}).ID() + assert.Equal(t, "gamma_fetch", service.ToolCallLabel(discovery, callToolName, "fetch"), "call_tool resolves to its bare target") + assert.Equal(t, "alpha_echo", service.ToolCallLabel(discovery, callToolName, "alpha_echo")) + assert.Equal(t, callToolName, service.ToolCallLabel(discovery, callToolName, "")) pinned := connectClient(t, gatewayURL+"/mcp/beta", nil).ID() - assert.Equal(t, "search", service.CanonicalToolName(pinned, "search"), "pinned endpoints use original names") - assert.Equal(t, "search", service.CanonicalToolName("unknown-session", "search")) + assert.Equal(t, "search", service.ToolCallLabel(pinned, "search", ""), "pinned endpoints use original names") + assert.Equal(t, callToolName, service.ToolCallLabel(pinned, callToolName, "search"), "without discovery call_tool is an upstream tool name") + assert.Equal(t, "search", service.ToolCallLabel("unknown-session", "search", "")) } func TestAggregatedEndpointRejectsAmbiguousBareToolName(t *testing.T) { @@ -551,9 +557,9 @@ func TestAuthorizeSessionFailsClosedWithoutBinding(t *testing.T) { err := service.authorizeSessionID("deleted-session", "alpha") require.ErrorIs(t, err, ErrServerNotVisible) - service.bindSession("live", "", "/staff", "", nil) + service.bindSession("live", "", "/staff", "", false, nil) require.NoError(t, service.authorizeSessionID("live", "alpha")) - service.bindSession("contractor", "", "/contractors/acme", "", nil) + service.bindSession("contractor", "", "/contractors/acme", "", false, nil) require.ErrorIs(t, service.authorizeSessionID("contractor", "alpha"), ErrServerNotVisible) } diff --git a/internal/server/mcp_service.go b/internal/server/mcp_service.go index 410cb7566..ce30abb66 100644 --- a/internal/server/mcp_service.go +++ b/internal/server/mcp_service.go @@ -42,8 +42,8 @@ func (s *mcpService) handle(c *echo.Context, pinnedServer string) error { // notification stream) and DELETE (session teardown) stay free. if c.Request().Method == http.MethodPost { sessionID := strings.TrimSpace(c.Request().Header.Get("Mcp-Session-Id")) - enrichMCPAuditEntry(c, s.logBodies, func(name string) string { - return s.gateway.CanonicalToolName(sessionID, name) + enrichMCPAuditEntry(c, s.logBodies, func(name, target string) string { + return s.gateway.ToolCallLabel(sessionID, name, target) }) release, err := enforceRateLimit(c, s.rateLimiter, rateLimitRoute{}) if err != nil { @@ -77,7 +77,7 @@ func (s *mcpService) handle(c *echo.Context, pinnedServer string) error { // the JSON-RPC frame as the request body when body logging is on. The body is // restored for the gateway handler; the body-limit middleware has already // bounded its size. -func enrichMCPAuditEntry(c *echo.Context, logBodies bool, resolveTool func(string) string) { +func enrichMCPAuditEntry(c *echo.Context, logBodies bool, resolveTool func(name, target string) string) { req := c.Request() if req.Body == nil { return @@ -98,29 +98,21 @@ func enrichMCPAuditEntry(c *echo.Context, logBodies bool, resolveTool func(strin } // mcpAuditLabel derives the request-log label from one JSON-RPC frame: the -// tool/prompt name for calls, otherwise the method. A search-discovery -// call_tool is labelled with the tool it runs, and resolveTool (optional) -// maps a bare tool name to the namespaced one, both matching usage entries. -// Empty means unlabelable (a bare response or malformed frame). -func mcpAuditLabel(body []byte, resolveTool func(string) string) string { +// tool/prompt name for calls, otherwise the method. resolveTool (optional) +// receives a tools/call name and its arguments.name and returns the tool the +// session actually runs, matching usage entries. Empty means unlabelable (a +// bare response or malformed frame). +func mcpAuditLabel(body []byte, resolveTool func(name, target string) string) string { method := strings.TrimSpace(gjson.GetBytes(body, "method").String()) if method == "" { return "" } if name := strings.TrimSpace(gjson.GetBytes(body, "params.name").String()); name != "" && (method == "tools/call" || method == "prompts/get") { - if method != "tools/call" { + if method != "tools/call" || resolveTool == nil { return name } - if name == mcpgateway.CallToolName { - if inner := strings.TrimSpace(gjson.GetBytes(body, "params.arguments.name").String()); inner != "" { - name = inner - } - } - if resolveTool != nil { - name = resolveTool(name) - } - return name + return resolveTool(name, strings.TrimSpace(gjson.GetBytes(body, "params.arguments.name").String())) } return method } diff --git a/internal/server/mcp_service_test.go b/internal/server/mcp_service_test.go index 7e95d7f8c..af83d9c63 100644 --- a/internal/server/mcp_service_test.go +++ b/internal/server/mcp_service_test.go @@ -32,13 +32,8 @@ func TestMCPAuditLabel(t *testing.T) { want: "github_create_issue", }, { - name: "discovery call_tool labels with the tool it runs", - body: `{"jsonrpc":"2.0","id":5,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"github_create_issue","arguments":{}}}}`, - want: "github_create_issue", - }, - { - name: "call_tool without a target keeps its own name", - body: `{"jsonrpc":"2.0","id":6,"method":"tools/call","params":{"name":"call_tool","arguments":{}}}`, + name: "call_tool keeps its own name without a session resolver", + body: `{"jsonrpc":"2.0","id":5,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"github_create_issue"}}}`, want: "call_tool", }, { @@ -81,7 +76,11 @@ func TestMCPAuditLabel(t *testing.T) { } func TestMCPAuditLabelResolvesToolNames(t *testing.T) { - resolve := func(name string) string { + // Mimics a discovery session on the aggregated endpoint. + resolve := func(name, target string) string { + if name == "call_tool" && target != "" { + name = target + } if name == "echo" { return "alpha_echo" } @@ -102,6 +101,11 @@ func TestMCPAuditLabelResolvesToolNames(t *testing.T) { body: `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"echo"}}}`, want: "alpha_echo", }, + { + name: "namespaced call_tool target", + body: `{"jsonrpc":"2.0","id":4,"method":"tools/call","params":{"name":"call_tool","arguments":{"name":"github_create_issue"}}}`, + want: "github_create_issue", + }, { name: "prompts are not tool names", body: `{"jsonrpc":"2.0","id":3,"method":"prompts/get","params":{"name":"echo"}}`,