feat: add opt-in Coder identity headers for MCP servers (#25153)

This commit is contained in:
Kyle Carberry
2026-05-12 08:54:53 -04:00
committed by GitHub
parent f1d160c7f4
commit b0b07536fc
22 changed files with 563 additions and 58 deletions
+1
View File
@@ -6977,6 +6977,7 @@ func (p *Server) runChat(
mcpTokens = p.refreshExpiredMCPTokens(ctx, logger, mcpConnectConfigs, mcpTokens)
mcpTools, mcpCleanup = mcpclient.ConnectAll(
ctx, logger, mcpConnectConfigs, mcpTokens, chat.OwnerID, p.oidcTokenSource,
chatprovider.CoderHeaders(chat),
)
return nil
})
@@ -0,0 +1,329 @@
package mcpclient_test
import (
"context"
"net/http"
"net/http/httptest"
"sync"
"testing"
"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/chatprovider"
"github.com/coder/coder/v2/coderd/x/chatd/mcpclient"
)
// newHeaderRecordingServer creates a streamable HTTP MCP server with a
// single "ping" tool. Every request's headers are appended to the
// returned slice so tests can assert which headers were forwarded.
func newHeaderRecordingServer(t *testing.T) (*httptest.Server, *sync.Mutex, *[]http.Header) {
t.Helper()
var (
mu sync.Mutex
headers []http.Header
)
srv := mcpserver.NewMCPServer("hdr-server", "1.0.0")
srv.AddTools(mcpserver.ServerTool{
Tool: mcp.NewTool("ping", mcp.WithDescription("records the request headers")),
Handler: func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
mu.Lock()
headers = append(headers, req.Header.Clone())
mu.Unlock()
return mcp.NewToolResultText("ok"), nil
},
})
httpSrv := mcpserver.NewStreamableHTTPServer(srv)
ts := httptest.NewServer(httpSrv)
t.Cleanup(ts.Close)
return ts, &mu, &headers
}
// TestConnectAll_ForwardCoderHeaders_DefaultOff is a regression guard
// that the Coder identity headers are NOT sent when the option is
// left at its default (false).
func TestConnectAll_ForwardCoderHeaders_DefaultOff(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
cfg := makeConfig("no-hdr", ts.URL)
assert.False(t, cfg.ForwardCoderHeaders, "default must be false")
coderHeaders := map[string]string{
chatprovider.HeaderCoderOwnerID: uuid.NewString(),
chatprovider.HeaderCoderChatID: uuid.NewString(),
chatprovider.HeaderCoderWorkspaceID: uuid.NewString(),
}
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "no-hdr__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
for _, h := range *recorded {
assert.Empty(t, h.Get(chatprovider.HeaderCoderOwnerID))
assert.Empty(t, h.Get(chatprovider.HeaderCoderChatID))
assert.Empty(t, h.Get(chatprovider.HeaderCoderSubchatID))
assert.Empty(t, h.Get(chatprovider.HeaderCoderWorkspaceID))
}
}
// TestConnectAll_ForwardCoderHeaders_Enabled verifies that when the
// option is enabled, the Coder identity headers are forwarded on every
// outgoing MCP request, including the subchat and workspace headers.
func TestConnectAll_ForwardCoderHeaders_Enabled(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
ownerID := uuid.New()
chatID := uuid.New()
workspaceID := uuid.New()
subchatID := uuid.New()
cfg := makeConfig("hdr", ts.URL)
cfg.ForwardCoderHeaders = true
// Subchat headers: parent's chat ID lives in X-Coder-Chat-Id, the
// subchat's own ID lives in X-Coder-Subchat-Id.
coderHeaders := chatprovider.CoderHeaders(database.Chat{
ID: subchatID,
OwnerID: ownerID,
ParentChatID: uuid.NullUUID{UUID: chatID, Valid: true},
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
})
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "hdr__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
last := (*recorded)[len(*recorded)-1]
assert.Equal(t, ownerID.String(), last.Get(chatprovider.HeaderCoderOwnerID))
assert.Equal(t, chatID.String(), last.Get(chatprovider.HeaderCoderChatID))
assert.Equal(t, subchatID.String(), last.Get(chatprovider.HeaderCoderSubchatID))
assert.Equal(t, workspaceID.String(), last.Get(chatprovider.HeaderCoderWorkspaceID))
}
// TestConnectAll_ForwardCoderHeaders_RootChat verifies that for a root
// chat (no parent), the chat's own ID is forwarded as
// X-Coder-Chat-Id and the X-Coder-Subchat-Id header is absent.
func TestConnectAll_ForwardCoderHeaders_RootChat(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
ownerID := uuid.New()
chatID := uuid.New()
cfg := makeConfig("hdr-root", ts.URL)
cfg.ForwardCoderHeaders = true
coderHeaders := chatprovider.CoderHeaders(database.Chat{
ID: chatID,
OwnerID: ownerID,
})
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "hdr-root__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
last := (*recorded)[len(*recorded)-1]
assert.Equal(t, ownerID.String(), last.Get(chatprovider.HeaderCoderOwnerID))
assert.Equal(t, chatID.String(), last.Get(chatprovider.HeaderCoderChatID))
assert.Empty(t, last.Get(chatprovider.HeaderCoderSubchatID))
assert.Empty(t, last.Get(chatprovider.HeaderCoderWorkspaceID))
}
// TestConnectAll_ForwardCoderHeaders_WithAPIKeyAuth verifies that the
// api_key auth header is preserved when Coder identity headers are
// forwarded alongside.
func TestConnectAll_ForwardCoderHeaders_WithAPIKeyAuth(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
ownerID := uuid.New()
chatID := uuid.New()
cfg := makeConfig("hdr-apikey", ts.URL)
cfg.AuthType = "api_key"
cfg.APIKeyHeader = "X-Api-Key"
cfg.APIKeyValue = "sekret"
cfg.ForwardCoderHeaders = true
coderHeaders := chatprovider.CoderHeaders(database.Chat{
ID: chatID,
OwnerID: ownerID,
})
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "hdr-apikey__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
last := (*recorded)[len(*recorded)-1]
assert.Equal(t, "sekret", last.Get("X-Api-Key"))
assert.Equal(t, ownerID.String(), last.Get(chatprovider.HeaderCoderOwnerID))
assert.Equal(t, chatID.String(), last.Get(chatprovider.HeaderCoderChatID))
}
// TestConnectAll_ForwardCoderHeaders_WithOAuth2 verifies that the
// oauth2 Authorization header is preserved when Coder identity
// headers are forwarded alongside, and that auth wins on a conflict.
func TestConnectAll_ForwardCoderHeaders_WithOAuth2(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
cfgID := uuid.New()
cfg := makeConfig("hdr-oauth", ts.URL)
cfg.ID = cfgID
cfg.AuthType = "oauth2"
cfg.ForwardCoderHeaders = true
token := database.MCPServerUserToken{
MCPServerConfigID: cfgID,
AccessToken: "oauth-token-xyz",
TokenType: "Bearer",
}
// Intentionally include an Authorization key to verify the auth
// header wins on conflict.
ownerID := uuid.NewString()
coderHeaders := map[string]string{
"Authorization": "Bearer should-be-overridden",
chatprovider.HeaderCoderOwnerID: ownerID,
}
tools, cleanup := mcpclient.ConnectAll(
ctx, logger,
[]database.MCPServerConfig{cfg},
[]database.MCPServerUserToken{token},
uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "hdr-oauth__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
last := (*recorded)[len(*recorded)-1]
assert.Equal(t, "Bearer oauth-token-xyz", last.Get("Authorization"))
assert.Equal(t, ownerID, last.Get(chatprovider.HeaderCoderOwnerID))
}
// TestConnectAll_ForwardCoderHeaders_WithCustomHeaders verifies that
// custom_headers admin-configured values are preserved when Coder
// identity headers are forwarded alongside, including the case where
// the admin configures a custom header whose name only differs from a
// Coder identity header by case. Conflict detection is case-
// insensitive because http.Header.Set canonicalizes header names.
func TestConnectAll_ForwardCoderHeaders_WithCustomHeaders(t *testing.T) {
t.Parallel()
ctx := t.Context()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ts, mu, recorded := newHeaderRecordingServer(t)
ownerID := uuid.New()
chatID := uuid.New()
cfg := makeConfig("hdr-custom", ts.URL)
cfg.AuthType = "custom_headers"
// Include both an unrelated custom header AND a case-variant of
// X-Coder-Owner-Id to exercise the case-insensitive conflict
// check. The admin-configured value MUST win.
cfg.CustomHeaders = `{"X-Tenant":"acme","x-coder-owner-id":"admin-controlled"}`
cfg.ForwardCoderHeaders = true
coderHeaders := chatprovider.CoderHeaders(database.Chat{
ID: chatID,
OwnerID: ownerID,
})
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil,
coderHeaders,
)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
_, err := tools[0].Run(ctx, fantasy.ToolCall{
ID: "call-1", Name: "hdr-custom__ping", Input: "{}",
})
require.NoError(t, err)
mu.Lock()
defer mu.Unlock()
require.NotEmpty(t, *recorded)
last := (*recorded)[len(*recorded)-1]
assert.Equal(t, "acme", last.Get("X-Tenant"))
// The admin's case-variant header must win, because HTTP header
// names are case-insensitive at the transport level.
assert.Equal(t, "admin-controlled", last.Get(chatprovider.HeaderCoderOwnerID))
assert.Equal(t, chatID.String(), last.Get(chatprovider.HeaderCoderChatID))
}
+24 -1
View File
@@ -74,6 +74,7 @@ func ConnectAll(
tokens []database.MCPServerUserToken,
userID uuid.UUID,
oidcSrc UserOIDCTokenSource,
coderHeaders map[string]string,
) ([]fantasy.AgentTool, func()) {
// Index tokens by server config ID so auth header
// construction is O(1) per server.
@@ -109,7 +110,7 @@ func ConnectAll(
eg.Go(func() error {
serverTools, mcpClient, connectErr := connectOne(
ctx, logger, cfg, tokensByConfigID, userID, oidcSrc,
ctx, logger, cfg, tokensByConfigID, userID, oidcSrc, coderHeaders,
)
if connectErr != nil {
logger.Warn(ctx,
@@ -175,9 +176,31 @@ func connectOne(
tokensByConfigID map[uuid.UUID]database.MCPServerUserToken,
userID uuid.UUID,
oidcSrc UserOIDCTokenSource,
coderHeaders map[string]string,
) ([]fantasy.AgentTool, *client.Client, error) {
headers := buildAuthHeaders(ctx, logger, cfg, tokensByConfigID, userID, oidcSrc)
// When opted-in, merge Coder identity headers BEFORE the
// transport is created so any auth header already set above
// wins on a conflict. Conflict detection uses
// http.CanonicalHeaderKey because the upstream transport applies
// http.Header.Set, which canonicalizes keys; without that, an
// admin-configured header that differs only in case from a Coder
// identity header would land in the request map twice and the
// surviving value would be non-deterministic.
if cfg.ForwardCoderHeaders {
canonicalAuth := make(map[string]struct{}, len(headers))
for k := range headers {
canonicalAuth[http.CanonicalHeaderKey(k)] = struct{}{}
}
for k, v := range coderHeaders {
if _, exists := canonicalAuth[http.CanonicalHeaderKey(k)]; exists {
continue
}
headers[k] = v
}
}
tr, err := createTransport(cfg, headers)
if err != nil {
return nil, nil, xerrors.Errorf(
+36 -25
View File
@@ -96,7 +96,7 @@ func TestConnectAll_DiscoverTools(t *testing.T) {
ts := newTestMCPServer(t, echoTool(), greetTool())
cfg := makeConfig("myserver", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
// Two tools should be discovered, namespaced with the server slug.
@@ -121,7 +121,7 @@ func TestConnectAll_CallTool(t *testing.T) {
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -147,7 +147,7 @@ func TestConnectAll_ToolAllowList(t *testing.T) {
// Only allow the "echo" tool.
cfg.ToolAllowList = []string{"echo"}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -165,7 +165,7 @@ func TestConnectAll_ToolDenyList(t *testing.T) {
// Deny the "greet" tool, so only "echo" remains.
cfg.ToolDenyList = []string{"greet"}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -179,7 +179,7 @@ func TestConnectAll_ConnectionFailure(t *testing.T) {
cfg := makeConfig("bad", "http://127.0.0.1:0/does-not-exist")
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
assert.Empty(t, tools, "no tools should be returned for an unreachable server")
@@ -201,6 +201,7 @@ func TestConnectAll_MultipleServers(t *testing.T) {
[]database.MCPServerConfig{cfg1, cfg2},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -227,6 +228,7 @@ func TestConnectAll_NoToolsAfterFiltering(t *testing.T) {
[]database.MCPServerConfig{cfg},
nil,
uuid.Nil, nil,
nil,
)
require.Empty(t, tools)
@@ -255,6 +257,7 @@ func TestConnectAll_DeterministicOrder(t *testing.T) {
},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -284,6 +287,7 @@ func TestConnectAll_DeterministicOrder(t *testing.T) {
},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -317,6 +321,7 @@ func TestConnectAll_DeterministicOrder(t *testing.T) {
[]database.MCPServerConfig{cfg1, cfg2},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -381,6 +386,7 @@ func TestConnectAll_AuthHeaders(t *testing.T) {
[]database.MCPServerConfig{cfg},
[]database.MCPServerUserToken{token},
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -435,7 +441,7 @@ func TestConnectAll_DisabledServer(t *testing.T) {
cfg := makeConfig("disabled", ts.URL)
cfg.Enabled = false
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
assert.Empty(t, tools)
}
@@ -450,7 +456,7 @@ func TestConnectAll_CallToolInvalidInput(t *testing.T) {
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -475,7 +481,7 @@ func TestConnectAll_ToolInfoParameters(t *testing.T) {
ts := newTestMCPServer(t, echoTool())
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -517,7 +523,7 @@ func TestConnectAll_NilRequiredBecomesEmptySlice(t *testing.T) {
ts := newTestMCPServer(t, noRequiredTool)
cfg := makeConfig("srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -570,6 +576,7 @@ func TestConnectAll_APIKeyAuth(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -627,6 +634,7 @@ func TestConnectAll_CustomHeadersAuth(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -664,6 +672,7 @@ func TestConnectAll_CustomHeadersInvalidJSON(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -722,7 +731,7 @@ func TestConnectAll_UserOIDCAuth(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
userID, src,
userID, src, nil,
)
t.Cleanup(cleanup)
@@ -781,7 +790,7 @@ func TestConnectAll_UserOIDCAuth_NoLink(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
uuid.New(), src,
uuid.New(), src, nil,
)
t.Cleanup(cleanup)
@@ -817,7 +826,7 @@ func TestConnectAll_UserOIDCAuth_NilSource(t *testing.T) {
tools, cleanup := mcpclient.ConnectAll(
ctx, logger, []database.MCPServerConfig{cfg}, nil,
uuid.New(), nil,
uuid.New(), nil, nil,
)
t.Cleanup(cleanup)
@@ -846,6 +855,7 @@ func TestConnectAll_ParallelConnections(t *testing.T) {
[]database.MCPServerConfig{cfg1, cfg2, cfg3},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -906,7 +916,7 @@ func TestConnectAll_ExpiredToken(t *testing.T) {
Expiry: sql.NullTime{Time: time.Now().Add(-1 * time.Hour), Valid: true},
}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token}, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token}, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
// The server accepts any auth, so the tool is still discovered
@@ -939,7 +949,7 @@ func TestConnectAll_EmptyAccessToken(t *testing.T) {
TokenType: "Bearer",
}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token}, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, []database.MCPServerUserToken{token}, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
// Tool is still discovered (server doesn't require auth), but
@@ -969,7 +979,7 @@ func TestConnectAll_MCPToolIdentifier(t *testing.T) {
Enabled: true,
}
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1016,6 +1026,7 @@ func TestConnectAll_MCPToolIdentifier_MultipleServers(t *testing.T) {
[]database.MCPServerConfig{cfg1, cfg2},
nil,
uuid.Nil, nil,
nil,
)
t.Cleanup(cleanup)
@@ -1072,7 +1083,7 @@ func TestConnectAll_EmbeddedResourceText(t *testing.T) {
t.Cleanup(ts.Close)
cfg := makeConfig("embed-txt", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1139,7 +1150,7 @@ func TestConnectAll_EmbeddedResourceBlob(t *testing.T) {
t.Cleanup(ts.Close)
cfg := makeConfig("embed-blob", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1219,7 +1230,7 @@ func TestConnectAll_ResourceLink(t *testing.T) {
t.Cleanup(ts.Close)
cfg := makeConfig("res-link", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1263,7 +1274,7 @@ func TestConnectAll_CallToolError(t *testing.T) {
t.Cleanup(ts.Close)
cfg := makeConfig("err-srv", ts.URL)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1287,7 +1298,7 @@ func TestModelIntent_Info_WrapsSchema(t *testing.T) {
cfg := makeConfig("intent-srv", ts.URL)
cfg.ModelIntent = true
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1323,7 +1334,7 @@ func TestModelIntent_Info_NoWrapWhenDisabled(t *testing.T) {
cfg := makeConfig("no-intent", ts.URL)
cfg.ModelIntent = false
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1346,7 +1357,7 @@ func TestModelIntent_Run_UnwrapsProperties(t *testing.T) {
cfg := makeConfig("unwrap-srv", ts.URL)
cfg.ModelIntent = true
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1371,7 +1382,7 @@ func TestModelIntent_Run_UnwrapsFlat(t *testing.T) {
cfg := makeConfig("flat-srv", ts.URL)
cfg.ModelIntent = true
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1396,7 +1407,7 @@ func TestModelIntent_Run_PassthroughWhenDisabled(t *testing.T) {
cfg := makeConfig("pass-srv", ts.URL)
cfg.ModelIntent = false
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)
@@ -1421,7 +1432,7 @@ func TestModelIntent_Run_FallbackOnBadJSON(t *testing.T) {
cfg := makeConfig("bad-srv", ts.URL)
cfg.ModelIntent = true
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil)
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil, uuid.Nil, nil, nil)
t.Cleanup(cleanup)
require.Len(t, tools, 1)