mirror of
https://github.com/coder/coder.git
synced 2026-09-01 14:53:15 +08:00
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:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
@@ -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})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,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
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user