mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(agent): gate MCP tool discovery on startup (#25034)
The first `/mcp/tools` request could race workspace startup and return an empty tool list before startup scripts had a chance to write `.mcp.json`. Chatd may only discover tools once for a turn, so that empty response could hide workspace MCP tools even though the agent loaded them later. Make the manager wait for startup to settle before treating missing MCP config files as a real empty state. Tool listing now goes through one manager-owned path that starts reload work independently of caller cancellation; caller contexts only bound that caller's wait. After the first reload body settles, transient reload errors return cached tools with the error so the HTTP handler can degrade to the last known tool set instead of returning `[]`. The handler is intentionally thin: it asks the manager for tools, logs any degraded path, and still returns the tool response shape callers already expect. Tests cover startup gating, caller-canceled waits, manager close, reload timeout via quartz, and cached-tool fallback after a later reload error.
This commit is contained in:
@@ -1413,6 +1413,7 @@ func (a *agent) handleManifest(manifestOK *checkpoint) func(ctx context.Context,
|
||||
// lifecycle transition to avoid delaying Ready.
|
||||
// This runs inside the tracked goroutine so it
|
||||
// is properly awaited on shutdown.
|
||||
a.mcpManager.MarkStartupSettled()
|
||||
if mcpErr := a.mcpManager.Reload(a.gracefulCtx, a.contextConfigAPI.MCPConfigFiles()); mcpErr != nil {
|
||||
a.logger.Warn(ctx, "failed to reload workspace MCP servers", slog.Error(mcpErr))
|
||||
}
|
||||
|
||||
+15
-48
@@ -22,20 +22,11 @@ type API struct {
|
||||
mcpConfigFiles func() []string
|
||||
}
|
||||
|
||||
// NewAPI creates a new MCP API handler backed by the given
|
||||
// manager. The mcpConfigFiles callback returns the current
|
||||
// resolved config file paths; it is called on every tool-list
|
||||
// NewAPI creates a new MCP API handler. mcpConfigFiles returns
|
||||
// the resolved .mcp.json paths and is called on every tool-list
|
||||
// request to detect config changes.
|
||||
func NewAPI(
|
||||
logger slog.Logger,
|
||||
manager *Manager,
|
||||
mcpConfigFiles func() []string,
|
||||
) *API {
|
||||
return &API{
|
||||
logger: logger,
|
||||
manager: manager,
|
||||
mcpConfigFiles: mcpConfigFiles,
|
||||
}
|
||||
func NewAPI(logger slog.Logger, m *Manager, mcpConfigFiles func() []string) *API {
|
||||
return &API{logger: logger, manager: m, mcpConfigFiles: mcpConfigFiles}
|
||||
}
|
||||
|
||||
// Routes returns the HTTP handler for MCP-related routes.
|
||||
@@ -46,50 +37,26 @@ func (api *API) Routes() http.Handler {
|
||||
return r
|
||||
}
|
||||
|
||||
// handleListTools checks whether any .mcp.json config file
|
||||
// has changed since the last reload, triggering a differential
|
||||
// reload if so, then returns the cached MCP tool definitions.
|
||||
// The ?refresh=true query parameter forces a tool re-scan
|
||||
// independent of config changes.
|
||||
// handleListTools returns the current MCP tool cache after the
|
||||
// manager performs startup-safe config synchronization.
|
||||
func (api *API) handleListTools(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
logger := api.logger.With(agentchat.Fields(ctx)...)
|
||||
|
||||
// Check config freshness and reload if changed.
|
||||
var reloaded bool
|
||||
paths := api.mcpConfigFiles()
|
||||
if api.manager.SnapshotChanged(paths) {
|
||||
if err := api.manager.Reload(ctx, paths); err != nil {
|
||||
// Categorize the error for operator debugging.
|
||||
switch {
|
||||
case errors.Is(err, context.Canceled):
|
||||
logger.Warn(ctx, "mcp reload canceled by caller", slog.Error(err))
|
||||
case errors.Is(err, context.DeadlineExceeded):
|
||||
logger.Warn(ctx, "mcp reload timed out", slog.Error(err))
|
||||
default:
|
||||
logger.Warn(ctx, "mcp reload failed", slog.Error(err))
|
||||
}
|
||||
// Fall through to return whatever tools we have.
|
||||
} else {
|
||||
reloaded = true
|
||||
tools, err := api.manager.Tools(ctx, api.mcpConfigFiles())
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, context.Canceled):
|
||||
logger.Warn(ctx, "mcp tool list canceled by caller", slog.Error(err))
|
||||
case errors.Is(err, context.DeadlineExceeded):
|
||||
logger.Warn(ctx, "mcp tool list timed out", slog.Error(err))
|
||||
default:
|
||||
logger.Warn(ctx, "mcp tool list failed", slog.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// Allow callers to force a tool re-scan before listing.
|
||||
// Skip if a config reload ran above, since it already
|
||||
// refreshes tools as part of the reload.
|
||||
if r.URL.Query().Get("refresh") == "true" && !reloaded {
|
||||
if err := api.manager.RefreshTools(ctx); err != nil {
|
||||
logger.Warn(ctx, "failed to refresh MCP tools", slog.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
tools := api.manager.Tools()
|
||||
// Ensure non-nil so JSON serialization returns [] not null.
|
||||
if tools == nil {
|
||||
tools = []workspacesdk.MCPToolInfo{}
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, workspacesdk.ListMCPToolsResponse{
|
||||
Tools: tools,
|
||||
})
|
||||
|
||||
@@ -1,11 +1,15 @@
|
||||
package agentmcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -70,6 +74,7 @@ func TestHandleListTools_ReloadOnChange(t *testing.T) {
|
||||
if tc.closeManager {
|
||||
require.NoError(t, m.Close())
|
||||
} else {
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
}
|
||||
|
||||
@@ -108,6 +113,7 @@ func TestHandleListTools_ReloadOnChange(t *testing.T) {
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv1": entry1})
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
@@ -145,7 +151,10 @@ func TestHandleListTools_ReloadOnChange(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestHandleListTools_RefreshParam(t *testing.T) {
|
||||
// TestHandleListTools_ReloadsAfterStartupSettled exercises the
|
||||
// cold-start path end-to-end against a real *Manager. Startup has
|
||||
// settled, so the handler may drive the first safe reload.
|
||||
func TestHandleListTools_ReloadsAfterStartupSettled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if os.Getenv("TEST_MCP_FAKE_SERVER") == "1" {
|
||||
@@ -153,76 +162,160 @@ func TestHandleListTools_RefreshParam(t *testing.T) {
|
||||
return
|
||||
}
|
||||
|
||||
t.Run("RefreshTrueUnchangedSnapshot", func(t *testing.T) {
|
||||
// Exercises the ?refresh=true code path when the config
|
||||
// snapshot is unchanged. Verifies the endpoint returns
|
||||
// tools without error.
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
// No prior m.Reload: snapshot empty and tools unset.
|
||||
require.Empty(t, m.cachedTools(), "manager should start with no tools")
|
||||
|
||||
api := NewAPI(logger, m, func() []string {
|
||||
return []string{configPath}
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/tools?refresh=true", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
api.Routes().ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var resp workspacesdk.ListMCPToolsResponse
|
||||
require.NoError(t, json.NewDecoder(rec.Body).Decode(&resp))
|
||||
// Tool should still be present after refresh.
|
||||
require.Len(t, resp.Tools, 1)
|
||||
assert.Contains(t, resp.Tools[0].Name, "echo")
|
||||
api := NewAPI(logger, m, func() []string {
|
||||
return []string{configPath}
|
||||
})
|
||||
|
||||
t.Run("RefreshTrueWithChangedConfig", func(t *testing.T) {
|
||||
// Exercises the ?refresh=true code path when the config
|
||||
// has also changed. The reload path already calls
|
||||
// RefreshTools, so the handler skips the redundant call.
|
||||
// This test covers the branch; it cannot observe the
|
||||
// skip without a mock.
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
req := httptest.NewRequest(http.MethodGet, "/tools", nil).WithContext(ctx)
|
||||
rec := httptest.NewRecorder()
|
||||
api.Routes().ServeHTTP(rec, req)
|
||||
|
||||
_, entry1 := fakeMCPServerConfig(t, "srv1")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv1": entry1})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var resp workspacesdk.ListMCPToolsResponse
|
||||
require.NoError(t, json.NewDecoder(rec.Body).Decode(&resp))
|
||||
require.Len(t, resp.Tools, 1)
|
||||
assert.Contains(t, resp.Tools[0].Name, "echo")
|
||||
}
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
func TestHandleListTools_WaitsForStartupSettled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
if os.Getenv("TEST_MCP_FAKE_SERVER") == "1" {
|
||||
runFakeMCPServer()
|
||||
return
|
||||
}
|
||||
|
||||
api := NewAPI(logger, m, func() []string {
|
||||
return []string{configPath}
|
||||
})
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
// Mutate config.
|
||||
_, entry2 := fakeMCPServerConfig(t, "srv2")
|
||||
writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv2": entry2})
|
||||
pathsRequested := make(chan struct{})
|
||||
var pathsOnce sync.Once
|
||||
api := NewAPI(logger, m, func() []string {
|
||||
pathsOnce.Do(func() { close(pathsRequested) })
|
||||
return []string{configPath}
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/tools?refresh=true", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/tools", nil).WithContext(ctx)
|
||||
rec := httptest.NewRecorder()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
api.Routes().ServeHTTP(rec, req)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-pathsRequested:
|
||||
case <-ctx.Done():
|
||||
t.Fatalf("handler did not request paths: %v", ctx.Err())
|
||||
}
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("handler returned before startup settled")
|
||||
default:
|
||||
}
|
||||
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
m.MarkStartupSettled()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-ctx.Done():
|
||||
t.Fatalf("handler did not return after startup settled: %v", ctx.Err())
|
||||
}
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var resp workspacesdk.ListMCPToolsResponse
|
||||
require.NoError(t, json.NewDecoder(rec.Body).Decode(&resp))
|
||||
require.Len(t, resp.Tools, 1)
|
||||
assert.Contains(t, resp.Tools[0].Name, "echo")
|
||||
}
|
||||
|
||||
func TestHandleListTools_LogsListErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var resp workspacesdk.ListMCPToolsResponse
|
||||
require.NoError(t, json.NewDecoder(rec.Body).Decode(&resp))
|
||||
require.Len(t, resp.Tools, 1)
|
||||
assert.Contains(t, resp.Tools[0].Name, "srv2")
|
||||
})
|
||||
cases := []struct {
|
||||
name string
|
||||
ctx func() context.Context
|
||||
closeManager bool
|
||||
message string
|
||||
}{
|
||||
{
|
||||
name: "Canceled",
|
||||
ctx: func() context.Context {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
return ctx
|
||||
},
|
||||
message: "mcp tool list canceled by caller",
|
||||
},
|
||||
{
|
||||
name: "DeadlineExceeded",
|
||||
ctx: func() context.Context {
|
||||
ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second))
|
||||
cancel()
|
||||
return ctx
|
||||
},
|
||||
message: "mcp tool list timed out",
|
||||
},
|
||||
{
|
||||
name: "ManagerClosed",
|
||||
ctx: context.Background,
|
||||
closeManager: true,
|
||||
message: "mcp tool list failed",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := tc.ctx()
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(context.Background(), logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
if tc.closeManager {
|
||||
require.NoError(t, m.Close())
|
||||
}
|
||||
|
||||
api := NewAPI(logger, m, func() []string {
|
||||
return []string{configPath}
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/tools", nil).WithContext(ctx)
|
||||
rec := httptest.NewRecorder()
|
||||
api.Routes().ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
entries := sink.Entries(func(e slog.SinkEntry) bool {
|
||||
return e.Message == tc.message
|
||||
})
|
||||
require.Len(t, entries, 1)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+216
-33
@@ -27,6 +27,7 @@ import (
|
||||
"github.com/coder/coder/v2/agent/usershell"
|
||||
"github.com/coder/coder/v2/buildinfo"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
// ToolNameSep separates the server name from the original tool name
|
||||
@@ -42,6 +43,10 @@ const connectTimeout = 30 * time.Second
|
||||
// take before being canceled.
|
||||
const toolCallTimeout = 60 * time.Second
|
||||
|
||||
// toolsReloadTimeout bounds how long Tools waits for a
|
||||
// post-startup reload to settle.
|
||||
const toolsReloadTimeout = 35 * time.Second
|
||||
|
||||
var (
|
||||
// ErrInvalidToolName is returned when the tool name format
|
||||
// is not "server__tool".
|
||||
@@ -49,6 +54,11 @@ var (
|
||||
// ErrUnknownServer is returned when no MCP server matches
|
||||
// the prefix in the tool name.
|
||||
ErrUnknownServer = xerrors.New("unknown MCP server")
|
||||
// ErrManagerClosed is returned by Reload and Tools after
|
||||
// Close. Close cancels the Manager's derived context, so this
|
||||
// sentinel keeps explicit Close distinguishable from parent
|
||||
// context cancellation.
|
||||
ErrManagerClosed = xerrors.New("manager closed")
|
||||
)
|
||||
|
||||
// fileSnapshot records the identity of a config file at the time
|
||||
@@ -59,22 +69,43 @@ type fileSnapshot struct {
|
||||
size int64
|
||||
}
|
||||
|
||||
type reloadResult = tailscalesingleflight.Result[struct{}]
|
||||
|
||||
// Manager manages connections to MCP servers discovered from a
|
||||
// workspace's .mcp.json file. It caches the aggregated tool list
|
||||
// and proxies tool calls to the appropriate server.
|
||||
type Manager struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
execer agentexec.Execer
|
||||
updateEnv func(current []string) ([]string, error)
|
||||
|
||||
mu sync.RWMutex
|
||||
logger slog.Logger
|
||||
clock quartz.Clock
|
||||
closed bool
|
||||
servers map[string]*serverEntry
|
||||
tools []workspacesdk.MCPToolInfo
|
||||
snapshot map[string]fileSnapshot
|
||||
serverGen uint64
|
||||
sf tailscalesingleflight.Group[string, struct{}]
|
||||
|
||||
// startupSettled is closed once startup scripts reach a terminal
|
||||
// state. Before that, missing MCP config files are unknown
|
||||
// because startup scripts may still create them.
|
||||
startupSettled chan struct{}
|
||||
startupOnce sync.Once
|
||||
|
||||
// firstSyncSettled records that a reload body reached a
|
||||
// terminal result, successful or not. It gates whether callers
|
||||
// may receive cached tools after reload errors.
|
||||
firstSyncSettled bool
|
||||
|
||||
// closedCh is closed by Close to unblock waiters that do not
|
||||
// otherwise observe Close (the parent ctx is owned by the
|
||||
// caller and may outlive Close).
|
||||
closedCh chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
// serverEntry pairs a server config with its connected client.
|
||||
@@ -93,55 +124,199 @@ func NewManager(
|
||||
execer agentexec.Execer,
|
||||
updateEnv func([]string) ([]string, error),
|
||||
) *Manager {
|
||||
managerCtx, cancel := context.WithCancel(ctx)
|
||||
return &Manager{
|
||||
ctx: ctx,
|
||||
logger: logger,
|
||||
execer: execer,
|
||||
updateEnv: updateEnv,
|
||||
servers: make(map[string]*serverEntry),
|
||||
snapshot: make(map[string]fileSnapshot),
|
||||
ctx: managerCtx,
|
||||
cancel: cancel,
|
||||
logger: logger,
|
||||
clock: quartz.NewReal(),
|
||||
execer: execer,
|
||||
updateEnv: updateEnv,
|
||||
servers: make(map[string]*serverEntry),
|
||||
snapshot: make(map[string]fileSnapshot),
|
||||
startupSettled: make(chan struct{}),
|
||||
closedCh: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Reload checks whether config files have changed and, if so,
|
||||
// performs a differential reconnect. Concurrent callers are
|
||||
// coalesced via singleflight; the reload body runs under the
|
||||
// Manager's lifetime context so it survives caller cancellation.
|
||||
// Reload ensures the tool cache reflects the current config.
|
||||
//
|
||||
// If config files differ from the last snapshot, a singleflight
|
||||
// differential reconnect is driven and Reload waits for it. If the
|
||||
// snapshot is current, Reload returns immediately.
|
||||
//
|
||||
// Starting and running the reload is manager-scoped. Caller contexts
|
||||
// may bound only that caller's wait for the reload result. They are
|
||||
// never passed to, and must not suppress, the reload body.
|
||||
func (m *Manager) Reload(ctx context.Context, paths []string) error {
|
||||
m.mu.RLock()
|
||||
closed := m.closed
|
||||
hasSnapshot := len(m.snapshot) > 0
|
||||
m.mu.RUnlock()
|
||||
if closed {
|
||||
return xerrors.New("manager closed")
|
||||
ch, started, err := m.startReloadIfNeeded(paths)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Double-check: another goroutine may have completed a
|
||||
// reload between the caller's SnapshotChanged and this
|
||||
// call. The singleflight body uses its own resolved paths.
|
||||
if hasSnapshot && !m.SnapshotChanged(paths) {
|
||||
if !started {
|
||||
return nil
|
||||
}
|
||||
return m.waitReload(ctx, ch, 0)
|
||||
}
|
||||
|
||||
// MarkStartupSettled marks startup scripts as terminal for MCP
|
||||
// config purposes. Missing config files after this point are a real
|
||||
// empty config, not an unknown startup state.
|
||||
func (m *Manager) MarkStartupSettled() {
|
||||
m.startupOnce.Do(func() { close(m.startupSettled) })
|
||||
}
|
||||
|
||||
// Tools returns the current MCP tool cache after startup-safe config
|
||||
// synchronization.
|
||||
//
|
||||
// Before startup has settled via MarkStartupSettled, Tools blocks until
|
||||
// settlement or ctx cancels. After settlement, it drives a config reload
|
||||
// bounded by toolsReloadTimeout.
|
||||
//
|
||||
// On error before the first sync settles, Tools returns nil tools and
|
||||
// the error. On error after a prior sync, it returns cached tools and
|
||||
// the error so callers can degrade gracefully.
|
||||
func (m *Manager) Tools(ctx context.Context, paths []string) ([]workspacesdk.MCPToolInfo, error) {
|
||||
if err := m.waitForStartupSettled(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ch, started, err := m.startReloadIfNeeded(paths)
|
||||
if err != nil {
|
||||
return m.toolsAfterReloadError(err)
|
||||
}
|
||||
if !started {
|
||||
return normalizeTools(m.cachedTools()), nil
|
||||
}
|
||||
|
||||
if err := m.waitReload(ctx, ch, toolsReloadTimeout); err != nil {
|
||||
return m.toolsAfterReloadError(err)
|
||||
}
|
||||
return normalizeTools(m.cachedTools()), nil
|
||||
}
|
||||
|
||||
func (m *Manager) waitForStartupSettled(ctx context.Context) error {
|
||||
select {
|
||||
case <-m.startupSettled:
|
||||
return nil
|
||||
default:
|
||||
}
|
||||
|
||||
select {
|
||||
case <-m.startupSettled:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-m.ctx.Done():
|
||||
if err := m.closeErr(); err != nil {
|
||||
return err
|
||||
}
|
||||
return m.ctx.Err()
|
||||
case <-m.closedCh:
|
||||
return ErrManagerClosed
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) toolsAfterReloadError(err error) ([]workspacesdk.MCPToolInfo, error) {
|
||||
m.mu.RLock()
|
||||
firstSyncSettled := m.firstSyncSettled
|
||||
tools := slices.Clone(m.tools)
|
||||
m.mu.RUnlock()
|
||||
if !firstSyncSettled {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeTools(tools), err
|
||||
}
|
||||
|
||||
func normalizeTools(tools []workspacesdk.MCPToolInfo) []workspacesdk.MCPToolInfo {
|
||||
if tools == nil {
|
||||
return []workspacesdk.MCPToolInfo{}
|
||||
}
|
||||
return tools
|
||||
}
|
||||
|
||||
// startReloadIfNeeded registers the reload with the singleflight group
|
||||
// using a fixed key so concurrent triggers share one body. The body
|
||||
// always runs under m.ctx. The returned channel yields the body's result
|
||||
// exactly once.
|
||||
//
|
||||
// All concurrent callers share one in-flight reload keyed by "reload".
|
||||
// If a concurrent caller resolves different paths, its paths are not
|
||||
// consulted. The next SnapshotChanged check after this reload completes
|
||||
// will detect the mismatch and trigger a fresh reload.
|
||||
func (m *Manager) startReloadIfNeeded(paths []string) (<-chan reloadResult, bool, error) {
|
||||
m.mu.RLock()
|
||||
closed := m.closed
|
||||
firstSyncSettled := m.firstSyncSettled
|
||||
m.mu.RUnlock()
|
||||
if closed {
|
||||
return nil, false, ErrManagerClosed
|
||||
}
|
||||
if err := m.ctx.Err(); err != nil {
|
||||
if closeErr := m.closeErr(); closeErr != nil {
|
||||
return nil, false, closeErr
|
||||
}
|
||||
return nil, false, err
|
||||
}
|
||||
if firstSyncSettled && !m.SnapshotChanged(paths) {
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
// All concurrent callers share one in-flight reload keyed
|
||||
// by "". If a concurrent caller resolves different paths
|
||||
// (e.g. after a manifest reconnect), its paths are not
|
||||
// consulted; the next SnapshotChanged check after this
|
||||
// reload completes will detect the mismatch and trigger
|
||||
// a fresh reload.
|
||||
ch := m.sf.DoChan("reload", func() (struct{}, error) {
|
||||
defer m.markFirstSyncSettled()
|
||||
err := m.doReload(m.ctx, paths)
|
||||
return struct{}{}, err
|
||||
})
|
||||
return ch, true, nil
|
||||
}
|
||||
|
||||
func (m *Manager) waitReload(ctx context.Context, ch <-chan reloadResult, timeout time.Duration) error {
|
||||
// Prefer caller cancellation when it already happened before the
|
||||
// wait. Otherwise select may choose a ready reload result instead.
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var timeoutC <-chan time.Time
|
||||
if timeout > 0 {
|
||||
timer := m.clock.NewTimer(timeout, "agentmcp", "tools_reload")
|
||||
defer timer.Stop()
|
||||
timeoutC = timer.C
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case res := <-ch:
|
||||
return res.Err
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-timeoutC:
|
||||
return xerrors.Errorf("tools reload timed out after %s: %w", timeout, context.DeadlineExceeded)
|
||||
case <-m.ctx.Done():
|
||||
if err := m.closeErr(); err != nil {
|
||||
return err
|
||||
}
|
||||
return m.ctx.Err()
|
||||
case <-m.closedCh:
|
||||
return ErrManagerClosed
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) closeErr() error {
|
||||
m.mu.RLock()
|
||||
closed := m.closed
|
||||
m.mu.RUnlock()
|
||||
if closed {
|
||||
return ErrManagerClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) markFirstSyncSettled() {
|
||||
m.mu.Lock()
|
||||
m.firstSyncSettled = true
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
// SnapshotChanged checks whether any config file has changed
|
||||
// since the last reload by comparing os.Stat results against
|
||||
// the stored snapshot.
|
||||
@@ -306,7 +481,7 @@ func (m *Manager) classifyServers(wanted map[string]ServerConfig) (*serverDiff,
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
if m.closed {
|
||||
return nil, xerrors.New("manager closed")
|
||||
return nil, ErrManagerClosed
|
||||
}
|
||||
|
||||
diff := &serverDiff{
|
||||
@@ -385,7 +560,7 @@ func (m *Manager) installServers(
|
||||
for _, cs := range connected {
|
||||
_ = cs.client.Close()
|
||||
}
|
||||
return nil, xerrors.New("manager closed")
|
||||
return nil, ErrManagerClosed
|
||||
}
|
||||
|
||||
newConnected := make(map[string]connectedServer, len(connected))
|
||||
@@ -442,8 +617,8 @@ func captureSnapshot(paths []string) map[string]fileSnapshot {
|
||||
return snap
|
||||
}
|
||||
|
||||
// Tools returns the cached tool list. Thread-safe.
|
||||
func (m *Manager) Tools() []workspacesdk.MCPToolInfo {
|
||||
// cachedTools returns the cached tool list. Thread-safe.
|
||||
func (m *Manager) cachedTools() []workspacesdk.MCPToolInfo {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
@@ -587,6 +762,7 @@ func (m *Manager) Close() error {
|
||||
defer m.mu.Unlock()
|
||||
|
||||
m.closed = true
|
||||
m.closeOnce.Do(func() { close(m.closedCh) })
|
||||
var errs []error
|
||||
for _, entry := range m.servers {
|
||||
if err := entry.client.Close(); err != nil {
|
||||
@@ -600,7 +776,14 @@ func (m *Manager) Close() error {
|
||||
}
|
||||
}
|
||||
m.servers = make(map[string]*serverEntry)
|
||||
// Prevent an in-flight RefreshTools from repopulating tools
|
||||
// after Close clears the cache.
|
||||
m.serverGen++
|
||||
m.tools = nil
|
||||
|
||||
// Cancel while holding the lock so waiters that observe
|
||||
// m.ctx.Done also observe m.closed when checking closeErr.
|
||||
m.cancel()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
|
||||
@@ -6,15 +6,21 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/agent/agentexec"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func TestSplitToolName(t *testing.T) {
|
||||
@@ -239,11 +245,320 @@ func TestConnectServer_StdioProcessSurvivesConnect(t *testing.T) {
|
||||
listCtx, listCancel := context.WithTimeout(ctx, testutil.WaitShort)
|
||||
defer listCancel()
|
||||
result, err := client.ListTools(listCtx, mcp.ListToolsRequest{})
|
||||
require.NoError(t, err, "ListTools should succeed — server must be alive after connect")
|
||||
require.NoError(t, err, "ListTools should succeed, server must be alive after connect")
|
||||
require.Len(t, result.Tools, 1)
|
||||
assert.Equal(t, "echo", result.Tools[0].Name)
|
||||
}
|
||||
|
||||
func TestManager_WaitReloadTimeout(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
clock := quartz.NewMock(t)
|
||||
timerTrap := clock.Trap().NewTimer("agentmcp", "tools_reload")
|
||||
defer timerTrap.Close()
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.clock = clock
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- m.waitReload(ctx, make(chan reloadResult), time.Minute)
|
||||
}()
|
||||
|
||||
call := timerTrap.MustWait(ctx)
|
||||
require.Equal(t, time.Minute, call.Duration)
|
||||
call.MustRelease(ctx)
|
||||
|
||||
clock.Advance(time.Minute).MustWait(ctx)
|
||||
err := testutil.RequireReceive(ctx, t, done)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
assert.Contains(t, err.Error(), "tools reload timed out after 1m0s")
|
||||
}
|
||||
|
||||
func TestManager_ToolsStartupGate(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if os.Getenv("TEST_MCP_FAKE_SERVER") == "1" {
|
||||
runFakeMCPServer()
|
||||
return
|
||||
}
|
||||
|
||||
t.Run("MissingBeforeStartupCanAppearBeforeSettlement", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
type result struct {
|
||||
tools []workspacesdk.MCPToolInfo
|
||||
err error
|
||||
}
|
||||
done := make(chan result, 1)
|
||||
go func() {
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
done <- result{tools: tools, err: err}
|
||||
}()
|
||||
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
m.MarkStartupSettled()
|
||||
|
||||
select {
|
||||
case got := <-done:
|
||||
require.NoError(t, got.err)
|
||||
require.Len(t, got.tools, 1)
|
||||
assert.Contains(t, got.tools[0].Name, "echo")
|
||||
case <-ctx.Done():
|
||||
t.Fatalf("Tools did not return after startup settled: %v", ctx.Err())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("MissingAfterStartupReturnsEmptyAndMarksFirstSync", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, tools)
|
||||
|
||||
m.mu.RLock()
|
||||
firstSyncSettled := m.firstSyncSettled
|
||||
m.mu.RUnlock()
|
||||
assert.True(t, firstSyncSettled)
|
||||
})
|
||||
|
||||
t.Run("ConfigAppearsAfterEmptySyncReloads", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, tools)
|
||||
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
|
||||
tools, err = m.Tools(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tools, 1)
|
||||
assert.Contains(t, tools[0].Name, "echo")
|
||||
})
|
||||
|
||||
t.Run("ConcurrentFirstListToolsCallsAllSucceed", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
const callers = 5
|
||||
var wg sync.WaitGroup
|
||||
errs := make([]error, callers)
|
||||
toolCounts := make([]int, callers)
|
||||
for i := range callers {
|
||||
wg.Go(func() {
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
errs[i] = err
|
||||
toolCounts[i] = len(tools)
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := range callers {
|
||||
assert.NoError(t, errs[i], "caller %d should not fail", i)
|
||||
assert.Equal(t, 1, toolCounts[i], "caller %d should see tools", i)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("CloseUnblocksStartupWait", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := m.Tools(ctx, []string{configPath})
|
||||
done <- err
|
||||
}()
|
||||
require.NoError(t, m.Close())
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrManagerClosed)
|
||||
case <-ctx.Done():
|
||||
t.Fatalf("Tools did not return after Close: %v", ctx.Err())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("CallerCanceledBeforeStartupReturnsNoTools", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
callerCtx, cancel := context.WithCancel(ctx)
|
||||
cancel()
|
||||
tools, err := m.Tools(callerCtx, []string{configPath})
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.Canceled)
|
||||
assert.Nil(t, tools)
|
||||
})
|
||||
|
||||
t.Run("ManagerCanceledBeforeStartupReturnsNoTools", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitLong))
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
cancel()
|
||||
tools, err := m.Tools(testutil.Context(t, testutil.WaitLong), []string{configPath})
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.Canceled)
|
||||
assert.Nil(t, tools)
|
||||
})
|
||||
|
||||
t.Run("ClosedBeforeFirstSyncReturnsNoTools", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
require.NoError(t, m.Close())
|
||||
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrManagerClosed)
|
||||
assert.Nil(t, tools)
|
||||
})
|
||||
|
||||
t.Run("CanceledBeforeFirstSyncStillStartsReload", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, ".mcp.json")
|
||||
paths := []string{configPath}
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
callerCtx, cancel := context.WithCancel(ctx)
|
||||
cancel()
|
||||
tools, err := m.Tools(callerCtx, paths)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.Canceled)
|
||||
assert.Empty(t, tools)
|
||||
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
m.mu.RLock()
|
||||
firstSyncSettled := m.firstSyncSettled
|
||||
m.mu.RUnlock()
|
||||
return firstSyncSettled && !m.SnapshotChanged(paths)
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
tools, err = m.Tools(ctx, paths)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, tools)
|
||||
})
|
||||
|
||||
t.Run("CanceledAfterFirstSyncNoopReturnsCachedTools", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
tools, err := m.Tools(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tools, 1)
|
||||
|
||||
callerCtx, cancel := context.WithCancel(ctx)
|
||||
cancel()
|
||||
tools, err = m.Tools(callerCtx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tools, 1)
|
||||
assert.Contains(t, tools[0].Name, "echo")
|
||||
})
|
||||
|
||||
t.Run("ManagerCanceledAfterFirstSyncReturnsCachedTools", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
paths := []string{configPath}
|
||||
|
||||
m := NewManager(ctx, logger, agentexec.DefaultExecer, nil)
|
||||
m.MarkStartupSettled()
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
|
||||
tools, err := m.Tools(ctx, paths)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tools, 1)
|
||||
|
||||
_, nextEntry := fakeMCPServerConfig(t, "srv2")
|
||||
writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv2": nextEntry})
|
||||
require.True(t, m.SnapshotChanged(paths))
|
||||
|
||||
m.cancel()
|
||||
tools, err = m.Tools(ctx, paths)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.Canceled)
|
||||
require.Len(t, tools, 1)
|
||||
assert.Contains(t, tools[0].Name, "echo")
|
||||
})
|
||||
}
|
||||
|
||||
// runFakeMCPServer implements a minimal JSON-RPC / MCP server over
|
||||
// stdin/stdout, just enough for initialize + tools/list.
|
||||
func runFakeMCPServer() {
|
||||
|
||||
@@ -220,7 +220,7 @@ func TestSnapshotChanged_MultipleConfigFiles(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Tools from both files should be present.
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 2, "should have tools from both config files")
|
||||
assert.Contains(t, tools[0].Name, "srv1",
|
||||
"first tool should be from first config")
|
||||
@@ -246,7 +246,7 @@ func TestReload(t *testing.T) {
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1, "should have one tool from the fake server")
|
||||
assert.Contains(t, tools[0].Name, "echo")
|
||||
|
||||
@@ -293,7 +293,7 @@ func TestReload(t *testing.T) {
|
||||
assert.NoError(t, err, "caller %d should not fail", i)
|
||||
}
|
||||
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1)
|
||||
})
|
||||
|
||||
@@ -302,9 +302,7 @@ func TestReload(t *testing.T) {
|
||||
mgrCtx := testutil.Context(t, testutil.WaitLong)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
dir := t.TempDir()
|
||||
|
||||
_, entry := fakeMCPServerConfig(t, "srv")
|
||||
configPath := writeMCPConfig(t, dir, map[string]mcpServerEntry{"srv": entry})
|
||||
paths := []string{filepath.Join(dir, ".mcp.json")}
|
||||
|
||||
m := NewManager(mgrCtx, logger, agentexec.DefaultExecer, nil)
|
||||
t.Cleanup(func() { _ = m.Close() })
|
||||
@@ -313,11 +311,18 @@ func TestReload(t *testing.T) {
|
||||
callerCtx, cancel := context.WithCancel(mgrCtx)
|
||||
cancel() // Cancel immediately.
|
||||
|
||||
err := m.Reload(callerCtx, []string{configPath})
|
||||
err := m.Reload(callerCtx, paths)
|
||||
// The caller context is already canceled, so Reload should
|
||||
// return the caller's context error.
|
||||
// return the caller's context error after starting the sync.
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.Canceled)
|
||||
|
||||
testutil.Eventually(mgrCtx, t, func(context.Context) bool {
|
||||
m.mu.RLock()
|
||||
firstSyncSettled := m.firstSyncSettled
|
||||
m.mu.RUnlock()
|
||||
return firstSyncSettled && !m.SnapshotChanged(paths)
|
||||
}, testutil.IntervalFast)
|
||||
})
|
||||
|
||||
t.Run("SequentialReloadsDiffDetect", func(t *testing.T) {
|
||||
@@ -335,7 +340,7 @@ func TestReload(t *testing.T) {
|
||||
// First reload.
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
tools1 := m.Tools()
|
||||
tools1 := m.cachedTools()
|
||||
require.Len(t, tools1, 1)
|
||||
assert.Contains(t, tools1[0].Name, "srv1")
|
||||
|
||||
@@ -347,7 +352,7 @@ func TestReload(t *testing.T) {
|
||||
assert.True(t, m.SnapshotChanged([]string{configPath}))
|
||||
err = m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
tools2 := m.Tools()
|
||||
tools2 := m.cachedTools()
|
||||
require.Len(t, tools2, 1)
|
||||
assert.Contains(t, tools2[0].Name, "srv2")
|
||||
})
|
||||
@@ -388,14 +393,14 @@ func TestReload(t *testing.T) {
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, m.Tools(), 1)
|
||||
require.Len(t, m.cachedTools(), 1)
|
||||
|
||||
// Delete config file.
|
||||
require.NoError(t, os.Remove(configPath))
|
||||
|
||||
err = m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, m.Tools(), "tools should be empty after config deleted")
|
||||
assert.Empty(t, m.cachedTools(), "tools should be empty after config deleted")
|
||||
|
||||
// Subsequent reload finds snapshot unchanged.
|
||||
assert.False(t, m.SnapshotChanged([]string{configPath}))
|
||||
@@ -446,7 +451,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
"unchanged server should reuse client pointer")
|
||||
|
||||
// Both servers should have tools.
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 2)
|
||||
})
|
||||
|
||||
@@ -500,7 +505,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, m.Tools(), 2)
|
||||
require.Len(t, m.cachedTools(), 2)
|
||||
|
||||
// Capture srvB's client before removal.
|
||||
m.mu.RLock()
|
||||
@@ -514,7 +519,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
err = m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1)
|
||||
assert.Contains(t, tools[0].Name, "srvA")
|
||||
|
||||
@@ -540,7 +545,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, m.Tools(), 1)
|
||||
require.Len(t, m.cachedTools(), 1)
|
||||
|
||||
m.mu.RLock()
|
||||
origClient := m.servers["srv"].client
|
||||
@@ -563,7 +568,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
"failed connect should retain old client")
|
||||
|
||||
// Tools should still work.
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1)
|
||||
})
|
||||
|
||||
@@ -581,7 +586,7 @@ func TestDifferentialReload(t *testing.T) {
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1)
|
||||
toolName := tools[0].Name
|
||||
|
||||
@@ -632,7 +637,7 @@ func TestReload_FirstBootPath(t *testing.T) {
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := m.Tools()
|
||||
tools := m.cachedTools()
|
||||
require.Len(t, tools, 1)
|
||||
assert.Contains(t, tools[0].Name, "echo")
|
||||
}
|
||||
@@ -668,6 +673,11 @@ func TestReload_NoopWhenUnchanged(t *testing.T) {
|
||||
err = m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
|
||||
callerCtx, cancel := context.WithCancel(ctx)
|
||||
cancel()
|
||||
err = m.Reload(callerCtx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
|
||||
m.mu.RLock()
|
||||
sameClient := m.servers["srv"].client
|
||||
m.mu.RUnlock()
|
||||
@@ -699,7 +709,7 @@ func TestClose_SuppressesSubprocessExitError(t *testing.T) {
|
||||
|
||||
err := m.Reload(ctx, []string{configPath})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, m.Tools(), 1, "server should be connected")
|
||||
require.Len(t, m.cachedTools(), 1, "server should be connected")
|
||||
|
||||
// Close kills the subprocess. The ExitError guard should
|
||||
// suppress the "signal: killed" error.
|
||||
|
||||
Reference in New Issue
Block a user