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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 17 additions & 13 deletions aibridge/intercept/messages/blocking.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package messages

import (
"context"
"encoding/base64"
"errors"
"fmt"
"net/http"
Expand All @@ -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"
Expand Down Expand Up @@ -257,42 +258,45 @@ 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,
},
})
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,
},
})
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:
Expand Down
30 changes: 17 additions & 13 deletions aibridge/intercept/messages/streaming.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package messages
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
Expand All @@ -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"
Expand Down Expand Up @@ -410,41 +411,44 @@ 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,
},
})
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:
Expand Down
39 changes: 23 additions & 16 deletions aibridge/internal/integrationtest/mockmcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
}
8 changes: 5 additions & 3 deletions aibridge/internal/testutil/mockserverproxier.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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
}
2 changes: 1 addition & 1 deletion aibridge/mcp/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
6 changes: 3 additions & 3 deletions aibridge/mcp/client_info.go
Original file line number Diff line number Diff line change
@@ -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(),
}
Expand Down
40 changes: 25 additions & 15 deletions aibridge/mcp/mcp_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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.
Expand All @@ -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)
Expand All @@ -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 changes: 33 additions & 0 deletions aibridge/mcp/mcphttpclient.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading