diff --git a/catalog/catalog_test.go b/catalog/catalog_test.go index 5698102..98df1a4 100644 --- a/catalog/catalog_test.go +++ b/catalog/catalog_test.go @@ -46,7 +46,7 @@ func TestCompileExhaustsPaginationSortsAndRoutesOriginalNames(t *testing.T) { if !ok { t.Fatalf("RouteTool(%q) not found", name) } - if route.Component.Name != "local files" || route.OriginalName != name { + if route.Component.Name != "local files" || route.Tool.Name != name { t.Fatalf("RouteTool(%q) = %+v", name, route) } } diff --git a/catalog/compiler.go b/catalog/compiler.go index 3c11385..ffb14b1 100644 --- a/catalog/compiler.go +++ b/catalog/compiler.go @@ -164,7 +164,7 @@ func compileTools(c *Catalog, server config.Server, prefix string, discovered [] clone.Description = override.OverrideDescription } c.tools = append(c.tools, &clone) - c.toolRoutes[name] = ToolRoute{Component: server, Prefix: prefix, OriginalName: tool.Name} + c.toolRoutes[name] = ToolRoute{Component: server, Prefix: prefix, Tool: tool} } return nil } diff --git a/catalog/compiler_test.go b/catalog/compiler_test.go index b092b08..0bf2fc8 100644 --- a/catalog/compiler_test.go +++ b/catalog/compiler_test.go @@ -209,7 +209,7 @@ func TestCompileAllFeaturesAppliesOverridesBeforeNamespace(t *testing.T) { if got := compiled.ResourceTemplates(); len(got) != 1 || got[0].URITemplate != "mmmcp+fancy_server:file:///public/{path}" || got[0].Name != "public files" || got[0].Description != "overridden template" { t.Fatalf("resource templates = %+v", got) } - if route, ok := compiled.RouteTool("fancy_server__find"); !ok || route.OriginalName != "search" { + if route, ok := compiled.RouteTool("fancy_server__find"); !ok || route.Tool.Name != "search" { t.Fatalf("tool route = %+v, %v", route, ok) } if route, ok := compiled.RoutePrompt("fancy_server__describe"); !ok || route.OriginalName != "explain" { @@ -248,7 +248,7 @@ func TestCompileSingleServerPreservesFeatureIdentities(t *testing.T) { if got := compiled.ResourceTemplates(); len(got) != 1 || got[0].URITemplate != "file:///{path}" { t.Fatalf("resource templates = %+v", got) } - if route, ok := compiled.RouteTool("search"); !ok || route.OriginalName != "search" || route.Prefix != "" { + if route, ok := compiled.RouteTool("search"); !ok || route.Tool.Name != "search" || route.Prefix != "" { t.Fatalf("tool route = %+v, %v", route, ok) } if route, ok := compiled.RoutePrompt("explain"); !ok || route.OriginalName != "explain" || route.Prefix != "" { diff --git a/catalog/routes.go b/catalog/routes.go index 23959f0..2ec86cf 100644 --- a/catalog/routes.go +++ b/catalog/routes.go @@ -1,15 +1,16 @@ package catalog import ( + "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/obot-platform/mmmcp/config" "github.com/yosida95/uritemplate/v3" ) // ToolRoute maps an exposed tool identity back to its component identity. type ToolRoute struct { - Component config.Server - Prefix string - OriginalName string + Component config.Server + Prefix string + Tool *mcp.Tool } // PromptRoute maps an exposed prompt identity back to its component identity. diff --git a/component/component.go b/component/component.go index e1d404e..7c6528f 100644 --- a/component/component.go +++ b/component/component.go @@ -78,6 +78,12 @@ type Callbacks struct { type requestHeadersContextKey struct{} type clientInfoContextKey struct{} +type toolCallContextKey struct{} + +type toolCallContext struct { + tool *mcp.Tool + arguments []byte +} type valueSuppressingContext struct { context.Context @@ -121,6 +127,18 @@ func RequestHeadersFromContext(ctx context.Context) http.Header { return headers.Clone() } +// ContextWithToolCall attaches the routed tool definition and raw arguments +// needed to generate transport-level parameter headers downstream. +func ContextWithToolCall(ctx context.Context, tool *mcp.Tool, arguments []byte) context.Context { + return context.WithValue(ctx, toolCallContextKey{}, toolCallContext{tool: tool, arguments: slices.Clone(arguments)}) +} + +// ToolCallFromContext returns the routed tool definition and raw arguments. +func ToolCallFromContext(ctx context.Context) (*mcp.Tool, []byte) { + call, _ := ctx.Value(toolCallContextKey{}).(toolCallContext) + return call.tool, slices.Clone(call.arguments) +} + // ContextWithClientInfo attaches an immutable snapshot of the frontend client identity for downstream connections. func ContextWithClientInfo(ctx context.Context, info *mcp.Implementation) context.Context { if info == nil { @@ -149,7 +167,7 @@ func cloneClientInfo(info mcp.Implementation) mcp.Implementation { func (c valueSuppressingContext) Value(key any) any { switch key.(type) { - case requestHeadersContextKey, clientInfoContextKey: + case requestHeadersContextKey, clientInfoContextKey, toolCallContextKey: return c.Context.Value(key) } return nil diff --git a/component/http/client.go b/component/http/client.go index f43ad4d..0d69ba2 100644 --- a/component/http/client.go +++ b/component/http/client.go @@ -3,9 +3,14 @@ package http import ( "context" + "encoding/base64" + "encoding/json" "errors" "fmt" + "math" "net/http" + "strconv" + "strings" "time" "github.com/modelcontextprotocol/go-sdk/auth" @@ -40,6 +45,11 @@ type headerTransport struct { passthroughHeaders []string } +type headerSchemaProperty struct { + Header json.RawMessage `json:"x-mcp-header"` + Properties map[string]headerSchemaProperty `json:"properties"` +} + // NewFactory creates a Streamable HTTP component factory. func NewFactory(opts FactoryOptions) *Factory { return &Factory{HTTPClient: opts.HTTPClient, OAuth: opts.OAuth, ClientInfo: opts.ClientInfo} @@ -320,5 +330,76 @@ func (t headerTransport) RoundTrip(req *http.Request) (*http.Response, error) { for name, value := range t.headers { clone.Header.Set(name, value) } + tool, arguments := component.ToolCallFromContext(req.Context()) + addToolParamHeaders(clone.Header, tool, arguments) return t.base.RoundTrip(clone) } + +func addToolParamHeaders(headers http.Header, tool *mcp.Tool, arguments []byte) { + if tool == nil || len(arguments) == 0 { + return + } + data, err := json.Marshal(tool.InputSchema) + if err != nil { + return + } + var schema struct { + Properties map[string]headerSchemaProperty `json:"properties"` + } + var args map[string]json.RawMessage + if json.Unmarshal(data, &schema) != nil || json.Unmarshal(arguments, &args) != nil { + return + } + addAnnotatedHeaders(headers, schema.Properties, args) +} + +func addAnnotatedHeaders(headers http.Header, properties map[string]headerSchemaProperty, arguments map[string]json.RawMessage) { + for name, property := range properties { + argument, ok := arguments[name] + if !ok || string(argument) == "null" { + continue + } + var header string + if json.Unmarshal(property.Header, &header) == nil && header != "" { + if value, ok := encodeParamHeader(argument); ok { + headers.Set("Mcp-Param-"+header, value) + } + } + if len(property.Properties) > 0 { + var nested map[string]json.RawMessage + if json.Unmarshal(argument, &nested) == nil { + addAnnotatedHeaders(headers, property.Properties, nested) + } + } + } +} + +func encodeParamHeader(raw json.RawMessage) (string, bool) { + var value any + if json.Unmarshal(raw, &value) != nil { + return "", false + } + var encoded string + switch value := value.(type) { + case string: + encoded = value + case bool: + encoded = strconv.FormatBool(value) + case float64: + if value != math.Trunc(value) || value < -(1<<53-1) || value > 1<<53-1 { + return "", false + } + encoded = strconv.FormatInt(int64(value), 10) + default: + return "", false + } + if strings.Trim(encoded, " \t") != encoded || (strings.HasPrefix(encoded, "=?base64?") && strings.HasSuffix(encoded, "?=")) { + return "=?base64?" + base64.StdEncoding.EncodeToString([]byte(encoded)) + "?=", true + } + for _, char := range encoded { + if char < 0x20 || char > 0x7e { + return "=?base64?" + base64.StdEncoding.EncodeToString([]byte(encoded)) + "?=", true + } + } + return encoded, true +} diff --git a/handler.go b/handler.go index 62b49a0..e768510 100644 --- a/handler.go +++ b/handler.go @@ -123,9 +123,10 @@ func (c *Composite) callTool(ctx context.Context, request mcp.Request, compiled if !ok { return nil, unknown("tool", req.Params.Name) } + ctx = component.ContextWithToolCall(ctx, route.Tool, req.Params.Arguments) params := &mcp.CallToolParams{ Meta: component.DownstreamMeta(req.Params.Meta), - Name: route.OriginalName, + Name: route.Tool.Name, Arguments: req.Params.Arguments, InputResponses: req.Params.InputResponses, RequestState: req.Params.RequestState, diff --git a/http_param_headers_test.go b/http_param_headers_test.go new file mode 100644 index 0000000..4066b4e --- /dev/null +++ b/http_param_headers_test.go @@ -0,0 +1,103 @@ +package mmmcp_test + +import ( + "context" + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/obot-platform/mmmcp" + "github.com/obot-platform/mmmcp/config" + "github.com/obot-platform/mmmcp/testserver" +) + +type toolCallHeaderRecorder struct { + base http.RoundTripper + mu sync.Mutex + headers http.Header +} + +func (r *toolCallHeaderRecorder) RoundTrip(request *http.Request) (*http.Response, error) { + if request.Method == http.MethodPost { + r.mu.Lock() + r.headers = request.Header.Clone() + r.mu.Unlock() + } + return r.base.RoundTrip(request) +} + +func (r *toolCallHeaderRecorder) Headers() http.Header { + r.mu.Lock() + defer r.mu.Unlock() + return r.headers.Clone() +} + +func TestCompositeGeneratesToolParameterHeaders(t *testing.T) { + for _, test := range []struct { + name string + protocol string + stateful bool + }{ + {name: "stateless", protocol: "2026-07-28"}, + {name: "stateful", protocol: "2025-11-25", stateful: true}, + } { + t.Run(test.name, func(t *testing.T) { + recorder := &toolCallHeaderRecorder{base: http.DefaultTransport} + fixture := testserver.New(t, testserver.Options{Stateful: test.stateful, Tools: []testserver.Tool{{ + Definition: &mcp.Tool{Name: "repository", InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "owner": map[string]any{"type": "string", "x-mcp-header": "owner"}, + "enabled": map[string]any{"type": "boolean", "x-mcp-header": "enabled"}, + "count": map[string]any{"type": "integer", "x-mcp-header": "count"}, + "scope": map[string]any{"type": "object", "properties": map[string]any{ + "region": map[string]any{"type": "string", "x-mcp-header": "region"}, + }}, + }, + }}, + Handler: func(context.Context, *mcp.CallToolRequest) (*mcp.CallToolResult, error) { + return &mcp.CallToolResult{}, nil + }, + }}}) + composite, err := mmmcp.New(t.Context(), &config.Config{Servers: []config.Server{{Name: "fixture", URL: fixture.URL}}}, mmmcp.Options{ + HTTPClient: &http.Client{Transport: recorder}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = composite.Close() }) + frontend := httptest.NewServer(composite.HTTPHandler()) + t.Cleanup(frontend.Close) + client := mcp.NewClient(&mcp.Implementation{Name: "test-client", Version: "1"}, nil) + session, err := client.Connect(t.Context(), &mcp.StreamableClientTransport{ + Endpoint: frontend.URL, + HTTPClient: frontend.Client(), + DisableStandaloneSSE: true, + }, &mcp.ClientSessionOptions{ProtocolVersion: test.protocol}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = session.Close() }) + + _, err = session.CallTool(t.Context(), &mcp.CallToolParams{Name: "repository", Arguments: map[string]any{ + "owner": " octocat ", "enabled": true, "count": 42, "scope": map[string]any{"region": "日本"}, + }}) + if err != nil { + t.Fatal(err) + } + headers := recorder.Headers() + for name, want := range map[string]string{ + "Mcp-Param-owner": "=?base64?IG9jdG9jYXQg?=", + "Mcp-Param-enabled": "true", + "Mcp-Param-count": "42", + "Mcp-Param-region": "=?base64?5pel5pys?=", + } { + if got := headers.Get(name); got != want { + t.Errorf("%s = %q, want %q", name, got, want) + } + } + }) + } +}