From 587cce04ae54214e5a1d34003ea340c9151ca289 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Wed, 12 Aug 2026 09:16:10 +0000 Subject: [PATCH] feat(aibridge): migrate injected-MCP proxy to official MCP Go SDK Replace the mark3labs client with the official SDK client for the deprecated injected-MCP proxy. The manual protocol-version handshake check and the 5s-close workaround are subsumed by SDK negotiation and session close. The proxy constructor takes an optional *http.Client instead of mark3labs transport options; auth headers ride an http.RoundTripper. Repeated Init now closes the previous session instead of leaking its transport. Test fixtures move to official SDK stateless servers, including the enterprise integration mock, which previously answered every POST with a canned initialize response that the stricter SDK client rejects. The mcpmock generate directive also pointed at the pre-vendoring aibridge module path; fix it and regenerate. --- aibridge/intercept/messages/blocking.go | 30 +++-- aibridge/intercept/messages/streaming.go | 30 +++-- aibridge/internal/integrationtest/mockmcp.go | 39 +++--- .../internal/testutil/mockserverproxier.go | 8 +- aibridge/mcp/api.go | 2 +- aibridge/mcp/client_info.go | 6 +- aibridge/mcp/mcp_test.go | 40 +++--- aibridge/mcp/mcphttpclient.go | 33 +++++ aibridge/mcp/proxy_streamable_http.go | 127 ++++++++++-------- aibridge/mcp/server_proxy_manager.go | 2 +- aibridge/mcp/tool.go | 17 +-- aibridge/mcpmock/doc.go | 2 +- aibridge/mcpmock/mcpmock.go | 6 +- coderd/aibridged/mcp.go | 1 + enterprise/aibridged_integration_test.go | 40 +++--- 15 files changed, 224 insertions(+), 159 deletions(-) diff --git a/aibridge/intercept/messages/blocking.go b/aibridge/intercept/messages/blocking.go index cba8a30f0cf..9bcffb40973 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -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: diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index b4b9f69bbd3..8322de557f2 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -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: diff --git a/aibridge/internal/integrationtest/mockmcp.go b/aibridge/internal/integrationtest/mockmcp.go index ffbd4fad19d..e93fc207a65 100644 --- a/aibridge/internal/integrationtest/mockmcp.go +++ b/aibridge/internal/integrationtest/mockmcp.go @@ -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 } diff --git a/aibridge/internal/testutil/mockserverproxier.go b/aibridge/internal/testutil/mockserverproxier.go index b962e825e74..6514c61220a 100644 --- a/aibridge/internal/testutil/mockserverproxier.go +++ b/aibridge/internal/testutil/mockserverproxier.go @@ -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 } diff --git a/aibridge/mcp/api.go b/aibridge/mcp/api.go index 1abd476a8cf..676beb089ca 100644 --- a/aibridge/mcp/api.go +++ b/aibridge/mcp/api.go @@ -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. diff --git a/aibridge/mcp/client_info.go b/aibridge/mcp/client_info.go index 04a4973a3e5..9a69b662653 100644 --- a/aibridge/mcp/client_info.go +++ b/aibridge/mcp/client_info.go @@ -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(), } diff --git a/aibridge/mcp/mcp_test.go b/aibridge/mcp/mcp_test.go index aeea86e72d2..b88b36287d1 100644 --- a/aibridge/mcp/mcp_test.go +++ b/aibridge/mcp/mcp_test.go @@ -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}) } diff --git a/aibridge/mcp/mcphttpclient.go b/aibridge/mcp/mcphttpclient.go index bc70a7f5abc..20270114e3d 100644 --- a/aibridge/mcp/mcphttpclient.go +++ b/aibridge/mcp/mcphttpclient.go @@ -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 diff --git a/aibridge/mcp/proxy_streamable_http.go b/aibridge/mcp/proxy_streamable_http.go index 8d9e3583c18..aa4d6765e1c 100644 --- a/aibridge/mcp/proxy_streamable_http.go +++ b/aibridge/mcp/proxy_streamable_http.go @@ -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) - } - - version := mcp.LATEST_PROTOCOL_VERSION - initReq := mcp.InitializeRequest{ - Params: mcp.InitializeParams{ - ProtocolVersion: version, - ClientInfo: GetClientInfo(), - }, + // 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 } - 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 +} diff --git a/aibridge/mcp/server_proxy_manager.go b/aibridge/mcp/server_proxy_manager.go index 9c9bdb12320..d2a86612a69 100644 --- a/aibridge/mcp/server_proxy_manager.go +++ b/aibridge/mcp/server_proxy_manager.go @@ -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" diff --git a/aibridge/mcp/tool.go b/aibridge/mcp/tool.go index bb13d626ef4..924407185a6 100644 --- a/aibridge/mcp/tool.go +++ b/aibridge/mcp/tool.go @@ -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 diff --git a/aibridge/mcpmock/doc.go b/aibridge/mcpmock/doc.go index 6b16ed44591..25c1838ca68 100644 --- a/aibridge/mcpmock/doc.go +++ b/aibridge/mcpmock/doc.go @@ -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 diff --git a/aibridge/mcpmock/mcpmock.go b/aibridge/mcpmock/mcpmock.go index 2678c733529..0804ca5a959 100644 --- a/aibridge/mcpmock/mcpmock.go +++ b/aibridge/mcpmock/mcpmock.go @@ -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" ) diff --git a/coderd/aibridged/mcp.go b/coderd/aibridged/mcp.go index ef6c8a26286..f5ad759d22d 100644 --- a/coderd/aibridged/mcp.go +++ b/coderd/aibridged/mcp.go @@ -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) diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 03f71996e7e..a5ae02f994d 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -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()