feat: migrate aibridge injected-MCP proxy to official MCP Go SDK (#28060)

## Stack Context

PR 5 of 6 in a stack that migrates every Coder MCP surface from the
archived `github.com/mark3labs/mcp-go` library to the official
`github.com/modelcontextprotocol/go-sdk` v1.7.0.

Stack: #28056 -> #28057 -> #28058 -> #28059 -> #28060 -> #28061

## Why

The aibridge injected-MCP proxy now owns an official `*mcp.Client`,
`*mcp.StreamableClientTransport`, and `*mcp.ClientSession`.

- The proxy constructor accepts an optional `*http.Client` instead of
mark3labs options; the header-injecting wrapper shallow-copies a
supplied client so its Timeout, Jar, and redirect policy survive.
- Manual protocol version negotiation and the mark3labs five-second
close workaround are removed; the SDK negotiates during `Connect` and
fails when no mutually supported version exists.
- Repeated `Init` closes the previous session, and a failed tool fetch
closes the just-created session so transports do not leak.
- Tool and intercept types use the official pointer content types;
embedded resource blobs are re-encoded to base64 for model-facing text
because the SDK decodes them into raw bytes.
- `aibridge/mcpmock` is regenerated, and its stale `go:generate` source
path is corrected.

> Mux created this PR on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-08-13 12:38:14 +02:00
committed by GitHub
parent 7720e283f5
commit c8e8b21a88
15 changed files with 224 additions and 159 deletions
+17 -13
View File
@@ -2,6 +2,7 @@ package messages
import (
"context"
"encoding/base64"
"errors"
"fmt"
"net/http"
@@ -10,7 +11,7 @@ import (
"github.com/anthropics/anthropic-sdk-go"
"github.com/anthropics/anthropic-sdk-go/option"
"github.com/google/uuid"
mcplib "github.com/mark3labs/mcp-go/mcp"
mcplib "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/tidwall/sjson"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
@@ -257,7 +258,7 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
var hasValidResult bool
for _, content := range res.Content {
switch cb := content.(type) {
case mcplib.TextContent:
case *mcplib.TextContent:
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: cb.Text,
@@ -265,20 +266,23 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
})
hasValidResult = true
// TODO: is there a more correct way of handling these non-text content responses?
case mcplib.EmbeddedResource:
switch resource := cb.Resource.(type) {
case mcplib.TextResourceContents:
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Text)
case *mcplib.EmbeddedResource:
resource := cb.Resource
switch {
case resource == nil:
i.logger.Warn(ctx, "embedded resource with no contents")
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: val,
Text: "Error: embedded resource with no contents",
},
})
toolResult.OfToolResult.IsError = anthropic.Bool(true)
hasValidResult = true
case mcplib.BlobResourceContents:
case resource.Blob != nil:
// The SDK decodes base64 during unmarshal; re-encode
// the bytes for model-facing text.
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Blob)
resource.MIMEType, resource.URI, base64.StdEncoding.EncodeToString(resource.Blob))
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: val,
@@ -286,13 +290,13 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
})
hasValidResult = true
default:
i.logger.Warn(ctx, "unknown embedded resource type", slog.F("type", fmt.Sprintf("%T", resource)))
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Text)
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: "Error: unknown embedded resource type",
Text: val,
},
})
toolResult.OfToolResult.IsError = anthropic.Bool(true)
hasValidResult = true
}
default:
+17 -13
View File
@@ -3,6 +3,7 @@ package messages
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
@@ -14,7 +15,7 @@ import (
"github.com/anthropics/anthropic-sdk-go/packages/ssestream"
"github.com/anthropics/anthropic-sdk-go/shared/constant"
"github.com/google/uuid"
mcplib "github.com/mark3labs/mcp-go/mcp"
mcplib "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/tidwall/sjson"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
@@ -410,27 +411,30 @@ newStream:
var hasValidResult bool
for _, content := range res.Content {
switch cb := content.(type) {
case mcplib.TextContent:
case *mcplib.TextContent:
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: cb.Text,
},
})
hasValidResult = true
case mcplib.EmbeddedResource:
switch resource := cb.Resource.(type) {
case mcplib.TextResourceContents:
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Text)
case *mcplib.EmbeddedResource:
resource := cb.Resource
switch {
case resource == nil:
logger.Warn(ctx, "embedded resource with no contents")
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: val,
Text: "Error: embedded resource with no contents",
},
})
toolResult.OfToolResult.IsError = anthropic.Bool(true)
hasValidResult = true
case mcplib.BlobResourceContents:
case resource.Blob != nil:
// The SDK decodes base64 during unmarshal; re-encode
// the bytes for model-facing text.
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Blob)
resource.MIMEType, resource.URI, base64.StdEncoding.EncodeToString(resource.Blob))
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: val,
@@ -438,13 +442,13 @@ newStream:
})
hasValidResult = true
default:
logger.Warn(ctx, "unknown embedded resource type", slog.F("type", fmt.Sprintf("%T", resource)))
val := fmt.Sprintf("Binary resource (MIME: %s, URI: %s): %s",
resource.MIMEType, resource.URI, resource.Text)
toolResult.OfToolResult.Content = append(toolResult.OfToolResult.Content, anthropic.ToolResultBlockParamContentUnion{
OfText: &anthropic.TextBlockParam{
Text: "Error: unknown embedded resource type",
Text: val,
},
})
toolResult.OfToolResult.IsError = anthropic.Bool(true)
hasValidResult = true
}
default:
+23 -16
View File
@@ -2,15 +2,14 @@ package integrationtest
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/mark3labs/mcp-go/client/transport"
mcplib "github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/trace"
"go.opentelemetry.io/otel/trace/noop"
@@ -63,7 +62,7 @@ func setupMCPForTestWithName(t *testing.T, name string, tracer trace.Tracer) *mo
httpTransport := &http.Transport{}
t.Cleanup(httpTransport.CloseIdleConnections)
httpClient := &http.Client{Transport: httpTransport}
proxy, err := mcp.NewStreamableHTTPServerProxy(name, mcpSrv.URL, nil, nil, nil, logger, tracer, transport.WithHTTPBasicClient(httpClient))
proxy, err := mcp.NewStreamableHTTPServerProxy(name, mcpSrv.URL, nil, nil, nil, logger, tracer, httpClient)
require.NoError(t, err)
mgr := mcp.NewServerProxyManager(map[string]mcp.ServerProxier{proxy.Name(): proxy}, tracer)
@@ -129,26 +128,34 @@ func (a *callAccumulator) getCallsByTool(name string) []any {
func createMockMCPSrv(t *testing.T) (http.Handler, *callAccumulator) {
t.Helper()
s := server.NewMCPServer(
"Mock coder MCP server",
"1.0.0",
server.WithToolCapabilities(true),
)
s := sdkmcp.NewServer(&sdkmcp.Implementation{
Name: "Mock coder MCP server",
Version: "1.0.0",
}, nil)
acc := newCallAccumulator()
for _, name := range []string{mockToolName, "coder_list_templates", "coder_template_version_parameters", "coder_get_authenticated_user", "coder_create_workspace_build", "coder_delete_template"} {
tool := mcplib.NewTool(name,
mcplib.WithDescription(fmt.Sprintf("Mock of the %s tool", name)),
)
s.AddTool(tool, func(_ context.Context, request mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
acc.addCall(request.Params.Name, request.Params.Arguments)
s.AddTool(&sdkmcp.Tool{
Name: name,
Description: fmt.Sprintf("Mock of the %s tool", name),
InputSchema: map[string]any{"type": "object"},
}, func(_ context.Context, request *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) {
var args any
if len(request.Params.Arguments) > 0 {
_ = json.Unmarshal(request.Params.Arguments, &args)
}
acc.addCall(request.Params.Name, args)
if errMsg, ok := acc.getToolError(request.Params.Name); ok {
return nil, xerrors.New(errMsg)
}
return mcplib.NewToolResultText("mock"), nil
return &sdkmcp.CallToolResult{
Content: []sdkmcp.Content{&sdkmcp.TextContent{Text: "mock"}},
}, nil
})
}
return server.NewStreamableHTTPServer(s), acc
// Stateless mode gives each POST an ephemeral server session, so
// no server-side goroutines outlive the request (goleak-clean).
return sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server { return s }, &sdkmcp.StreamableHTTPOptions{Stateless: true}), acc
}
@@ -3,7 +3,7 @@ package testutil
import (
"context"
mcpgo "github.com/mark3labs/mcp-go/mcp"
mcpgo "github.com/modelcontextprotocol/go-sdk/mcp"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge/mcp"
@@ -59,6 +59,8 @@ func (*MockServerProxier) CallTool(context.Context, string, any) (*mcpgo.CallToo
// StubToolCaller is a minimal tool client that returns a fixed text result.
type StubToolCaller struct{}
func (StubToolCaller) CallTool(_ context.Context, _ mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) {
return mcpgo.NewToolResultText("tool result"), nil
func (StubToolCaller) CallTool(_ context.Context, _ *mcpgo.CallToolParams) (*mcpgo.CallToolResult, error) {
return &mcpgo.CallToolResult{
Content: []mcpgo.Content{&mcpgo.TextContent{Text: "tool result"}},
}, nil
}
+1 -1
View File
@@ -3,7 +3,7 @@ package mcp
import (
"context"
"github.com/mark3labs/mcp-go/mcp"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// ServerProxier provides an abstraction to communicate with MCP Servers regardless of their transport.
+3 -3
View File
@@ -1,15 +1,15 @@
package mcp
import (
"github.com/mark3labs/mcp-go/mcp"
"github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/coder/coder/v2/buildinfo"
)
// GetClientInfo returns the MCP client information to use when initializing MCP connections.
// This provides a consistent way for all proxy implementations to report client information.
func GetClientInfo() mcp.Implementation {
return mcp.Implementation{
func GetClientInfo() *mcp.Implementation {
return &mcp.Implementation{
Name: "coder/aibridge",
Version: buildinfo.Version(),
}
+25 -15
View File
@@ -10,8 +10,7 @@ import (
"strings"
"testing"
mcplib "github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"go.uber.org/goleak"
@@ -308,9 +307,9 @@ func TestToolInjectionOrder(t *testing.T) {
tracer := otel.Tracer("forTesting")
// When: creating two MCP server proxies, both listing the same tools by name but under different server namespaces.
proxy, err := mcp.NewStreamableHTTPServerProxy("coder", mcpSrv.URL, nil, nil, nil, logger, tracer)
proxy, err := mcp.NewStreamableHTTPServerProxy("coder", mcpSrv.URL, nil, nil, nil, logger, tracer, nil)
require.NoError(t, err)
proxy2, err := mcp.NewStreamableHTTPServerProxy("shmoder", mcpSrv.URL, nil, nil, nil, logger, tracer)
proxy2, err := mcp.NewStreamableHTTPServerProxy("shmoder", mcpSrv.URL, nil, nil, nil, logger, tracer, nil)
require.NoError(t, err)
// Then: initialize both proxies.
@@ -327,6 +326,13 @@ func TestToolInjectionOrder(t *testing.T) {
"shmoder": proxy2,
}, otel.GetTracerProvider().Tracer("test"))
require.NoError(t, mgr.Init(ctx))
// Close the sessions before the httptest server's own cleanup,
// which blocks until all client connections are gone.
t.Cleanup(func() {
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), testutil.WaitShort)
defer shutdownCancel()
require.NoError(t, mgr.Shutdown(shutdownCtx))
})
// Then: the tools from both servers should be collectively sorted stably.
validateToolOrder(t, mgr)
@@ -352,20 +358,24 @@ func validateToolOrder(t *testing.T, proxy mcp.ServerProxier) {
func createMockMCPSrv(t *testing.T) http.Handler {
t.Helper()
s := server.NewMCPServer(
"Mock coder MCP server",
"1.0.0",
server.WithToolCapabilities(true),
)
s := sdkmcp.NewServer(&sdkmcp.Implementation{
Name: "Mock coder MCP server",
Version: "1.0.0",
}, nil)
for _, name := range []string{"coder_list_workspaces", "coder_list_templates", "coder_template_version_parameters", "coder_get_authenticated_user"} {
tool := mcplib.NewTool(name,
mcplib.WithDescription(fmt.Sprintf("Mock of the %s tool", name)),
)
s.AddTool(tool, func(ctx context.Context, request mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
return mcplib.NewToolResultText("mock"), nil
s.AddTool(&sdkmcp.Tool{
Name: name,
Description: fmt.Sprintf("Mock of the %s tool", name),
InputSchema: map[string]any{"type": "object"},
}, func(ctx context.Context, request *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) {
return &sdkmcp.CallToolResult{
Content: []sdkmcp.Content{&sdkmcp.TextContent{Text: "mock"}},
}, nil
})
}
return server.NewStreamableHTTPServer(s)
// Stateless mode gives each POST an ephemeral server session, so
// no server-side goroutines outlive the request (goleak-clean).
return sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server { return s }, &sdkmcp.StreamableHTTPOptions{Stateless: true})
}
+33
View File
@@ -5,6 +5,39 @@ import (
"net/http"
)
// withHeaders shallow-copies base so client-level settings such as
// Timeout and Jar survive the transport wrap.
func withHeaders(base *http.Client, headers map[string]string) *http.Client {
client := &http.Client{}
if base != nil {
clone := *base
client = &clone
}
if client.Transport == nil {
client.Transport = http.DefaultTransport
}
if len(headers) > 0 {
client.Transport = &headerRoundTripper{
base: client.Transport,
headers: headers,
}
}
return client
}
type headerRoundTripper struct {
base http.RoundTripper
headers map[string]string
}
func (h *headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
clone := req.Clone(req.Context())
for k, v := range h.headers {
clone.Header.Set(k, v)
}
return h.base.RoundTrip(clone)
}
// mcpHTTPClient returns an isolated *http.Client when running
// inside tests, or nil for production. During tests,
// httptest.Server.Close() calls
+72 -55
View File
@@ -2,13 +2,12 @@ package mcp
import (
"context"
"net/http"
"regexp"
"slices"
"strings"
"github.com/mark3labs/mcp-go/client"
"github.com/mark3labs/mcp-go/client/transport"
"github.com/mark3labs/mcp-go/mcp"
"github.com/modelcontextprotocol/go-sdk/mcp"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
"golang.org/x/exp/maps"
@@ -21,9 +20,11 @@ import (
var _ ServerProxier = &StreamableHTTPServerProxy{}
type StreamableHTTPServerProxy struct {
client *client.Client
logger slog.Logger
tracer trace.Tracer
client *mcp.Client
tr *mcp.StreamableClientTransport
session *mcp.ClientSession
logger slog.Logger
tracer trace.Tracer
allowlistPattern *regexp.Regexp
denylistPattern *regexp.Regexp
@@ -33,32 +34,23 @@ type StreamableHTTPServerProxy struct {
tools map[string]*Tool
}
func NewStreamableHTTPServerProxy(serverName, serverURL string, headers map[string]string, allowlist, denylist *regexp.Regexp, logger slog.Logger, tracer trace.Tracer, opts ...transport.StreamableHTTPCOption) (*StreamableHTTPServerProxy, error) {
// nit: headers should be passed in as an option instead of a separate parameter. Not changed as this would be a breaking change.
if headers != nil {
opts = append(opts, transport.WithHTTPHeaders(headers))
func NewStreamableHTTPServerProxy(serverName, serverURL string, headers map[string]string, allowlist, denylist *regexp.Regexp, logger slog.Logger, tracer trace.Tracer, httpClient *http.Client) (*StreamableHTTPServerProxy, error) {
if httpClient == nil {
httpClient = mcpHTTPClient()
}
httpClient = withHeaders(httpClient, headers)
// Prepend an isolated HTTP client when running in tests so
// httptest.Server.Close() does not disrupt this proxy's
// connections via http.DefaultTransport.CloseIdleConnections().
// Caller-provided WithHTTPBasicClient in opts overrides this
// (last-wins).
if c := mcpHTTPClient(); c != nil {
opts = append([]transport.StreamableHTTPCOption{
transport.WithHTTPBasicClient(c),
}, opts...)
}
mcpClient, err := client.NewStreamableHttpClient(serverURL, opts...)
if err != nil {
return nil, xerrors.Errorf("create streamable http client: %w", err)
tr := &mcp.StreamableClientTransport{
Endpoint: serverURL,
HTTPClient: httpClient,
}
mcpClient := mcp.NewClient(GetClientInfo(), nil)
return &StreamableHTTPServerProxy{
serverName: serverName,
serverURL: serverURL,
client: mcpClient,
tr: tr,
logger: logger,
tracer: tracer,
allowlistPattern: allowlist,
@@ -74,34 +66,32 @@ func (p *StreamableHTTPServerProxy) Init(ctx context.Context) (outErr error) {
ctx, span := p.tracer.Start(ctx, "StreamableHTTPServerProxy.Init", trace.WithAttributes(p.traceAttributes()...))
defer tracing.EndSpanErr(span, &outErr)
if err := p.client.Start(ctx); err != nil {
return xerrors.Errorf("start client: %w", err)
// Init may be called again (e.g. via ServerProxyManager); close
// the previous session so its transport does not leak.
if p.session != nil {
if err := p.session.Close(); err != nil {
p.logger.Debug(ctx, "failed to close previous MCP session", slog.Error(err))
}
p.session = nil
}
version := mcp.LATEST_PROTOCOL_VERSION
initReq := mcp.InitializeRequest{
Params: mcp.InitializeParams{
ProtocolVersion: version,
ClientInfo: GetClientInfo(),
},
}
result, err := p.client.Initialize(ctx, initReq)
// The SDK negotiates the protocol version during Connect and
// fails when no mutually supported version exists.
session, err := p.client.Connect(ctx, p.tr, nil)
if err != nil {
return xerrors.Errorf("init MCP client: %w", err)
}
p.session = session
if !slices.Contains(mcp.ValidProtocolVersions, result.ProtocolVersion) {
if err := p.client.Close(); err != nil {
p.logger.Debug(ctx, "failed to close MCP client on unsuccessful version negotiation", slog.Error(err))
}
return xerrors.Errorf("MCP version negotiation failed; requested %q, accepts %q, received %q", version, strings.Join(mcp.ValidProtocolVersions, ","), result.ProtocolVersion)
}
result := session.InitializeResult()
p.logger.Debug(ctx, "mcp client initialized", slog.F("name", result.ServerInfo.Name), slog.F("server_version", result.ServerInfo.Version))
tools, err := p.fetchTools(ctx)
if err != nil {
if closeErr := session.Close(); closeErr != nil {
p.logger.Debug(ctx, "failed to close MCP session after fetch tools error", slog.Error(closeErr))
}
p.session = nil
return xerrors.Errorf("fetch tools: %w", err)
}
@@ -136,11 +126,13 @@ func (p *StreamableHTTPServerProxy) CallTool(ctx context.Context, name string, i
return nil, xerrors.Errorf("%q tool not known", name)
}
return p.client.CallTool(ctx, mcp.CallToolRequest{
Params: mcp.CallToolParams{
Name: tool.Name,
Arguments: input,
},
if p.session == nil {
return nil, xerrors.New("proxy not initialized")
}
return p.session.CallTool(ctx, &mcp.CallToolParams{
Name: tool.Name,
Arguments: input,
})
}
@@ -148,7 +140,7 @@ func (p *StreamableHTTPServerProxy) fetchTools(ctx context.Context) (_ map[strin
ctx, span := p.tracer.Start(ctx, "StreamableHTTPServerProxy.Init.fetchTools", trace.WithAttributes(p.traceAttributes()...))
defer tracing.EndSpanErr(span, &outErr)
tools, err := p.client.ListTools(ctx, mcp.ListToolsRequest{})
tools, err := p.session.ListTools(ctx, nil)
if err != nil {
return nil, xerrors.Errorf("list MCP tools: %w", err)
}
@@ -166,14 +158,14 @@ func (p *StreamableHTTPServerProxy) fetchTools(ctx context.Context) (_ map[strin
)
}
out[encodedID] = &Tool{
Client: p.client,
Client: p.session,
ID: encodedID,
Name: tool.Name,
ServerName: p.serverName,
ServerURL: p.serverURL,
Description: tool.Description,
Params: tool.InputSchema.Properties,
Required: tool.InputSchema.Required,
Params: toolParams(tool.InputSchema),
Required: toolRequired(tool.InputSchema),
Logger: p.logger,
}
}
@@ -182,13 +174,11 @@ func (p *StreamableHTTPServerProxy) fetchTools(ctx context.Context) (_ map[strin
}
func (p *StreamableHTTPServerProxy) Shutdown(_ context.Context) error {
if p.client == nil {
if p.session == nil {
return nil
}
// NOTE: as of v0.38.0 the lib doesn't allow an outside context to be passed in;
// it has an internal timeout of 5s, though.
return p.client.Close()
return p.session.Close()
}
func (p *StreamableHTTPServerProxy) traceAttributes() []attribute.KeyValue {
@@ -198,3 +188,30 @@ func (p *StreamableHTTPServerProxy) traceAttributes() []attribute.KeyValue {
attribute.String(tracing.MCPServerURL, p.serverURL),
}
}
func toolParams(schema any) map[string]any {
m, ok := schema.(map[string]any)
if !ok {
return nil
}
properties, _ := m["properties"].(map[string]any)
return properties
}
func toolRequired(schema any) []string {
m, ok := schema.(map[string]any)
if !ok {
return nil
}
rawRequired, ok := m["required"].([]any)
if !ok {
return nil
}
var required []string
for _, r := range rawRequired {
if str, ok := r.(string); ok {
required = append(required, str)
}
}
return required
}
+1 -1
View File
@@ -6,7 +6,7 @@ import (
"strings"
"sync"
"github.com/mark3labs/mcp-go/mcp"
"github.com/modelcontextprotocol/go-sdk/mcp"
"go.opentelemetry.io/otel/trace"
"golang.org/x/xerrors"
+7 -10
View File
@@ -7,7 +7,7 @@ import (
"strings"
"time"
"github.com/mark3labs/mcp-go/mcp"
"github.com/modelcontextprotocol/go-sdk/mcp"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
"golang.org/x/xerrors"
@@ -42,11 +42,10 @@ func SanitizeToolName(name string) string {
return toolNameSanitizer.ReplaceAllString(name, "_")
}
// ToolCaller is the narrowest interface which describes the behavior required from [mcp.Client],
// which will normally be passed into [Tool] for interaction with an MCP server.
// TODO: don't expose github.com/mark3labs/mcp-go outside this package.
// ToolCaller is the subset of [mcp.ClientSession] used by [Tool].
// TODO: avoid exposing MCP SDK types from this package.
type ToolCaller interface {
CallTool(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error)
CallTool(ctx context.Context, params *mcp.CallToolParams) (*mcp.CallToolResult, error)
}
type Tool struct {
@@ -92,11 +91,9 @@ func (t *Tool) Call(ctx context.Context, input any, tracer trace.Tracer) (_ *mcp
start := time.Now()
var res *mcp.CallToolResult
res, outErr = t.Client.CallTool(ctx, mcp.CallToolRequest{
Params: mcp.CallToolParams{
Name: t.Name,
Arguments: input,
},
res, outErr = t.Client.CallTool(ctx, &mcp.CallToolParams{
Name: t.Name,
Arguments: input,
})
logFn := t.Logger.Debug
+1 -1
View File
@@ -1,3 +1,3 @@
package mcpmock
//go:generate go tool mockgen -destination ./mcpmock.go -package mcpmock github.com/coder/aibridge/mcp ServerProxier
//go:generate go tool mockgen -destination ./mcpmock.go -package mcpmock github.com/coder/coder/v2/aibridge/mcp ServerProxier
+3 -3
View File
@@ -1,9 +1,9 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: github.com/coder/aibridge/mcp (interfaces: ServerProxier)
// Source: github.com/coder/coder/v2/aibridge/mcp (interfaces: ServerProxier)
//
// Generated by this command:
//
// mockgen -destination ./mcpmock.go -package mcpmock github.com/coder/aibridge/mcp ServerProxier
// mockgen -destination ./mcpmock.go -package mcpmock github.com/coder/coder/v2/aibridge/mcp ServerProxier
//
// Package mcpmock is a generated GoMock package.
@@ -14,7 +14,7 @@ import (
reflect "reflect"
mcp "github.com/coder/coder/v2/aibridge/mcp"
mcp0 "github.com/mark3labs/mcp-go/mcp"
mcp0 "github.com/modelcontextprotocol/go-sdk/mcp"
gomock "go.uber.org/mock/gomock"
)
+1
View File
@@ -196,6 +196,7 @@ func (m *MCPProxyFactory) newStreamableHTTPServerProxy(cfg *proto.MCPServerConfi
denylist,
m.logger.Named(fmt.Sprintf("mcp-server-proxy-%s", cfg.GetId())),
m.tracer,
nil,
)
if err != nil {
return nil, xerrors.Errorf("create streamable HTTP MCP server proxy: %w", err)
+15 -25
View File
@@ -7,9 +7,11 @@ import (
"net/http"
"net/http/httptest"
"slices"
"sync/atomic"
"testing"
"time"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/prometheus/client_golang/prometheus"
promtest "github.com/prometheus/client_golang/prometheus/testutil"
"github.com/stretchr/testify/require"
@@ -67,33 +69,19 @@ func TestIntegration(t *testing.T) {
tracer := tp.Tracer(t.Name())
defer func() { _ = tp.Shutdown(t.Context()) }()
// Create mock MCP server.
var mcpTokenReceived string
var mcpTokenReceived atomic.Pointer[string]
mcpHandler := sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server {
return sdkmcp.NewServer(&sdkmcp.Implementation{
Name: "test-mcp-server",
Version: "1.0.0",
}, nil)
}, &sdkmcp.StreamableHTTPOptions{Stateless: true})
mockMCPServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
t.Logf("Mock MCP server received request: %s %s", r.Method, r.URL.Path)
if r.Method == http.MethodPost && r.URL.Path == "/" {
// Mark that init was called.
mcpTokenReceived = r.Header.Get("Authorization")
t.Log("MCP init request received")
// Return a basic MCP init response.
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Mcp-Session-Id", "test-session-123")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{
"jsonrpc": "2.0",
"id": 1,
"result": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"serverInfo": {
"name": "test-mcp-server",
"version": "1.0.0"
}
}
}`))
if auth := r.Header.Get("Authorization"); auth != "" {
mcpTokenReceived.Store(&auth)
}
mcpHandler.ServeHTTP(w, r)
}))
t.Cleanup(mockMCPServer.Close)
t.Logf("Mock MCP server running at: %s", mockMCPServer.URL)
@@ -292,7 +280,9 @@ func TestIntegration(t *testing.T) {
require.False(t, tools[0].Injected)
// Then: the MCP server was initialized.
require.Contains(t, mcpTokenReceived, authLink.OAuthAccessToken, "mock MCP server not requested")
gotMCPToken := mcpTokenReceived.Load()
require.NotNil(t, gotMCPToken, "mock MCP server not requested")
require.Contains(t, *gotMCPToken, authLink.OAuthAccessToken)
// Then: verify tracing spans were recorded.
spans := sr.Ended()