mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add opt-in Coder identity headers for MCP servers (#25153)
This commit is contained in:
@@ -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))
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user