Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion catalog/catalog_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
Expand Down
2 changes: 1 addition & 1 deletion catalog/compiler.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
4 changes: 2 additions & 2 deletions catalog/compiler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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" {
Expand Down Expand Up @@ -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 != "" {
Expand Down
7 changes: 4 additions & 3 deletions catalog/routes.go
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
20 changes: 19 additions & 1 deletion component/component.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand Down
81 changes: 81 additions & 0 deletions component/http/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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}
Expand Down Expand Up @@ -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
}
3 changes: 2 additions & 1 deletion handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
103 changes: 103 additions & 0 deletions http_param_headers_test.go
Original file line number Diff line number Diff line change
@@ -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{
Comment thread
thedadams marked this conversation as resolved.
"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)
}
}
})
}
}
Loading