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()