chore: move chatd and related packages to /x/ subpackage (#23445)

- Moves `coderd/chatd/`, `coderd/gitsync/`, `enterprise/coderd/chatd/`
under `x/` parent directories to signal instability
- Adds `Experimental:` glue code comments in `coderd/coderd.go`

> 🤖 This PR was created with the help of Coder Agents, and was
reviewed by my human. 🧑‍💻
This commit is contained in:
Cian Johnston
2026-03-23 17:34:43 +00:00
committed by GitHub
parent 86d8b6daee
commit 80a172f932
64 changed files with 92 additions and 90 deletions
+542
View File
@@ -0,0 +1,542 @@
package mcpclient
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/url"
"strings"
"sync"
"time"
"charm.land/fantasy"
"github.com/google/uuid"
"github.com/mark3labs/mcp-go/client"
"github.com/mark3labs/mcp-go/client/transport"
"github.com/mark3labs/mcp-go/mcp"
"golang.org/x/sync/errgroup"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/buildinfo"
"github.com/coder/coder/v2/coderd/database"
)
// toolNameSep separates the server slug from the original tool
// name in prefixed tool names. Double underscore avoids collisions
// with tool names that may contain single underscores.
//
// TODO: tool names that themselves contain "__" produce ambiguous
// prefixed names (e.g. "srv__my__tool" is indistinguishable from
// slug "srv" + tool "my__tool" vs slug "srv__my" + tool "tool").
// This doesn't affect tool invocation since originalName is used
// directly when calling the remote server.
const toolNameSep = "__"
// connectTimeout bounds how long we wait for a single MCP server
// to start its transport and complete initialization. Servers that
// take longer are skipped so one slow server cannot block the
// entire chat startup.
const connectTimeout = 10 * time.Second
// toolCallTimeout bounds how long a single tool invocation may
// take before being canceled.
const toolCallTimeout = 60 * time.Second
// ConnectAll connects to all configured MCP servers, discovers
// their tools, and returns them as fantasy.AgentTool values. It
// skips servers that fail to connect and logs warnings. The
// returned cleanup function must be called to close all
// connections.
func ConnectAll(
ctx context.Context,
logger slog.Logger,
configs []database.MCPServerConfig,
tokens []database.MCPServerUserToken,
) ([]fantasy.AgentTool, func()) {
// Index tokens by server config ID so auth header
// construction is O(1) per server.
tokensByConfigID := make(
map[uuid.UUID]database.MCPServerUserToken, len(tokens),
)
for _, tok := range tokens {
tokensByConfigID[tok.MCPServerConfigID] = tok
}
var (
mu sync.Mutex
clients []*client.Client
tools []fantasy.AgentTool
)
// Build cleanup eagerly so it always closes any clients
// that connected, even if a later connection fails.
cleanup := func() {
mu.Lock()
defer mu.Unlock()
for _, c := range clients {
_ = c.Close()
}
clients = nil
}
var eg errgroup.Group
for _, cfg := range configs {
if !cfg.Enabled {
continue
}
eg.Go(func() error {
serverTools, mcpClient, connectErr := connectOne(
ctx, logger, cfg, tokensByConfigID,
)
if connectErr != nil {
logger.Warn(ctx,
"skipping MCP server due to connection failure",
slog.F("server_slug", cfg.Slug),
slog.F("server_url", RedactURL(cfg.Url)),
slog.F("error", redactErrorURL(connectErr)),
)
// Connection failures are not propagated — the
// LLM simply won't have this server's tools.
return nil
}
mu.Lock()
clients = append(clients, mcpClient)
tools = append(tools, serverTools...)
mu.Unlock()
return nil
})
}
// All goroutines return nil; error is intentionally
// discarded.
_ = eg.Wait()
return tools, cleanup
}
// connectOne establishes a connection to a single MCP server,
// discovers its tools, and wraps each one as an AgentTool with
// the server slug prefix applied.
func connectOne(
ctx context.Context,
logger slog.Logger,
cfg database.MCPServerConfig,
tokensByConfigID map[uuid.UUID]database.MCPServerUserToken,
) ([]fantasy.AgentTool, *client.Client, error) {
headers := buildAuthHeaders(ctx, logger, cfg, tokensByConfigID)
tr, err := createTransport(cfg, headers)
if err != nil {
return nil, nil, xerrors.Errorf(
"create transport: %w", err,
)
}
mcpClient := client.NewClient(tr)
// The timeout covers the entire connect+init+list sequence,
// not each phase individually.
connectCtx, cancel := context.WithTimeout(
ctx, connectTimeout,
)
defer cancel()
if err := mcpClient.Start(connectCtx); err != nil {
_ = mcpClient.Close()
return nil, nil, xerrors.Errorf(
"start transport: %w", err,
)
}
_, err = mcpClient.Initialize(
connectCtx,
mcp.InitializeRequest{
Params: mcp.InitializeParams{
ProtocolVersion: mcp.LATEST_PROTOCOL_VERSION,
ClientInfo: mcp.Implementation{
Name: "coder",
Version: buildinfo.Version(),
},
},
},
)
if err != nil {
// Best-effort close so we don't leak the transport.
_ = mcpClient.Close()
return nil, nil, xerrors.Errorf("initialize: %w", err)
}
toolsResult, err := mcpClient.ListTools(
connectCtx, mcp.ListToolsRequest{},
)
if err != nil {
_ = mcpClient.Close()
return nil, nil, xerrors.Errorf("list tools: %w", err)
}
var tools []fantasy.AgentTool
for _, mcpTool := range toolsResult.Tools {
if !isToolAllowed(
mcpTool.Name,
cfg.ToolAllowList,
cfg.ToolDenyList,
) {
logger.Debug(ctx, "skipping denied MCP tool",
slog.F("server_slug", cfg.Slug),
slog.F("tool_name", mcpTool.Name),
)
continue
}
tools = append(
tools, newMCPTool(cfg.Slug, mcpTool, mcpClient),
)
}
// If no tools passed filtering, close the client early
// to avoid holding an idle connection.
if len(tools) == 0 {
_ = mcpClient.Close()
return nil, nil, nil
}
return tools, mcpClient, nil
}
// createTransport builds the appropriate mcp-go transport based
// on the server's configured transport type.
func createTransport(
cfg database.MCPServerConfig,
headers map[string]string,
) (transport.Interface, error) {
switch cfg.Transport {
case "sse":
return transport.NewSSE(
cfg.Url,
transport.WithHeaders(headers),
)
case "", "streamable_http":
// Default to streamable HTTP, the newer transport.
return transport.NewStreamableHTTP(
cfg.Url,
transport.WithHTTPHeaders(headers),
)
default:
return nil, xerrors.Errorf(
"unsupported transport %q", cfg.Transport,
)
}
}
// buildAuthHeaders constructs HTTP headers for authenticating
// with the MCP server based on the configured auth type.
func buildAuthHeaders(
ctx context.Context,
logger slog.Logger,
cfg database.MCPServerConfig,
tokensByConfigID map[uuid.UUID]database.MCPServerUserToken,
) map[string]string {
// Using map[string]string rather than http.Header because
// the mcp-go transport options accept map[string]string.
// MCP servers typically don't require multi-valued headers.
headers := make(map[string]string)
switch cfg.AuthType {
case "oauth2":
tok, ok := tokensByConfigID[cfg.ID]
if !ok {
logger.Warn(ctx,
"no oauth2 token found for MCP server",
slog.F("server_slug", cfg.Slug),
)
break
}
if tok.Expiry.Valid && tok.Expiry.Time.Before(time.Now()) {
logger.Warn(ctx,
"oauth2 token for MCP server is expired",
slog.F("server_slug", cfg.Slug),
slog.F("expired_at", tok.Expiry.Time),
)
}
if tok.AccessToken == "" {
logger.Warn(ctx,
"oauth2 token record has empty access token",
slog.F("server_slug", cfg.Slug),
)
break
}
tokenType := tok.TokenType
if tokenType == "" {
tokenType = "Bearer"
}
headers["Authorization"] = tokenType + " " + tok.AccessToken
case "api_key":
if cfg.APIKeyHeader != "" && cfg.APIKeyValue != "" {
headers[cfg.APIKeyHeader] = cfg.APIKeyValue
}
case "custom_headers":
if cfg.CustomHeaders != "" {
var custom map[string]string
if err := json.Unmarshal(
[]byte(cfg.CustomHeaders), &custom,
); err != nil {
logger.Warn(ctx,
"failed to parse custom headers JSON",
slog.F("server_slug", cfg.Slug),
slog.Error(err),
)
} else {
for k, v := range custom {
headers[k] = v
}
}
}
case "none", "":
// No auth headers needed.
}
return headers
}
// isToolAllowed checks a tool name against the allow and deny
// lists. When the allow list is non-empty only tools in it are
// permitted and the deny list is ignored. When the allow list
// is empty and the deny list is non-empty, tools in the deny
// list are rejected. Both lists use exact string matching
// against the original (non-prefixed) tool name.
func isToolAllowed(
toolName string,
allowList []string,
denyList []string,
) bool {
if len(allowList) > 0 {
for _, allowed := range allowList {
if allowed == toolName {
return true
}
}
// Allow list is set but the tool isn't in it.
return false
}
for _, denied := range denyList {
if denied == toolName {
return false
}
}
return true
}
// RedactURL strips userinfo and query parameters from a URL
// to avoid logging embedded credentials. Query params are
// removed because API keys are sometimes passed as
// ?api_key=sk-... in server URLs.
func RedactURL(rawURL string) string {
u, err := url.Parse(rawURL)
if err != nil {
return rawURL
}
u.User = nil
u.RawQuery = ""
u.Fragment = ""
return u.String()
}
// redactErrorURL rewrites URLs in an error string to strip
// credentials. Go's net/http embeds the full request URL in
// *url.Error messages, which can leak userinfo.
func redactErrorURL(err error) string {
if err == nil {
return ""
}
var urlErr *url.Error
if errors.As(err, &urlErr) {
urlErr.URL = RedactURL(urlErr.URL)
return urlErr.Error()
}
return err.Error()
}
// mcpToolWrapper adapts a single MCP tool into a
// fantasy.AgentTool. It stores the prefixed name for Info() but
// strips the prefix when forwarding calls to the remote server.
type mcpToolWrapper struct {
prefixedName string
originalName string
description string
parameters map[string]any
required []string
client *client.Client
providerOptions fantasy.ProviderOptions
}
// newMCPTool creates an mcpToolWrapper from an mcp.Tool
// discovered on a remote server.
func newMCPTool(
serverSlug string,
tool mcp.Tool,
mcpClient *client.Client,
) *mcpToolWrapper {
return &mcpToolWrapper{
prefixedName: serverSlug + toolNameSep + tool.Name,
originalName: tool.Name,
description: tool.Description,
parameters: tool.InputSchema.Properties,
required: tool.InputSchema.Required,
client: mcpClient,
}
}
func (t *mcpToolWrapper) Info() fantasy.ToolInfo {
return fantasy.ToolInfo{
Name: t.prefixedName,
Description: t.description,
Parameters: t.parameters,
Required: t.required,
Parallel: true,
}
}
func (t *mcpToolWrapper) Run(
ctx context.Context,
params fantasy.ToolCall,
) (fantasy.ToolResponse, error) {
var args map[string]any
if params.Input != "" {
if err := json.Unmarshal(
[]byte(params.Input), &args,
); err != nil {
return fantasy.NewTextErrorResponse(
"invalid JSON input: " + err.Error(),
), nil
}
}
callCtx, cancel := context.WithTimeout(ctx, toolCallTimeout)
defer cancel()
result, err := t.client.CallTool(
callCtx,
mcp.CallToolRequest{
Params: mcp.CallToolParams{
Name: t.originalName,
Arguments: args,
},
},
)
if err != nil {
return fantasy.NewTextErrorResponse(err.Error()), nil
}
return convertCallResult(result), nil
}
func (t *mcpToolWrapper) ProviderOptions() fantasy.ProviderOptions {
return t.providerOptions
}
func (t *mcpToolWrapper) SetProviderOptions(
opts fantasy.ProviderOptions,
) {
t.providerOptions = opts
}
// convertCallResult translates an MCP CallToolResult into a
// fantasy.ToolResponse. The fantasy response model supports a
// single content type per response, so we prioritize text. All
// text items are collected first. Binary items (image or audio)
// are only returned when no text content is available.
func convertCallResult(
result *mcp.CallToolResult,
) fantasy.ToolResponse {
if result == nil {
return fantasy.NewTextResponse("")
}
var (
textParts []string
binaryResult *fantasy.ToolResponse
)
for _, item := range result.Content {
switch c := item.(type) {
case mcp.TextContent:
textParts = append(textParts, c.Text)
case mcp.ImageContent:
data, err := base64.StdEncoding.DecodeString(
c.Data,
)
if err != nil {
textParts = append(textParts,
"[image decode error: "+err.Error()+"]",
)
continue
}
if binaryResult == nil {
r := fantasy.ToolResponse{
Type: "image",
Data: data,
MediaType: c.MIMEType,
IsError: result.IsError,
}
binaryResult = &r
}
case mcp.AudioContent:
data, err := base64.StdEncoding.DecodeString(
c.Data,
)
if err != nil {
textParts = append(textParts,
"[audio decode error: "+err.Error()+"]",
)
continue
}
if binaryResult == nil {
r := fantasy.ToolResponse{
Type: "media",
Data: data,
MediaType: c.MIMEType,
IsError: result.IsError,
}
binaryResult = &r
}
default:
textParts = append(textParts,
fmt.Sprintf("[unsupported content type: %T]", c),
)
}
}
// If structured content is present, marshal it to JSON and
// append as a text part so the data is preserved for the LLM.
if result.StructuredContent != nil {
data, err := json.Marshal(result.StructuredContent)
if err != nil {
textParts = append(textParts,
"[structured content marshal error: "+
err.Error()+"]",
)
} else {
textParts = append(textParts, string(data))
}
}
// Prefer text content. Only fall back to binary when no
// text was collected.
if len(textParts) > 0 {
resp := fantasy.NewTextResponse(
strings.Join(textParts, "\n"),
)
resp.IsError = result.IsError
return resp
}
if binaryResult != nil {
return *binaryResult
}
return fantasy.NewTextResponse("")
}
+659
View File
@@ -0,0 +1,659 @@
package mcpclient_test
import (
"context"
"database/sql"
"encoding/json"
"net/http/httptest"
"sync"
"testing"
"time"
"charm.land/fantasy"
"github.com/google/uuid"
"github.com/mark3labs/mcp-go/mcp"
mcpserver "github.com/mark3labs/mcp-go/server"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/x/chatd/mcpclient"
)
// newTestMCPServer creates a streamable HTTP MCP server with the
// given tools. The caller must close the returned *httptest.Server.
func newTestMCPServer(t *testing.T, tools ...mcpserver.ServerTool) *httptest.Server {
t.Helper()
srv := mcpserver.NewMCPServer("test-server", "1.0.0")
srv.AddTools(tools...)
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
return ts
}
// echoTool returns a ServerTool that echoes its "input" argument
// prefixed with "echo: ".
func echoTool() mcpserver.ServerTool {
return mcpserver.ServerTool{
Tool: mcp.NewTool("echo",
mcp.WithDescription("Echoes the input"),
mcp.WithString("input", mcp.Description("The input"), mcp.Required()),
),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
input, _ := req.GetArguments()["input"].(string)
return mcp.NewToolResultText("echo: " + input), nil
},
}
}
// greetTool returns a ServerTool that greets by name.
func greetTool() mcpserver.ServerTool {
return mcpserver.ServerTool{
Tool: mcp.NewTool("greet",
mcp.WithDescription("Greets the user"),
mcp.WithString("name", mcp.Description("Name to greet"), mcp.Required()),
),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
name, _ := req.GetArguments()["name"].(string)
return mcp.NewToolResultText("hello " + name), nil
},
}
}
// makeConfig builds a database.MCPServerConfig suitable for tests.
func makeConfig(slug, url string) database.MCPServerConfig {
return database.MCPServerConfig{
ID: uuid.New(),
Slug: slug,
DisplayName: slug,
Url: url,
Transport: "streamable_http",
AuthType: "none",
Enabled: true,
}
}
func TestConnectAll_DiscoverTools(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool(), greetTool())
cfg := makeConfig("myserver", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
// Two tools should be discovered, namespaced with the server slug.
require.Len(t, tools, 2)
names := toolNames(tools)
assert.Contains(t, names, "myserver__echo")
assert.Contains(t, names, "myserver__greet")
// Verify the description is preserved.
foundEcho := findTool(tools, "myserver__echo")
require.NotNilf(t, foundEcho, "expected to find myserver__echo")
echoInfo := foundEcho.Info()
assert.Equal(t, "Echoes the input", echoInfo.Description)
}
func TestConnectAll_CallTool(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
tool := tools[0]
resp, err := tool.Run(ctx, fantasy.ToolCall{
ID: "call-1",
Name: "srv__echo",
Input: `{"input":"hello world"}`,
})
require.NoError(t, err)
assert.False(t, resp.IsError)
assert.Equal(t, "echo: hello world", resp.Content)
}
func TestConnectAll_ToolAllowList(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool(), greetTool())
cfg := makeConfig("filtered", ts.URL)
// Only allow the "echo" tool.
cfg.ToolAllowList = []string{"echo"}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
assert.Equal(t, "filtered__echo", tools[0].Info().Name)
}
func TestConnectAll_ToolDenyList(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool(), greetTool())
cfg := makeConfig("filtered", ts.URL)
// Deny the "greet" tool, so only "echo" remains.
cfg.ToolDenyList = []string{"greet"}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
assert.Equal(t, "filtered__echo", tools[0].Info().Name)
}
func TestConnectAll_ConnectionFailure(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
cfg := makeConfig("bad", "http://127.0.0.1:0/does-not-exist")
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
assert.Empty(t, tools, "no tools should be returned for an unreachable server")
}
func TestConnectAll_MultipleServers(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts1 := newTestMCPServer(t, echoTool())
ts2 := newTestMCPServer(t, greetTool())
cfg1 := makeConfig("alpha", ts1.URL)
cfg2 := makeConfig("beta", ts2.URL)
tools, cleanup := mcpclient.ConnectAll(
ctx, logger,
[]database.MCPServerConfig{cfg1, cfg2},
nil,
)
t.Cleanup(cleanup)
require.Len(t, tools, 2)
names := toolNames(tools)
assert.Contains(t, names, "alpha__echo")
assert.Contains(t, names, "beta__greet")
}
func TestConnectAll_AuthHeaders(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
// Create a server whose tool handler records the Authorization
// header it receives on each request.
var (
mu sync.Mutex
seenHeaders []string
)
srv := mcpserver.NewMCPServer("auth-server", "1.0.0")
srv.AddTools(mcpserver.ServerTool{
Tool: mcp.NewTool("whoami",
mcp.WithDescription("Returns the auth header"),
),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
auth := req.Header.Get("Authorization")
mu.Lock()
seenHeaders = append(seenHeaders, auth)
mu.Unlock()
return mcp.NewToolResultText("auth:" + auth), nil
},
})
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
configID := uuid.New()
cfg := database.MCPServerConfig{
ID: configID,
Slug: "auth-srv",
DisplayName: "Auth Server",
Url: ts.URL,
Transport: "streamable_http",
AuthType: "oauth2",
Enabled: true,
}
token := database.MCPServerUserToken{
MCPServerConfigID: configID,
AccessToken: "test-token-abc",
TokenType: "Bearer",
}
tools, cleanup := mcpclient.ConnectAll(
ctx, logger,
[]database.MCPServerConfig{cfg},
[]database.MCPServerUserToken{token},
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
// Call the tool and verify the response includes the auth header
// that was sent.
resp, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-auth",
Name: "auth-srv__whoami",
Input: "{}",
})
require.NoError(t, err)
assert.False(t, resp.IsError)
assert.Equal(t, "auth:Bearer test-token-abc", resp.Content)
// Also verify the handler actually observed the header.
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, seenHeaders)
assert.Equal(t, "Bearer test-token-abc", seenHeaders[len(seenHeaders)-1])
}
// --- helpers ---
func toolNames(tools []fantasy.AgentTool) []string {
names := make([]string, 0, len(tools))
for _, t := range tools {
names = append(names, t.Info().Name)
}
return names
}
func findTool(tools []fantasy.AgentTool, name string) fantasy.AgentTool {
for _, t := range tools {
if t.Info().Name == name {
return t
}
}
return nil
}
// TestConnectAll_DisabledServer verifies that disabled configs are
// silently skipped.
func TestConnectAll_DisabledServer(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("disabled", ts.URL)
cfg.Enabled = false
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
assert.Empty(t, tools)
}
// TestConnectAll_CallToolInvalidInput verifies that malformed JSON
// input returns an error response rather than a Go error.
func TestConnectAll_CallToolInvalidInput(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
// Pass syntactically invalid JSON as tool input.
resp, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-bad",
Name: "srv__echo",
Input: `{not json`,
})
require.NoError(t, err, "Run should not return a Go error for bad input")
assert.True(t, resp.IsError)
assert.Contains(t, resp.Content, "invalid JSON input")
}
// TestConnectAll_ToolInfoParameters verifies that tool input schema
// parameters are propagated to the ToolInfo.
func TestConnectAll_ToolInfoParameters(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
info := tools[0].Info()
// The echo tool has a required "input" string parameter.
require.NotNil(t, info.Parameters)
_, hasInput := info.Parameters["input"]
assert.True(t, hasInput, "parameters should contain 'input'")
// The "input" field should also appear in Required.
inputProp, ok := info.Parameters["input"].(map[string]any)
assert.True(t, ok, "input parameter should be a map")
if ok {
propBytes, _ := json.Marshal(inputProp)
assert.Contains(t, string(propBytes), "string")
}
assert.Contains(t, info.Required, "input")
}
// TestConnectAll_APIKeyAuth verifies that api_key auth sends the
// configured header and value on every request.
func TestConnectAll_APIKeyAuth(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
var (
mu sync.Mutex
seenHeaders []string
)
srv := mcpserver.NewMCPServer("apikey-server", "1.0.0")
srv.AddTools(mcpserver.ServerTool{
Tool: mcp.NewTool("check",
mcp.WithDescription("Returns the API key header"),
),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
val := req.Header.Get("X-API-Key")
mu.Lock()
seenHeaders = append(seenHeaders, val)
mu.Unlock()
return mcp.NewToolResultText("key:" + val), nil
},
})
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
cfg := makeConfig("apikey", ts.URL)
cfg.AuthType = "api_key"
cfg.APIKeyHeader = "X-API-Key"
cfg.APIKeyValue = "secret-123"
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
resp, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-apikey",
Name: "apikey__check",
Input: "{}",
})
require.NoError(t, err)
assert.False(t, resp.IsError)
assert.Equal(t, "key:secret-123", resp.Content)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, seenHeaders)
assert.Equal(t, "secret-123", seenHeaders[len(seenHeaders)-1])
}
// TestConnectAll_CustomHeadersAuth verifies that custom_headers
// auth sends the configured headers on every request.
func TestConnectAll_CustomHeadersAuth(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
var (
mu sync.Mutex
seenHeaders []string
)
srv := mcpserver.NewMCPServer("custom-server", "1.0.0")
srv.AddTools(mcpserver.ServerTool{
Tool: mcp.NewTool("check",
mcp.WithDescription("Returns the custom auth header"),
),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
val := req.Header.Get("X-Custom-Auth")
mu.Lock()
seenHeaders = append(seenHeaders, val)
mu.Unlock()
return mcp.NewToolResultText("custom:" + val), nil
},
})
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
cfg := makeConfig("custom", ts.URL)
cfg.AuthType = "custom_headers"
cfg.CustomHeaders = `{"X-Custom-Auth":"custom-val"}`
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
resp, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-custom",
Name: "custom__check",
Input: "{}",
})
require.NoError(t, err)
assert.False(t, resp.IsError)
assert.Equal(t, "custom:custom-val", resp.Content)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, seenHeaders)
assert.Equal(t, "custom-val", seenHeaders[len(seenHeaders)-1])
}
// TestConnectAll_CustomHeadersInvalidJSON verifies that invalid
// JSON in CustomHeaders does not prevent the server from
// connecting. The auth headers are silently skipped.
func TestConnectAll_CustomHeadersInvalidJSON(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("badjson", ts.URL)
cfg.AuthType = "custom_headers"
cfg.CustomHeaders = "{not json}"
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
)
t.Cleanup(cleanup)
// The server should still connect; only auth headers are
// skipped.
require.Len(t, tools, 1)
assert.Equal(t, "badjson__echo", tools[0].Info().Name)
}
// TestConnectAll_ParallelConnections verifies that connecting to
// multiple MCP servers simultaneously returns all discovered
// tools with the correct server slug prefixes.
func TestConnectAll_ParallelConnections(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts1 := newTestMCPServer(t, echoTool())
ts2 := newTestMCPServer(t, greetTool())
ts3 := newTestMCPServer(t, echoTool())
cfg1 := makeConfig("srv1", ts1.URL)
cfg2 := makeConfig("srv2", ts2.URL)
cfg3 := makeConfig("srv3", ts3.URL)
tools, cleanup := mcpclient.ConnectAll(
ctx, logger,
[]database.MCPServerConfig{cfg1, cfg2, cfg3},
nil,
)
t.Cleanup(cleanup)
require.Len(t, tools, 3)
names := toolNames(tools)
assert.Contains(t, names, "srv1__echo")
assert.Contains(t, names, "srv2__greet")
assert.Contains(t, names, "srv3__echo")
}
func TestRedactURL(t *testing.T) {
t.Parallel()
tests := []struct {
name string
input string
expected string
}{
{"plain", "https://mcp.example.com/v1", "https://mcp.example.com/v1"},
{"with userinfo", "https://user:secret@mcp.example.com/v1", "https://mcp.example.com/v1"},
{"with query params", "https://mcp.example.com/v1?api_key=sk-123", "https://mcp.example.com/v1"},
{"with both", "https://user:pass@host/p?key=val", "https://host/p"},
{"invalid url", "://not-a-url", "://not-a-url"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
got := mcpclient.RedactURL(tt.input)
assert.Equal(t, tt.expected, got)
})
}
}
func TestConnectAll_ExpiredToken(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
configID := uuid.New()
cfg := database.MCPServerConfig{
ID: configID,
Slug: "expired-srv",
DisplayName: "Expired Server",
Url: ts.URL,
Transport: "streamable_http",
AuthType: "oauth2",
Enabled: true,
}
// Token exists but is expired.
token := database.MCPServerUserToken{
MCPServerConfigID: configID,
AccessToken: "expired-token",
TokenType: "Bearer",
Expiry: sql.NullTime{Time: time.Now().Add(-1 * time.Hour), Valid: true},
}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token})
t.Cleanup(cleanup)
// The server accepts any auth, so the tool is still discovered
// despite the expired token. The important thing is that the
// warning is logged (verified via IgnoreErrors: true in slogtest).
require.NotEmpty(t, tools)
}
func TestConnectAll_EmptyAccessToken(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts := newTestMCPServer(t, echoTool())
configID := uuid.New()
cfg := database.MCPServerConfig{
ID: configID,
Slug: "empty-tok",
DisplayName: "Empty Token Server",
Url: ts.URL,
Transport: "streamable_http",
AuthType: "oauth2",
Enabled: true,
}
// Token record exists but AccessToken is empty.
token := database.MCPServerUserToken{
MCPServerConfigID: configID,
AccessToken: "",
TokenType: "Bearer",
}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token})
t.Cleanup(cleanup)
// Tool is still discovered (server doesn't require auth), but
// no Authorization header was sent. The warning about empty
// access token is logged.
require.NotEmpty(t, tools)
}
func TestConnectAll_CallToolError(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
// Server with a tool that always returns an error result.
srv := mcpserver.NewMCPServer("error-server", "1.0.0")
srv.AddTools(mcpserver.ServerTool{
Tool: mcp.NewTool("fail_tool",
mcp.WithDescription("Always fails"),
),
Handler: func(_ context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
return &mcp.CallToolResult{
Content: []mcp.Content{mcp.NewTextContent("something broke")},
IsError: true,
}, nil
},
})
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
cfg := makeConfig("err-srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
resp, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-err",
Name: "err-srv__fail_tool",
Input: "{}",
})
require.NoError(t, err, "Run should not return a Go error for MCP-level errors")
assert.True(t, resp.IsError, "response should be flagged as error")
assert.Contains(t, resp.Content, "something broke")
}