diff --git a/desktop/src/__tests__/mcpApi.test.ts b/desktop/src/__tests__/mcpApi.test.ts new file mode 100644 index 00000000..e96ed1e8 --- /dev/null +++ b/desktop/src/__tests__/mcpApi.test.ts @@ -0,0 +1,34 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' +import { mcpApi } from '../api/mcp' + +const fetchMock = vi.fn() + +afterEach(() => vi.unstubAllGlobals()) + +describe('MCP toggle transport receipt', () => { + it('sends the selected project/session and preserves a failed runtime receipt', async () => { + const response = { + server: { name: 'project/echo', enabled: false, status: 'disabled' }, + sessionSync: { applied: false, reason: 'failed', error: 'control timed out' }, + } + fetchMock.mockResolvedValue(new Response(JSON.stringify(response), { status: 200 })) + vi.stubGlobal('fetch', fetchMock) + + expect(await mcpApi.toggle('project/echo', '/tmp/qa005-project', 'chat-1')).toEqual(response) + const [url, options] = fetchMock.mock.lastCall! + expect(url).toContain('/api/mcp/project%2Fecho/toggle') + expect(JSON.parse(options.body)).toEqual({ cwd: '/tmp/qa005-project', sessionId: 'chat-1' }) + }) + + it('keeps no-session acknowledgement distinct from an applied control', async () => { + const response = { + server: { name: 'echo', enabled: true, status: 'connected' }, + sessionSync: { applied: false, reason: 'no_session' }, + } + fetchMock.mockResolvedValue(new Response(JSON.stringify(response), { status: 200 })) + vi.stubGlobal('fetch', fetchMock) + + expect(await mcpApi.toggle('echo')).toEqual(response) + expect(JSON.parse(fetchMock.mock.lastCall![1].body)).toEqual({}) + }) +}) diff --git a/desktop/src/__tests__/mcpSettings.test.tsx b/desktop/src/__tests__/mcpSettings.test.tsx index a76de70c..225de69a 100644 --- a/desktop/src/__tests__/mcpSettings.test.tsx +++ b/desktop/src/__tests__/mcpSettings.test.tsx @@ -5,6 +5,7 @@ import '@testing-library/jest-dom' import { McpSettings } from '../pages/McpSettings' import { sessionsApi } from '../api/sessions' import { mcpApi } from '../api/mcp' +import { useUIStore } from '../stores/uiStore' import { useMcpStore } from '../stores/mcpStore' import { useSessionStore } from '../stores/sessionStore' import { useSettingsStore } from '../stores/settingsStore' @@ -31,6 +32,8 @@ vi.mock('../api/mcp', async (importOriginal) => { } }) +const originalAddToast = useUIStore.getState().addToast + async function renderLoadedMcpSettings() { const result = render() await waitFor(() => { @@ -41,6 +44,7 @@ async function renderLoadedMcpSettings() { describe('McpSettings', () => { beforeEach(() => { + useUIStore.setState({ addToast: originalAddToast, toasts: [] }) vi.mocked(sessionsApi.getRecentProjects).mockResolvedValue({ projects: [{ projectPath: '/workspace/selected-project', @@ -455,7 +459,7 @@ describe('McpSettings', () => { expect(screen.getByText('Configured')).toBeInTheDocument() expect(screen.getByText('Not loaded in the current chat. Open a chat in this project to use it.')).toBeInTheDocument() expect(screen.queryByText('Connected')).not.toBeInTheDocument() - expect(screen.getByText('Connected in current chat').parentElement?.parentElement).toHaveTextContent('0') + expect(screen.getByText('Connection checks for current project').parentElement?.parentElement).toHaveTextContent('0') expect(refreshServerStatus).not.toHaveBeenCalled() await act(async () => { @@ -545,7 +549,7 @@ describe('McpSettings', () => { }) it('uses the active cwd when toggling a server', async () => { - const toggleServer = vi.fn().mockResolvedValue(undefined) + const toggleServer = vi.fn() const server = { name: 'global-user', scope: 'user', @@ -562,6 +566,7 @@ describe('McpSettings', () => { config: { type: 'http', url: 'https://example.com/mcp', headers: {} }, } as const + toggleServer.mockResolvedValue({ server: { ...server, enabled: false }, sessionSync: { applied: true } }) useMcpStore.setState({ servers: [server], toggleServer, @@ -576,6 +581,57 @@ describe('McpSettings', () => { expect(toggleServer).toHaveBeenCalledWith(server, '/workspace/project', 'session-1') }) + it.each([ + [{ applied: false, reason: 'failed', error: 'Control timeout' }, 'warning', 'Control timeout'], + [{ applied: false, reason: 'not_running' }, 'warning', 'chat is not running'], + [{ applied: false, reason: 'different_project' }, 'warning', 'different project'], + [{ applied: false, reason: 'no_session' }, 'warning', 'No chat was selected'], + [undefined, 'warning', 'not confirmed'], + [{ applied: true }, 'success', 'Disabled MCP server'], + ] as const)('reports the runtime receipt after a saved toggle: %j', async (sessionSync, type, message) => { + const server = { + name: 'echo', scope: 'project', transport: 'stdio', enabled: true, + status: 'connected', statusLabel: 'Connected', configLocation: '/tmp/config', summary: 'echo', + canEdit: true, canRemove: true, canReconnect: true, canToggle: true, + config: { type: 'stdio', command: 'echo', args: [], env: {} }, + } as const + const addToast = vi.fn() + useUIStore.setState({ addToast }) + useMcpStore.setState({ + servers: [server], + toggleServer: vi.fn().mockResolvedValue({ + server: { ...server, enabled: false, status: 'disabled' }, sessionSync, + }), + }) + await renderLoadedMcpSettings() + await act(async () => { fireEvent.click(screen.getByRole('switch')) }) + expect(addToast).toHaveBeenCalledWith({ type, message: expect.stringContaining(message) }) + if (type === 'warning') { + expect(addToast).not.toHaveBeenCalledWith(expect.objectContaining({ type: 'success' })) + } + }) + + it('warns when enabling saved settings could not connect the server', async () => { + const server = { + name: 'echo', scope: 'project', transport: 'stdio', enabled: false, + status: 'disabled', statusLabel: 'Disabled', configLocation: '/tmp/config', summary: 'echo', + canEdit: true, canRemove: true, canReconnect: true, canToggle: true, + config: { type: 'stdio', command: 'echo', args: [], env: {} }, + } as const + const addToast = vi.fn() + useUIStore.setState({ addToast }) + useMcpStore.setState({ + servers: [server], + toggleServer: vi.fn().mockResolvedValue({ + server: { ...server, enabled: true, status: 'failed', statusDetail: 'Command is unavailable' }, + sessionSync: { applied: true }, + }), + }) + await renderLoadedMcpSettings() + await act(async () => { fireEvent.click(screen.getByRole('switch')) }) + expect(addToast).toHaveBeenCalledWith({ type: 'warning', message: 'Command is unavailable' }) + }) + it('clears the selected MCP server when returning to the list', async () => { const selectServer = vi.fn() const server = { diff --git a/desktop/src/__tests__/mcpStoreKnownProjects.test.ts b/desktop/src/__tests__/mcpStoreKnownProjects.test.ts index 4eb88cd7..a3699bc1 100644 --- a/desktop/src/__tests__/mcpStoreKnownProjects.test.ts +++ b/desktop/src/__tests__/mcpStoreKnownProjects.test.ts @@ -152,7 +152,7 @@ describe('fetchServersForKnownProjects', () => { const server = useMcpStore.getState().servers[0]! const toggled = await useMcpStore.getState().toggleServer(server, server.projectPath) - expect(toggled.activeInCurrentContext).toBe(true) + expect(toggled.server.activeInCurrentContext).toBe(true) expect(useMcpStore.getState().servers).toHaveLength(1) expect(useMcpStore.getState().servers[0]).toMatchObject({ enabled: false, @@ -164,6 +164,38 @@ describe('fetchServersForKnownProjects', () => { expect(useMcpStore.getState().servers[0]?.enabled).toBe(false) }) + it('preserves the saved state and failed runtime receipt after toggling', async () => { + const server = record('echo', 'project', { projectPath: '/tmp/qa-005-project' }) + const disabled = { ...server, enabled: false, status: 'disabled' as const } + const sessionSync = { applied: false, reason: 'failed' as const, error: 'Control timeout' } + vi.mocked(mcpApi.toggle).mockResolvedValue({ server: disabled, sessionSync }) + useMcpStore.setState({ servers: [server], selectedServer: server }) + + const result = await useMcpStore.getState().toggleServer(server, server.projectPath, 'session-1') + + expect(result).toEqual({ server: disabled, sessionSync }) + expect(useMcpStore.getState().servers).toEqual([disabled]) + expect(useMcpStore.getState().selectedServer).toEqual(disabled) + expect(mcpApi.toggle).toHaveBeenCalledWith('echo', server.projectPath, 'session-1') + }) + + it('does not let an older connection check overwrite a completed disable', async () => { + const server = record('echo', 'project', { projectPath: '/tmp/qa-005-project' }) + let finishCheck!: (response: { server: McpServerRecord }) => void + vi.mocked(mcpApi.status).mockReturnValue(new Promise((resolve) => { finishCheck = resolve })) + const disabled = { ...server, enabled: false, status: 'disabled' as const } + vi.mocked(mcpApi.toggle).mockResolvedValue({ server: disabled, sessionSync: { applied: true } }) + useMcpStore.setState({ servers: [server], selectedServer: server }) + + const checking = useMcpStore.getState().refreshServerStatus(server, server.projectPath) + await useMcpStore.getState().toggleServer(server, server.projectPath, 'session-1') + finishCheck({ server: { ...server, status: 'connected' } }) + + expect(await checking).toMatchObject({ enabled: false, status: 'disabled' }) + expect(useMcpStore.getState().servers).toEqual([disabled]) + expect(useMcpStore.getState().selectedServer).toEqual(disabled) + }) + it('keeps an inherited project server active when a later context finds the same declaration', async () => { vi.mocked(sessionsApi.getRecentProjects).mockResolvedValue({ projects: [{ realPath: '/repo/root/packages/desktop' }], diff --git a/desktop/src/api/mcp.ts b/desktop/src/api/mcp.ts index 904fbe3a..cd8ce7a3 100644 --- a/desktop/src/api/mcp.ts +++ b/desktop/src/api/mcp.ts @@ -1,5 +1,5 @@ import { api } from './client' -import type { McpServerRecord, McpUpsertPayload } from '../types/mcp' +import type { McpServerRecord, McpToggleResult, McpUpsertPayload } from '../types/mcp' export const mcpApi = { list: (cwd?: string) => { @@ -39,7 +39,7 @@ export const mcpApi = { }, toggle: (name: string, cwd?: string, sessionId?: string) => { - return api.post<{ server: McpServerRecord }>( + return api.post( `/api/mcp/${encodeURIComponent(name)}/toggle`, { ...(cwd ? { cwd } : {}), diff --git a/desktop/src/i18n/locales/en.ts b/desktop/src/i18n/locales/en.ts index a4965b30..8570216d 100644 --- a/desktop/src/i18n/locales/en.ts +++ b/desktop/src/i18n/locales/en.ts @@ -1060,7 +1060,7 @@ Row 9, all 8 cells: continuing from straight down, turning left through lower-le 'settings.mcp.empty': 'No MCP servers configured yet', 'settings.mcp.emptyHint': 'Add a custom stdio, HTTP, or SSE MCP server to start extending tool access.', 'settings.mcp.stats.total': 'Total servers', - 'settings.mcp.stats.connected': 'Connected in current chat', + 'settings.mcp.stats.connected': 'Connection checks for current project', 'settings.mcp.stats.attention': 'Need attention', 'settings.mcp.status.configured': 'Configured', 'settings.mcp.status.configuredElsewhere': 'Not loaded in the current chat. Open a chat in this project to use it.', @@ -1140,6 +1140,11 @@ Row 9, all 8 cells: continuing from straight down, turning left through lower-le 'settings.mcp.toast.deleteFailed': 'Failed to delete MCP server', 'settings.mcp.toast.toggleFailed': 'Failed to update MCP server state', 'settings.mcp.toast.reconnectFailed': 'Failed to reconnect MCP server', + 'settings.mcp.toast.syncFailed': 'Chat synchronization failed: {error}', + 'settings.mcp.toast.syncNotRunning': 'The current chat is not running; the setting will apply when it starts.', + 'settings.mcp.toast.syncDifferentProject': 'The current chat uses a different project; its tools were not updated.', + 'settings.mcp.toast.syncNoSession': 'No chat was selected; settings were saved.', + 'settings.mcp.toast.syncUnconfirmed': 'The current chat update was not confirmed.', // Settings > Agents 'settings.tab.agents': 'Agents', diff --git a/desktop/src/i18n/locales/jp.ts b/desktop/src/i18n/locales/jp.ts index 7260f3d5..3e080595 100644 --- a/desktop/src/i18n/locales/jp.ts +++ b/desktop/src/i18n/locales/jp.ts @@ -1062,7 +1062,7 @@ export const jp: Record = { 'settings.mcp.empty': 'MCP サーバーはまだ設定されていません', 'settings.mcp.emptyHint': 'カスタムの stdio、HTTP、または SSE の MCP サーバーを追加して、ツールアクセスの拡張を始めましょう。', 'settings.mcp.stats.total': 'サーバー総数', - 'settings.mcp.stats.connected': '現在のチャットで接続中', + 'settings.mcp.stats.connected': '現在のプロジェクトの接続確認', 'settings.mcp.stats.attention': '要対応', 'settings.mcp.status.configured': '設定済み', 'settings.mcp.status.configuredElsewhere': '現在のチャットでは読み込まれていません。このプロジェクトでチャットを開いて使用してください。', @@ -1142,6 +1142,11 @@ export const jp: Record = { 'settings.mcp.toast.deleteFailed': 'MCP サーバーの削除に失敗しました', 'settings.mcp.toast.toggleFailed': 'MCP サーバーの状態の更新に失敗しました', 'settings.mcp.toast.reconnectFailed': 'MCP サーバーの再接続に失敗しました', + 'settings.mcp.toast.syncFailed': 'チャットの同期に失敗しました:{error}', + 'settings.mcp.toast.syncNotRunning': '現在のチャットは実行中ではありません。開始時に設定が適用されます。', + 'settings.mcp.toast.syncDifferentProject': '現在のチャットは別のプロジェクトを使用しているため、ツールは更新されませんでした。', + 'settings.mcp.toast.syncNoSession': 'チャットが選択されていません。設定を保存しました。', + 'settings.mcp.toast.syncUnconfirmed': '現在のチャットの更新を確認できませんでした。', // Settings > Agents 'settings.tab.agents': 'エージェント', diff --git a/desktop/src/i18n/locales/kr.ts b/desktop/src/i18n/locales/kr.ts index 8aff4a27..ebbca181 100644 --- a/desktop/src/i18n/locales/kr.ts +++ b/desktop/src/i18n/locales/kr.ts @@ -1062,7 +1062,7 @@ export const kr: Record = { 'settings.mcp.empty': '아직 구성된 MCP 서버가 없습니다', 'settings.mcp.emptyHint': '사용자 지정 stdio, HTTP 또는 SSE MCP 서버를 추가하여 도구 액세스 확장을 시작하세요.', 'settings.mcp.stats.total': '총 서버', - 'settings.mcp.stats.connected': '현재 채팅에 연결됨', + 'settings.mcp.stats.connected': '현재 프로젝트 연결 확인', 'settings.mcp.stats.attention': '주의 필요', 'settings.mcp.status.configured': '구성됨', 'settings.mcp.status.configuredElsewhere': '현재 채팅에는 로드되지 않았습니다. 이 프로젝트에서 채팅을 열어 사용하세요.', @@ -1142,6 +1142,11 @@ export const kr: Record = { 'settings.mcp.toast.deleteFailed': 'MCP 서버를 삭제하지 못했습니다', 'settings.mcp.toast.toggleFailed': 'MCP 서버 상태를 업데이트하지 못했습니다', 'settings.mcp.toast.reconnectFailed': 'MCP 서버를 다시 연결하지 못했습니다', + 'settings.mcp.toast.syncFailed': '채팅 동기화 실패: {error}', + 'settings.mcp.toast.syncNotRunning': '현재 채팅이 실행 중이 아닙니다. 시작 시 설정이 적용됩니다.', + 'settings.mcp.toast.syncDifferentProject': '현재 채팅은 다른 프로젝트를 사용하므로 도구가 업데이트되지 않았습니다.', + 'settings.mcp.toast.syncNoSession': '선택한 채팅이 없습니다. 설정을 저장했습니다.', + 'settings.mcp.toast.syncUnconfirmed': '현재 채팅 업데이트가 확인되지 않았습니다.', // Settings > Agents 'settings.tab.agents': '에이전트', diff --git a/desktop/src/i18n/locales/zh-TW.ts b/desktop/src/i18n/locales/zh-TW.ts index 4e758f35..8c52295d 100644 --- a/desktop/src/i18n/locales/zh-TW.ts +++ b/desktop/src/i18n/locales/zh-TW.ts @@ -1061,7 +1061,7 @@ export const zh: Record = { 'settings.mcp.empty': '還沒有配置 MCP 服務', 'settings.mcp.emptyHint': '先新增一個自定義的 STDIO、HTTP 或 SSE MCP 服務。', 'settings.mcp.stats.total': '服務總數', - 'settings.mcp.stats.connected': '目前聊天已連線', + 'settings.mcp.stats.connected': '目前專案連線檢查', 'settings.mcp.stats.attention': '需要處理', 'settings.mcp.status.configured': '已設定', 'settings.mcp.status.configuredElsewhere': '目前聊天未載入;請在這個專案中開啟聊天後使用。', @@ -1141,6 +1141,11 @@ export const zh: Record = { 'settings.mcp.toast.deleteFailed': '刪除 MCP 服務失敗', 'settings.mcp.toast.toggleFailed': '更新 MCP 服務狀態失敗', 'settings.mcp.toast.reconnectFailed': '重連 MCP 服務失敗', + 'settings.mcp.toast.syncFailed': '聊天同步失敗:{error}', + 'settings.mcp.toast.syncNotRunning': '目前聊天未執行;設定將在聊天啟動時套用。', + 'settings.mcp.toast.syncDifferentProject': '目前聊天屬於其他專案,其工具未更新。', + 'settings.mcp.toast.syncNoSession': '未選擇聊天;設定已儲存。', + 'settings.mcp.toast.syncUnconfirmed': '尚未確認目前聊天同步成功。', // Settings > Agents 'settings.tab.agents': 'Agents', diff --git a/desktop/src/i18n/locales/zh.ts b/desktop/src/i18n/locales/zh.ts index 4f71d1ba..10ceaa7b 100644 --- a/desktop/src/i18n/locales/zh.ts +++ b/desktop/src/i18n/locales/zh.ts @@ -1061,7 +1061,7 @@ export const zh: Record = { 'settings.mcp.empty': '还没有配置 MCP 服务', 'settings.mcp.emptyHint': '先添加一个自定义的 STDIO、HTTP 或 SSE MCP 服务。', 'settings.mcp.stats.total': '服务总数', - 'settings.mcp.stats.connected': '当前聊天已连接', + 'settings.mcp.stats.connected': '当前项目连接检查', 'settings.mcp.stats.attention': '需要处理', 'settings.mcp.status.configured': '已配置', 'settings.mcp.status.configuredElsewhere': '当前聊天未加载;请在这个项目中打开聊天后使用。', @@ -1141,6 +1141,11 @@ export const zh: Record = { 'settings.mcp.toast.deleteFailed': '删除 MCP 服务失败', 'settings.mcp.toast.toggleFailed': '更新 MCP 服务状态失败', 'settings.mcp.toast.reconnectFailed': '重连 MCP 服务失败', + 'settings.mcp.toast.syncFailed': '聊天同步失败:{error}', + 'settings.mcp.toast.syncNotRunning': '当前聊天未运行;设置将在聊天启动时应用。', + 'settings.mcp.toast.syncDifferentProject': '当前聊天属于其他项目,其工具未更新。', + 'settings.mcp.toast.syncNoSession': '未选择聊天;设置已保存。', + 'settings.mcp.toast.syncUnconfirmed': '尚未确认当前聊天同步成功。', // Settings > Agents 'settings.tab.agents': 'Agents', diff --git a/desktop/src/pages/McpSettings.tsx b/desktop/src/pages/McpSettings.tsx index 31bc23d6..da7cefe7 100644 --- a/desktop/src/pages/McpSettings.tsx +++ b/desktop/src/pages/McpSettings.tsx @@ -599,7 +599,27 @@ export function McpSettings() { const handleToggle = async (server: McpServerRecord) => { setBusyServerKey(getMcpServerIdentityKey(server)) try { - const updated = await toggleServer(server, resolveOperationCwd(server), activeSessionId ?? undefined) + const { server: updated, sessionSync } = await toggleServer(server, resolveOperationCwd(server), activeSessionId ?? undefined) + if (!sessionSync?.applied) { + const detail = sessionSync?.reason === 'failed' + ? t('settings.mcp.toast.syncFailed', { error: sessionSync.error || t('settings.mcp.toast.toggleFailed') }) + : sessionSync?.reason === 'not_running' + ? t('settings.mcp.toast.syncNotRunning') + : sessionSync?.reason === 'different_project' + ? t('settings.mcp.toast.syncDifferentProject') + : sessionSync?.reason === 'no_session' || !activeSessionId + ? t('settings.mcp.toast.syncNoSession') + : t('settings.mcp.toast.syncUnconfirmed') + addToast({ + type: 'warning', + message: `${t('settings.mcp.toast.saved', { name: server.name })}. ${detail}`, + }) + return + } + if (updated.enabled && (updated.status === 'failed' || updated.status === 'needs-auth')) { + addToast({ type: 'warning', message: updated.statusDetail || updated.statusLabel }) + return + } addToast({ type: 'success', message: updated.enabled ? t('settings.mcp.toast.enabled', { name: server.name }) : t('settings.mcp.toast.disabled', { name: server.name }), diff --git a/desktop/src/stores/mcpStore.ts b/desktop/src/stores/mcpStore.ts index 1062cf78..6328faa7 100644 --- a/desktop/src/stores/mcpStore.ts +++ b/desktop/src/stores/mcpStore.ts @@ -7,7 +7,7 @@ import { isSameMcpServer, mcpProjectPathKey, } from '../lib/mcpIdentity' -import type { McpServerRecord, McpUpsertPayload } from '../types/mcp' +import type { McpServerRecord, McpToggleResult, McpUpsertPayload } from '../types/mcp' type McpStore = { servers: McpServerRecord[] @@ -19,7 +19,7 @@ type McpStore = { createServer: (name: string, payload: McpUpsertPayload, cwd?: string) => Promise updateServer: (server: McpServerRecord, payload: McpUpsertPayload, cwd?: string) => Promise deleteServer: (server: McpServerRecord, cwd?: string) => Promise - toggleServer: (server: McpServerRecord, cwd?: string, sessionId?: string) => Promise + toggleServer: (server: McpServerRecord, cwd?: string, sessionId?: string) => Promise reconnectServer: (server: McpServerRecord, cwd?: string) => Promise refreshServerStatus: (server: McpServerRecord, cwd?: string) => Promise selectServer: (server: McpServerRecord | null) => void @@ -214,7 +214,7 @@ export const useMcpStore = create((set, get) => ({ selectedServer: state.selectedServer && isSameMcpServer(state.selectedServer, server) ? updated : state.selectedServer, error: null, })) - return updated + return { ...response, server: updated } }, reconnectServer: async (server, cwd) => { @@ -232,7 +232,12 @@ export const useMcpStore = create((set, get) => ({ }, refreshServerStatus: async (server, cwd) => { + const snapshot = get().servers.find((item) => isSameMcpServer(item, server)) const response = await mcpApi.status(server.name, cwd) + const current = get().servers.find((item) => isSameMcpServer(item, server)) + // A toggle, edit, or refresh completed while this probe was in flight. + // Keep its newer state instead of restoring an old enabled connection. + if (current !== snapshot) return current ?? server const updated = preserveCurrentContextActivity( attachProjectPath(response.server, cwd ?? server.projectPath), server, diff --git a/desktop/src/types/mcp.ts b/desktop/src/types/mcp.ts index ec5f5cf9..9eb069e4 100644 --- a/desktop/src/types/mcp.ts +++ b/desktop/src/types/mcp.ts @@ -44,3 +44,14 @@ export type McpUpsertPayload = { scope: McpWritableScope config: McpEditableConfig } + +export type McpSessionSync = { + applied: boolean + reason?: 'not_running' | 'different_project' | 'failed' | 'no_session' + error?: string +} + +export type McpToggleResult = { + server: McpServerRecord + sessionSync?: McpSessionSync +} diff --git a/src/cli/print.mcpReconnect.test.ts b/src/cli/print.mcpReconnect.test.ts index 3d09d500..f2c39163 100644 --- a/src/cli/print.mcpReconnect.test.ts +++ b/src/cli/print.mcpReconnect.test.ts @@ -14,8 +14,16 @@ process.env.ANTHROPIC_API_KEY = 'test-key' const mcpClient = { ...await import('../services/mcp/client.js') } const mcpConfig = { ...await import('../services/mcp/config.js') } +const contextModule = { ...await import('../commands/context/context-noninteractive.js') } +mock.module('../commands/context/context-noninteractive.js', () => ({ + ...contextModule, + collectContextData: async ({ options }: { options: { tools: Tool[] } }) => ({ toolNames: options.tools.map(tool => tool.name) }), +})) let isDisabled = false +let sdkSetupEnabled = false +let resolveSdkSetup: (() => void) | undefined +const sdkCleanup = mock(async () => {}) let hasConfig = true let resolveReconnect: ((result: ReturnType) => void) | undefined const cleanup = mock(async () => {}) @@ -69,6 +77,13 @@ mock.module('../services/mcp/client.js', () => ({ resolveReconnect = resolve }), clearServerCache, + setupSdkMcpClients: async () => { + await new Promise(resolve => { resolveSdkSetup = resolve }) + return { + clients: [{ ...connectedClient(), cleanup: sdkCleanup }], + tools: [{ name: 'unprefixed-sdk-tool', mcpInfo: { serverName: 'test-server', toolName: 'sdk' } } as Tool], + } + }, })) mock.module('../services/mcp/config.js', () => ({ @@ -81,6 +96,7 @@ mock.module('../services/mcp/config.js', () => ({ } : undefined, isMcpServerDisabled: () => isDisabled, + isMcpServerDisabledForExecution: () => isDisabled, setMcpServerEnabled: (_name: string, enabled: boolean) => { isDisabled = !enabled }, @@ -88,7 +104,7 @@ mock.module('../services/mcp/config.js', () => ({ const { __runHeadlessStreamingForTests } = await import('./print.js') -function startHeadless(input: Stream, initialClient?: MCPServerConnection) { +function startHeadless(input: Stream, initialClient?: MCPServerConnection, initialTools: Tool[] = []) { const io = new StructuredIO(input) let state = getDefaultAppState() state = { @@ -107,7 +123,7 @@ function startHeadless(input: Stream, initialClient?: MCPServerConnectio config: { type: 'stdio', command: 'other' }, }, ], - tools: [{ name: 'mcp__test-server__old' } as Tool], + tools: [{ name: 'mcp__test-server__old', isReadOnly: () => true } as Tool], commands: [ { name: 'mcp__test-server__old', description: '', argumentHint: '' }, ], @@ -118,10 +134,10 @@ function startHeadless(input: Stream, initialClient?: MCPServerConnectio io, [], [], - [], + initialTools, [], (() => undefined) as unknown as CanUseToolFn, - {}, + sdkSetupEnabled ? { 'test-server': { type: 'sdk', name: 'test-server' } } : {}, () => state, update => { state = update(state) @@ -129,7 +145,7 @@ function startHeadless(input: Stream, initialClient?: MCPServerConnectio [], { outputFormat: 'stream-json' }, ) - return { io, output, getState: () => state } + return { io, output, getState: () => state, setState: (next: typeof state) => { state = next } } } async function nextControlResponse(output: AsyncIterable) { @@ -142,6 +158,9 @@ async function nextControlResponse(output: AsyncIterable) { afterEach(() => { hasConfig = true isDisabled = false + sdkSetupEnabled = false + resolveSdkSetup = undefined + sdkCleanup.mockClear() resolveReconnect = undefined clearServerCache.mockClear() cleanup.mockClear() @@ -150,6 +169,7 @@ afterEach(() => { afterAll(() => { mock.module('../services/mcp/client.js', () => mcpClient) mock.module('../services/mcp/config.js', () => mcpConfig) + mock.module('../commands/context/context-noninteractive.js', () => contextModule) if (originalAnthropicApiKey === undefined) { delete process.env.ANTHROPIC_API_KEY } else { @@ -392,3 +412,107 @@ test.each(['mcp_reconnect', 'mcp_toggle'] as const)( } }, ) + +async function controlReader(output: AsyncIterable) { + const iterator = output[Symbol.asyncIterator]() + return async () => { + while (true) { + const { value, done } = await iterator.next() + if (done) throw new Error('Missing control response') + if ((value as { type?: string }).type === 'control_response') return value + } + } +} + +function enqueueControl(input: Stream, requestId: string, request: Record) { + input.enqueue(`${JSON.stringify({ type: 'control_request', request_id: requestId, request })}\n`) +} + +test('removes startup and dynamic tools after disable and restores only fresh tools on enable', async () => { + const input = new Stream() + const { output } = startHeadless(input, connectedClient(), [{ name: 'mcp__test-server__startup' } as Tool]) + const next = await controlReader(output) + try { + enqueueControl(input, 'reconnect', { subtype: 'mcp_reconnect', serverName: 'test-server' }) + await Bun.sleep(0) + resolveReconnect?.(reconnectResult(connectedClient(), true)) + await next() + enqueueControl(input, 'disable', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: false }) + await next() + enqueueControl(input, 'disabled-pool', { subtype: 'get_context_usage', estimateOnly: true }) + const disabled = await next() as { response: { response: { toolNames: string[] } } } + expect(disabled.response.response.toolNames.filter(name => name.startsWith('mcp__test-server__'))).toEqual([]) + enqueueControl(input, 'enable', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: true }) + await Bun.sleep(0) + resolveReconnect?.(reconnectResult(connectedClient(), true)) + await next() + enqueueControl(input, 'enabled-pool', { subtype: 'get_context_usage', estimateOnly: true }) + const enabled = await next() as { response: { response: { toolNames: string[] } } } + expect(enabled.response.response.toolNames.filter(name => name.startsWith('mcp__test-server__'))).toEqual(['mcp__test-server__lookup']) + } finally { input.done() } +}) + +test('invalidates a pending connection when disabling without a connected client', async () => { + const input = new Stream() + const { output } = startHeadless(input, { name: 'test-server', type: 'pending', config: connectedClient().config }) + enqueueControl(input, 'disable-pending', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: false }) + await nextControlResponse(output) + expect(clearServerCache).toHaveBeenCalledTimes(1) + input.done() +}) + +test('does not overwrite a newer persisted disable with a delayed enable control', async () => { + isDisabled = true + const input = new Stream() + const { output, getState } = startHeadless(input) + enqueueControl(input, 'stale-enable', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: true, alreadyPersisted: true }) + await Bun.sleep(0) + // Resolve the buggy reconnect so this assertion fails without timing out. + resolveReconnect?.(reconnectResult(connectedClient(), true)) + await nextControlResponse(output) + expect(isDisabled).toBe(true) + expect(resolveReconnect).toBeUndefined() + expect(getState().mcp.clients[0]?.type).toBe('disabled') + input.done() +}) + +test('filters a late initialization result from tool assembly and status after persisted disable', async () => { + const input = new Stream() + const { output, getState, setState } = startHeadless(input, connectedClient()) + const next = await controlReader(output) + const lateState = getState() + enqueueControl(input, 'disable-before-init', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: false }) + await next() + setState(lateState) + enqueueControl(input, 'late-pool', { subtype: 'get_context_usage', estimateOnly: true }) + const pool = await next() as { response: { response: { toolNames: string[] } } } + expect(pool.response.response.toolNames.filter(name => name.startsWith('mcp__test-server__'))).toEqual([]) + enqueueControl(input, 'late-status', { subtype: 'mcp_status' }) + expect(await next()).toMatchObject({ response: { response: { mcpServers: [ + { name: 'test-server', status: 'disabled' }, { name: 'other-server' }, + ] } } }) + input.done() +}) + +test.each([false, true])('cleans SDK tools and transports when initialization finishes after disable: %s', async late => { + sdkSetupEnabled = true + const input = new Stream() + const { output, getState } = startHeadless(input, connectedClient()) + const next = await controlReader(output) + if (!late) { + resolveSdkSetup?.() + await Bun.sleep(0) + } + enqueueControl(input, 'disable-sdk', { subtype: 'mcp_toggle', serverName: 'test-server', enabled: false }) + await next() + if (late) { + resolveSdkSetup?.() + await Bun.sleep(0) + } + expect(sdkCleanup).toHaveBeenCalledTimes(1) + expect(getState().mcp.tools.some(tool => tool.mcpInfo?.serverName === 'test-server')).toBe(false) + enqueueControl(input, 'sdk-pool', { subtype: 'get_context_usage', estimateOnly: true }) + const pool = await next() as { response: { response: { toolNames: string[] } } } + expect(pool.response.response.toolNames).not.toContain('unprefixed-sdk-tool') + input.done() +}) diff --git a/src/cli/print.ts b/src/cli/print.ts index f70fb7b3..26f2b323 100644 --- a/src/cli/print.ts +++ b/src/cli/print.ts @@ -235,7 +235,7 @@ import { import { filterMcpServersByPolicy, getMcpConfigByName, - isMcpServerDisabled, + isMcpServerDisabledForExecution, setMcpServerEnabled, } from 'src/services/mcp/config.js' import { @@ -1456,12 +1456,24 @@ function runHeadlessStreaming( // Re-initialize all SDK MCP servers with current config const sdkSetup = await setupSdkMcpClients( - sdkMcpConfigs, + Object.fromEntries(Object.entries(sdkMcpConfigs).filter(([name]) => + !isMcpServerDisabledForExecution(name), + )), (serverName, message) => structuredIO.sendMcpMessage(serverName, message), ) - sdkClients = sdkSetup.clients - sdkTools = sdkSetup.tools + const disabledSdkNames = new Set(sdkSetup.clients + .filter(client => isMcpServerDisabledForExecution(client.name)) + .map(client => client.name)) + sdkClients = sdkSetup.clients.map(client => { + if (!disabledSdkNames.has(client.name)) return client + if (client.type === 'connected') void client.cleanup() + return { name: client.name, type: 'disabled' as const, config: client.config } + }) + sdkTools = sdkSetup.tools.filter(tool => + !(tool.mcpInfo && disabledSdkNames.has(tool.mcpInfo.serverName)) && + ![...disabledSdkNames].some(name => tool.name.startsWith(getMcpPrefix(name))), + ) // Store SDK MCP tools in appState so subagents can access them via // assembleToolPool. Only tools are stored here — SDK clients are already @@ -1516,6 +1528,20 @@ function runHeadlessStreaming( ), 'name', ) + const serverNames = uniq([ + ...[ + ...mcpClients, ...sdkClients, ...dynamicMcpState.clients, ...appState.mcp.clients, + ].map(client => client.name), + ...allTools.flatMap(tool => tool.mcpInfo ? [tool.mcpInfo.serverName] : []), + ]) + const disabledServers = new Set( + serverNames.filter(name => isMcpServerDisabledForExecution(name)), + ) + const disabledPrefixes = [...disabledServers].map(getMcpPrefix) + allTools = allTools.filter(tool => + !(tool.mcpInfo && disabledServers.has(tool.mcpInfo.serverName)) && + !disabledPrefixes.some(prefix => tool.name.startsWith(prefix)), + ) if (options.permissionPromptToolName) { allTools = allTools.filter( tool => !toolMatchesName(tool, options.permissionPromptToolName!), @@ -1656,7 +1682,10 @@ function runHeadlessStreaming( ...currentMcpClients, ...sdkClients, ...dynamicMcpState.clients.filter(c => !existingNames.has(c.name)), - ].map(connection => { + ].map(current => { + const connection = isMcpServerDisabledForExecution(current.name) + ? { name: current.name, type: 'disabled' as const, config: current.config } + : current let config if ( connection.config.type === 'sse' || @@ -3152,7 +3181,7 @@ function runHeadlessStreaming( const result = await reconnectMcpServerImpl(serverName, config) // If the server was disabled while the reconnect was in flight, // close any fresh connection and keep the disabled state - if (isMcpServerDisabled(serverName)) { + if (isMcpServerDisabledForExecution(serverName)) { if (result.client.type === 'connected') { void result.client.cleanup() } @@ -3251,7 +3280,12 @@ function runHeadlessStreaming( } } else if (message.request.subtype === 'mcp_toggle') { const currentAppState = getAppState() - const { serverName, enabled } = message.request + const { serverName, alreadyPersisted } = message.request + // API requests already wrote project settings. A delayed control + // request must apply current state rather than undo a newer toggle. + const enabled = alreadyPersisted + ? !isMcpServerDisabledForExecution(serverName) + : message.request.enabled elicitationRegistered.delete(serverName) // Gate must match the client-lookup spread below (which // includes sdkClients and dynamicMcpState.clients). Same fix as @@ -3267,6 +3301,18 @@ function runHeadlessStreaming( const markDisabled = (cfg: NonNullable) => { const prefix = getMcpPrefix(serverName) + const keepTool = (tool: Tool) => + tool.mcpInfo?.serverName !== serverName && !tool.name?.startsWith(prefix) + tools = tools.filter(keepTool) + sdkTools = sdkTools.filter(keepTool) + const disabled = { name: serverName, type: 'disabled' as const, config: cfg } + mcpClients = mcpClients.map(client => client.name === serverName ? disabled : client) + sdkClients = sdkClients.map(client => client.name === serverName ? disabled : client) + dynamicMcpState = { + ...dynamicMcpState, + clients: [...dynamicMcpState.clients.filter(client => client.name !== serverName), disabled], + tools: dynamicMcpState.tools.filter(keepTool), + } setAppState(prev => ({ ...prev, mcp: { @@ -3276,7 +3322,7 @@ function runHeadlessStreaming( ? { name: serverName, type: 'disabled' as const, config: cfg } : c, ), - tools: reject(prev.mcp.tools, t => t.name?.startsWith(prefix)), + tools: prev.mcp.tools.filter(keepTool), commands: reject(prev.mcp.commands, c => commandBelongsToServer(c, serverName), ), @@ -3289,25 +3335,23 @@ function runHeadlessStreaming( sendControlResponseError(message, `Server not found: ${serverName}`) } else if (!enabled) { // Disabling: persist + disconnect (matches TUI toggleMcpServer behavior) - setMcpServerEnabled(serverName, false) - const client = [ - ...mcpClients, - ...sdkClients, - ...dynamicMcpState.clients, - ...currentAppState.mcp.clients, - ].find(c => c.name === serverName) + if (!alreadyPersisted) setMcpServerEnabled(serverName, false) + const sdkConnections = sdkClients.filter(client => client.name === serverName) markDisabled(config) - if (client && client.type === 'connected') { - await clearServerCache(serverName, config) - } + // Pending transports also need invalidation: they may finish their + // handshake after the disabled state has already been published. + await Promise.all([ + clearServerCache(serverName, config), + ...sdkConnections.map(client => client.type === 'connected' ? client.cleanup() : undefined), + ]) sendControlResponseSuccess(message) } else { // Enabling: persist + reconnect - setMcpServerEnabled(serverName, true) + if (!alreadyPersisted) setMcpServerEnabled(serverName, true) const result = await reconnectMcpServerImpl(serverName, config) // If the server was disabled while the reconnect was in flight, // close any fresh connection and keep the disabled state - if (isMcpServerDisabled(serverName)) { + if (isMcpServerDisabledForExecution(serverName)) { if (result.client.type === 'connected') { void result.client.cleanup() } @@ -3345,6 +3389,18 @@ function runHeadlessStreaming( : omit(prev.mcp.resources, serverName), }, })) + dynamicMcpState = { + ...dynamicMcpState, + clients: [...dynamicMcpState.clients.filter(client => client.name !== serverName), result.client], + tools: [ + ...dynamicMcpState.tools.filter(tool => + tool.mcpInfo?.serverName !== serverName && !tool.name?.startsWith(prefix), + ), + ...result.tools, + ], + } + mcpClients = mcpClients.map(client => client.name === serverName ? result.client : client) + sdkClients = sdkClients.map(client => client.name === serverName ? result.client : client) if (result.client.type === 'connected') { registerElicitationHandlers([result.client]) reregisterChannelHandlerAfterReconnect(result.client) @@ -3443,7 +3499,7 @@ function runHeadlessStreaming( const fullFlowPromise = oauthPromise .then(async () => { // Don't reconnect if the server was disabled during the OAuth flow - if (isMcpServerDisabled(serverName)) { + if (isMcpServerDisabledForExecution(serverName)) { return } // Skip reconnect if the manual callback path was used — @@ -3460,7 +3516,7 @@ function runHeadlessStreaming( // If the server was disabled while the reconnect was in // flight, close any fresh connection and keep the disabled // state - if (isMcpServerDisabled(serverName)) { + if (isMcpServerDisabledForExecution(serverName)) { setAppState(prev => { const disabled = disableStaleMcpReconnect( serverName, diff --git a/src/entrypoints/sdk/controlSchemas.mcpToggle.test.ts b/src/entrypoints/sdk/controlSchemas.mcpToggle.test.ts new file mode 100644 index 00000000..5e46bf9b --- /dev/null +++ b/src/entrypoints/sdk/controlSchemas.mcpToggle.test.ts @@ -0,0 +1,16 @@ +import { expect, test } from 'bun:test' +import { SDKControlRequestSchema } from './controlSchemas.js' + +test.each([undefined, true, false])('preserves MCP toggle persistence ownership: %s', alreadyPersisted => { + const request = { + type: 'control_request', + request_id: 'toggle', + request: { + subtype: 'mcp_toggle', + serverName: 'local-fixture', + enabled: false, + ...(alreadyPersisted === undefined ? {} : { alreadyPersisted }), + }, + } + expect(SDKControlRequestSchema().parse(request)).toEqual(request) +}) diff --git a/src/entrypoints/sdk/controlSchemas.ts b/src/entrypoints/sdk/controlSchemas.ts index 5fe10e35..50ff0729 100644 --- a/src/entrypoints/sdk/controlSchemas.ts +++ b/src/entrypoints/sdk/controlSchemas.ts @@ -511,6 +511,7 @@ export const SDKControlMcpToggleRequestSchema = lazySchema(() => subtype: z.literal('mcp_toggle'), serverName: z.string(), enabled: z.boolean(), + alreadyPersisted: z.boolean().optional(), }) .describe('Enables or disables an MCP server.'), ) diff --git a/src/server/__tests__/mcp.test.ts b/src/server/__tests__/mcp.test.ts index 8f62480b..d5a66d2b 100644 --- a/src/server/__tests__/mcp.test.ts +++ b/src/server/__tests__/mcp.test.ts @@ -1,4 +1,6 @@ +import '../../../preload.ts' import { afterEach, beforeEach, describe, expect, it, mock, spyOn } from 'bun:test' +import { Client } from '@modelcontextprotocol/sdk/client/index.js' import * as fs from 'fs/promises' import * as os from 'os' import * as path from 'path' @@ -7,6 +9,7 @@ import * as mcpClient from '../../services/mcp/client.js' import * as mcpConfig from '../../services/mcp/config.js' import { _setGlobalConfigCacheForTesting, getProjectPathForConfig } from '../../utils/config.js' import { getGlobalClaudeFile } from '../../utils/env.js' +import { runWithCwdOverride } from '../../utils/cwd.js' import { normalizePathForConfigKey } from '../../utils/path.js' import * as mcpHostPreflight from '../services/mcpHostPreflight.js' import { handleMcpApi } from '../api/mcp.js' @@ -23,6 +26,7 @@ let reconnectSpy: ReturnType | undefined let hostPreflightSpy: ReturnType | undefined let originalRequestControl: typeof conversationService.requestControl let originalHasSession: typeof conversationService.hasSession +let originalGetActiveSessions: typeof conversationService.getActiveSessions let originalGetSessionWorkDir: typeof conversationService.getSessionWorkDir function clearConfigPathCaches() { @@ -82,6 +86,7 @@ describe('MCP API', () => { originalRequestControl = conversationService.requestControl.bind(conversationService) originalHasSession = conversationService.hasSession.bind(conversationService) originalGetSessionWorkDir = conversationService.getSessionWorkDir.bind(conversationService) + originalGetActiveSessions = conversationService.getActiveSessions.bind(conversationService) connectSpy = spyOn(mcpClient, 'connectToServer').mockImplementation(async (name, config) => ({ name, @@ -109,6 +114,7 @@ describe('MCP API', () => { conversationService.requestControl = originalRequestControl conversationService.hasSession = originalHasSession conversationService.getSessionWorkDir = originalGetSessionWorkDir + conversationService.getActiveSessions = originalGetActiveSessions await teardown() }) @@ -468,6 +474,66 @@ describe('MCP API', () => { }) }) + it('retries status after a failed connection instead of reusing its cached failure', async () => { + connectSpy?.mockRestore() + connectSpy = undefined + const connect = spyOn(Client.prototype, 'connect') + .mockRejectedValueOnce(new Error('fixture temporarily unavailable')) + .mockResolvedValue(undefined) + try { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'retry-probe', scope: 'local', + config: { type: 'sse', url: 'http://127.0.0.1:1/mcp' }, + }) + await handleMcpApi(create.req, create.url, create.segments) + const status = makeRequest('GET', `/api/mcp/retry-probe/status?cwd=${encodeURIComponent(projectRoot)}`) + expect((await (await handleMcpApi(status.req, status.url, status.segments)).json()).server.status).toBe('failed') + expect((await (await handleMcpApi(status.req, status.url, status.segments)).json()).server.status).toBe('connected') + expect(connect).toHaveBeenCalledTimes(2) + } finally { + await runWithCwdOverride(projectRoot, () => mcpClient.clearServerCache('retry-probe', mcpConfig.getMcpConfigByName('retry-probe')!)) + connect.mockRestore() + } + }) + + it('does not discard a replacement connection when an older status probe fails', async () => { + connectSpy?.mockRestore() + connectSpy = undefined + let failOld!: (error: Error) => void + let markStarted!: () => void + const started = new Promise(resolve => { markStarted = resolve }) + const connect = spyOn(Client.prototype, 'connect') + .mockImplementationOnce(() => { + markStarted() + return new Promise((_resolve, reject) => { failOld = reject }) + }) + .mockResolvedValue(undefined) + let config: NonNullable> + try { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'replaced-probe', scope: 'local', + config: { type: 'sse', url: 'http://127.0.0.1:1/mcp' }, + }) + await handleMcpApi(create.req, create.url, create.segments) + config = runWithCwdOverride(projectRoot, () => mcpConfig.getMcpConfigByName('replaced-probe')!) + const status = makeRequest('GET', `/api/mcp/replaced-probe/status?cwd=${encodeURIComponent(projectRoot)}`) + const probing = handleMcpApi(status.req, status.url, status.segments) + await started + await runWithCwdOverride(projectRoot, () => mcpClient.clearServerCache('replaced-probe', config)) + const replacement = await runWithCwdOverride(projectRoot, () => mcpClient.connectToServer('replaced-probe', config)) + expect(replacement.type).toBe('connected') + failOld(new Error('old probe failed')) + expect((await (await probing).json()).server.status).toBe('failed') + const retained = await runWithCwdOverride(projectRoot, () => mcpClient.connectToServer('replaced-probe', config)) + expect(retained).toBe(replacement) + expect(connect).toHaveBeenCalledTimes(2) + } finally { + failOld?.(new Error('fixture cleanup')) + await runWithCwdOverride(projectRoot, () => mcpClient.clearServerCache('replaced-probe', mcpConfig.getMcpConfigByName('replaced-probe')!)) + connect.mockRestore() + } + }) + it('checks a single server status on demand', async () => { const create = makeRequest('POST', '/api/mcp', { cwd: projectRoot, @@ -845,11 +911,97 @@ describe('MCP API', () => { expect(disableRes.status).toBe(200) expect(requestControl).toHaveBeenCalledWith( 'session-1', - { subtype: 'mcp_toggle', serverName: 'session-sync', enabled: false }, + { subtype: 'mcp_toggle', serverName: 'session-sync', enabled: false, alreadyPersisted: true }, 120_000, ) }) + it('reports no selected session and clears every running session in the same project', async () => { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'all-sessions', scope: 'local', + config: { type: 'stdio', command: 'mock', args: [], env: {} }, + }) + await handleMcpApi(create.req, create.url, create.segments) + conversationService.getActiveSessions = () => ['first', 'second', 'other'] + conversationService.hasSession = () => true + conversationService.getSessionWorkDir = id => id === 'other' ? tmpDir : projectRoot + const requestControl = mock(async () => ({})) + conversationService.requestControl = requestControl as typeof conversationService.requestControl + const toggle = makeRequest('POST', '/api/mcp/all-sessions/toggle', { cwd: projectRoot }) + const response = await handleMcpApi(toggle.req, toggle.url, toggle.segments) + expect((await response.json()).sessionSync).toEqual({ applied: false, reason: 'no_session' }) + expect(requestControl.mock.calls.map(call => call[0])).toEqual(['first', 'second']) + }) + + it('reports a failed background session sync while keeping the persisted disable', async () => { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'failed-sync', scope: 'local', + config: { type: 'stdio', command: 'mock', args: [], env: {} }, + }) + await handleMcpApi(create.req, create.url, create.segments) + conversationService.getActiveSessions = () => ['first', 'second'] + conversationService.hasSession = () => true + conversationService.getSessionWorkDir = () => projectRoot + conversationService.requestControl = (async (id: string) => { + if (id === 'second') throw new Error('control transport closed') + return {} + }) as typeof conversationService.requestControl + const toggle = makeRequest('POST', '/api/mcp/failed-sync/toggle', { cwd: projectRoot, sessionId: 'first' }) + const response = await handleMcpApi(toggle.req, toggle.url, toggle.segments) + const body = await response.json() + expect(body.server.enabled).toBe(false) + expect(body.sessionSync).toEqual({ applied: false, reason: 'failed', error: 'control transport closed' }) + }) + + it('applies disable immediately while a slow enable probe is pending', async () => { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'slow-toggle', scope: 'local', + config: { type: 'stdio', command: 'mock', args: [], env: {} }, + }) + await handleMcpApi(create.req, create.url, create.segments) + const toggle = makeRequest('POST', '/api/mcp/slow-toggle/toggle', { cwd: projectRoot }) + await handleMcpApi(toggle.req, toggle.url, toggle.segments) + let release!: () => void + let probeStarted!: () => void + const started = new Promise(resolve => { probeStarted = resolve }) + const pending = new Promise(resolve => { release = resolve }) + connectSpy!.mockImplementation(async (name, config) => { + probeStarted() + await pending + return { name, type: 'connected', config, client: {} as never, capabilities: {}, cleanup: async () => {} } + }) + const enable = makeRequest('POST', '/api/mcp/slow-toggle/toggle', { cwd: projectRoot }) + const enabling = handleMcpApi(enable.req, enable.url, enable.segments) + await started + const disable = makeRequest('POST', '/api/mcp/slow-toggle/toggle', { cwd: projectRoot }) + const disabling = handleMcpApi(disable.req, disable.url, disable.segments) + try { + const disabled = await Promise.race([disabling, Bun.sleep(100).then(() => null)]) + expect(disabled).not.toBeNull() + expect((await disabled!.json()).server).toMatchObject({ enabled: false, status: 'disabled' }) + } finally { release() } + expect((await (await enabling).json()).server).toMatchObject({ enabled: false, status: 'disabled' }) + expect(connectSpy).toHaveBeenCalledTimes(1) + }) + + it('retains the not-running sync receipt when enable preflight fails', async () => { + const create = makeRequest('POST', '/api/mcp', { + cwd: projectRoot, name: 'preflight-sync', scope: 'local', + config: { type: 'stdio', command: 'mock', args: [], env: {} }, + }) + await handleMcpApi(create.req, create.url, create.segments) + conversationService.hasSession = () => false + const disable = makeRequest('POST', '/api/mcp/preflight-sync/toggle', { cwd: projectRoot }) + await handleMcpApi(disable.req, disable.url, disable.segments) + hostPreflightSpy!.mockResolvedValue({ ok: false, message: 'Missing mock executable' }) + const enable = makeRequest('POST', '/api/mcp/preflight-sync/toggle', { cwd: projectRoot, sessionId: 'stopped' }) + const response = await handleMcpApi(enable.req, enable.url, enable.segments) + expect(await response.json()).toMatchObject({ + server: { enabled: true, status: 'failed' }, + sessionSync: { applied: false, reason: 'not_running' }, + }) + }) + it('does not sync a project-specific toggle into a session from another project', async () => { const otherProject = path.join(tmpDir, 'other-project') await fs.mkdir(otherProject, { recursive: true }) diff --git a/src/server/api/mcp.ts b/src/server/api/mcp.ts index 8d49ed83..714656cf 100644 --- a/src/server/api/mcp.ts +++ b/src/server/api/mcp.ts @@ -6,6 +6,7 @@ import { import { clearServerCache, connectToServer, + getServerCacheKey, reconnectMcpServerImpl, } from '../../services/mcp/client.js' import { @@ -14,6 +15,7 @@ import { getAllMcpConfigs, getMcpConfigByName, isMcpServerDisabled, + isMcpServerDisabledForExecution, projectDirDeclaresMcpServers, registerCwdProjectIfDeclaresMcpServers, removeMcpConfig, @@ -84,7 +86,7 @@ type McpMutationBody = { type McpSessionSyncDto = { applied: boolean - reason?: 'not_running' | 'different_project' | 'failed' + reason?: 'not_running' | 'different_project' | 'failed' | 'no_session' error?: string } @@ -121,12 +123,18 @@ async function syncMcpToggleToSession( sessionId: string | undefined, server: McpServerIdentity, enabled: boolean, -): Promise { - if (!sessionId) return undefined +): Promise { + if (!sessionId) return { applied: false, reason: 'no_session' } if (!conversationService.hasSession(sessionId)) { return { applied: false, reason: 'not_running' } } + const sessionCwd = conversationService.getSessionWorkDir(sessionId) + if (!sessionCwd || normalizePathForConfigKey(getProjectPathForConfig(sessionCwd)) !== + normalizePathForConfigKey(getProjectPathForConfig(getCwd()))) { + return { applied: false, reason: 'different_project' } + } + if (server.scope === 'local' || server.scope === 'project' || server.scope === 'user') { const sessionWorkDir = conversationService.getSessionWorkDir(sessionId) const sessionServer = sessionWorkDir @@ -144,7 +152,7 @@ async function syncMcpToggleToSession( try { await conversationService.requestControl( sessionId, - { subtype: 'mcp_toggle', serverName: server.name, enabled }, + { subtype: 'mcp_toggle', serverName: server.name, enabled, alreadyPersisted: true }, 120_000, ) return { applied: true } @@ -334,9 +342,13 @@ async function inspectServerStatus( return hostPreflightStatus } + const cacheKey = getServerCacheKey(name, config) + let cachedProbe: ReturnType | undefined try { - const client = await connectToServer(name, config) - await clearServerCache(name, config).catch(() => {}) + const connecting = connectToServer(name, config) + cachedProbe = connectToServer.cache?.get(cacheKey) + const client = await connecting + if (client.type === 'connected') await client.cleanup().catch(() => {}) const status: McpServerDto['status'] = client.type === 'connected' @@ -351,12 +363,17 @@ async function inspectServerStatus( statusDetail: 'error' in client ? client.error : undefined, } } catch (error) { - await clearServerCache(name, config).catch(() => {}) return { status: 'failed', statusLabel: getStatusLabel('failed'), statusDetail: error instanceof Error ? error.message : String(error), } + } finally { + // Failed probes must be retried, but a slow old probe may have been replaced + // by a reconnect. Check its cached promise immediately before detaching it. + if (cachedProbe && connectToServer.cache?.get(cacheKey) === cachedProbe) { + await clearServerCache(name, config).catch(() => {}) + } } } @@ -365,7 +382,8 @@ function buildServerDto( config: ScopedMcpServerConfig, status: Pick, ): McpServerDto { - const enabled = !isMcpServerDisabled(name) + const enabled = !isMcpServerDisabledForExecution(name) + if (!enabled) status = getInitialStatus(false) const transport = config.type ?? 'stdio' const canEdit = EDITABLE_SCOPES.has(config.scope) && (transport === 'stdio' || transport === 'http' || transport === 'sse') @@ -373,7 +391,7 @@ function buildServerDto( name, scope: config.scope, transport, - enabled: !isMcpServerDisabled(name), + enabled, status: status.status, statusLabel: status.statusLabel, statusDetail: status.statusDetail, @@ -648,6 +666,24 @@ async function deleteServer(name: string, url: URL): Promise { return Response.json({ ok: true }) } +// Serialize enable probes so a stale probe's cleanup cannot close the next +// enable's connection. Persisting a disable never waits behind a slow probe. +const enableProbeQueues = new Map>() + +async function syncMcpToggleToSessions( + sessionId: string | undefined, + server: McpServerIdentity, + enabled: boolean, +): Promise { + const sessionIds = [...new Set([ + ...(sessionId ? [sessionId] : []), + ...conversationService.getActiveSessions(), + ])] + const results = await Promise.all(sessionIds.map(id => syncMcpToggleToSession(id, server, enabled))) + const failure = results.find(result => result.reason === 'failed') + return failure ?? (sessionId ? results[sessionIds.indexOf(sessionId)]! : { applied: false, reason: 'no_session' }) +} + async function toggleServer(name: string, sessionId?: string): Promise { const existing = await resolveServerForRuntimeAction(name) if (!existing) { @@ -655,38 +691,31 @@ async function toggleServer(name: string, sessionId?: string): Promise } const serverIdentity = getServerIdentity(name, existing) - const enabled = isMcpServerDisabled(name) + const enabled = isMcpServerDisabledForExecution(name) setMcpServerEnabled(name, enabled) - const sessionSync = await syncMcpToggleToSession(sessionId, serverIdentity, enabled) + if (!enabled) await clearServerCache(name, existing).catch(() => {}) + const sessionSync = await syncMcpToggleToSessions(sessionId, serverIdentity, enabled) if (!enabled) { - await clearServerCache(name, existing).catch(() => {}) const updated = serializeServerSnapshot(name, existing) - return Response.json({ server: updated, ...(sessionSync ? { sessionSync } : {}) }) + return Response.json({ server: updated, sessionSync }) } - const hostPreflightStatus = await getHostPreflightStatus(existing, true) - if (hostPreflightStatus) { + const key = JSON.stringify([getProjectPathForConfig(getCwd()), name]) + const previous = enableProbeQueues.get(key) + const next = (previous ?? Promise.resolve()).catch(() => {}).then(async () => { + if (isMcpServerDisabledForExecution(name)) { + return Response.json({ server: serializeServerSnapshot(name, existing), sessionSync }) + } await clearServerCache(name, existing).catch(() => {}) - return Response.json({ - server: buildServerDto(name, existing, hostPreflightStatus), - }) - } - - const result = await reconnectMcpServerImpl(name, existing) - await clearServerCache(name, existing).catch(() => {}) - - const updated = await serializeServerWithLiveStatus(name, existing) - const statusDetail = - result.client.type === 'failed' && 'error' in result.client ? result.client.error : undefined - - return Response.json({ - server: { - ...updated, - ...(statusDetail ? { statusDetail } : {}), - }, - ...(sessionSync ? { sessionSync } : {}), + const updated = await serializeServerWithLiveStatus(name, existing) + return Response.json({ server: updated, sessionSync }) }) + enableProbeQueues.set(key, next) + void next.finally(() => { + if (enableProbeQueues.get(key) === next) enableProbeQueues.delete(key) + }).catch(() => {}) + return next } async function reconnectServer(name: string): Promise { diff --git a/src/services/mcp/client.disabled.test.ts b/src/services/mcp/client.disabled.test.ts new file mode 100644 index 00000000..12d1e66e --- /dev/null +++ b/src/services/mcp/client.disabled.test.ts @@ -0,0 +1,188 @@ +import '../../../preload.ts' +import { afterEach, beforeEach, describe, expect, test } from 'bun:test' +import { mkdtemp, mkdir, readFile, rm, writeFile } from 'node:fs/promises' +import { join } from 'node:path' +import { tmpdir } from 'node:os' +import { clearServerCache, connectToServer, ensureConnectedClient, fetchToolsForClient } from './client.js' +import { setMcpServerEnabled } from './config.js' +import { getGlobalClaudeFile } from '../../utils/env.js' +import { _setGlobalConfigCacheForTesting, enableConfigs, getProjectPathForConfig } from '../../utils/config.js' +import { runWithCwdOverride } from '../../utils/cwd.js' +import type { ScopedMcpServerConfig } from './types.js' + +let root: string +let config: ScopedMcpServerConfig +const name = 'disabled-stdio-regression' +let previousConfigDir: string | undefined + +function inProject(fn: () => T) { return runWithCwdOverride(root, fn) } +async function instances() { + return (await readFile(join(root, 'instances'), 'utf8').catch(() => '')).trim().split('\n').filter(Boolean) +} +async function waitForInstance(count: number) { + for (let i = 0; i < 200; i++) { + if ((await instances()).length === count) return + await Bun.sleep(10) + } + throw new Error('stdio fixture did not start') +} +async function connected() { + const client = await inProject(() => connectToServer(name, config)) + if (client.type !== 'connected') throw new Error(`Fixture connection failed: ${client.type}`) + return client +} +async function invoke(tool: Awaited>[number]) { + return inProject(() => tool!.call( + { text: 'echo' }, + { abortController: new AbortController(), setAppState: () => {} } as never, + undefined as never, + { message: { content: [] } } as never, + )) +} + +beforeEach(async () => { + root = await mkdtemp(join(tmpdir(), 'qa005-stdio-')) + previousConfigDir = process.env.CLAUDE_CONFIG_DIR + process.env.CLAUDE_CONFIG_DIR = root + getGlobalClaudeFile.cache.clear?.() + getProjectPathForConfig.cache.clear?.() + _setGlobalConfigCacheForTesting(null) + enableConfigs() + await writeFile(join(root, 'server.cjs'), ` +const fs = require('node:fs') +const readline = require('node:readline') +const instance = String(process.pid) +fs.appendFileSync(process.argv[2], instance + '\\n') +readline.createInterface({ input: process.stdin }).on('line', async line => { + const req = JSON.parse(line) + if (req.id === undefined) return + let result + if (req.method === 'initialize') { + while (process.argv[3] && !fs.existsSync(process.argv[3])) await new Promise(r => setTimeout(r, 10)) + result = { protocolVersion: req.params.protocolVersion, capabilities: { tools: {} }, serverInfo: { name: 'fixture', version: '1' } } + } else if (req.method === 'tools/list') result = { tools: [{ name: 'echo', inputSchema: { type: 'object' } }] } + else if (req.method === 'tools/call') result = { content: [{ type: 'text', text: instance }] } + else result = {} + process.stdout.write(JSON.stringify({ jsonrpc: '2.0', id: req.id, result }) + '\\n') +}) +`) + config = { type: 'stdio', command: process.execPath, args: [join(root, 'server.cjs'), join(root, 'instances')], scope: 'project' } +}) + +afterEach(async () => { + await inProject(() => clearServerCache(name, config)) + if (previousConfigDir === undefined) delete process.env.CLAUDE_CONFIG_DIR + else process.env.CLAUDE_CONFIG_DIR = previousConfigDir + getGlobalClaudeFile.cache.clear?.() + getProjectPathForConfig.cache.clear?.() + _setGlobalConfigCacheForTesting(null) + await rm(root, { recursive: true, force: true }) +}) + +describe('disabled MCP execution boundary', () => { + test('old tool closure cannot spawn after disable; re-enable starts a fresh instance', async () => { + const old = await connected() + const [tool] = await fetchToolsForClient(old) + await invoke(tool!) + expect(await instances()).toHaveLength(1) + inProject(() => setMcpServerEnabled(name, false)) + await inProject(() => clearServerCache(name, config)) + await expect(invoke(tool!)).rejects.toThrow('disabled') + expect(await instances()).toHaveLength(1) + inProject(() => setMcpServerEnabled(name, true)) + await invoke(tool!) + expect(await instances()).toHaveLength(2) + expect(new Set(await instances()).size).toBe(2) + }) + + test('cached connection and direct connection entry obey disable without control delivery', async () => { + const old = await connected() + inProject(() => setMcpServerEnabled(name, false)) + await expect(inProject(() => ensureConnectedClient(old))).rejects.toThrow('disabled') + await expect(inProject(() => ensureConnectedClient({ ...old, config: { type: 'sdk', name, scope: 'dynamic' } }))).rejects.toThrow('disabled') + expect((await inProject(() => connectToServer(name, config))).type).toBe('disabled') + expect(await instances()).toHaveLength(1) + }) + + test('a retained closure keeps its project disabled state when invoked from another project', async () => { + const old = await connected() + const [tool] = await fetchToolsForClient(old) + inProject(() => setMcpServerEnabled(name, false)) + const otherProject = join(root, 'other') + await mkdir(otherProject) + await expect(runWithCwdOverride(otherProject, () => ensureConnectedClient(old))).rejects.toThrow('disabled') + await expect(invoke(tool!)).rejects.toThrow('disabled') + expect(await instances()).toHaveLength(1) + }) + + test('same-name tools remain isolated across projects, including cache cleanup', async () => { + const first = await connected() + const firstTools = await fetchToolsForClient(first) + const otherProject = join(root, 'other') + await mkdir(otherProject) + try { + const second = await runWithCwdOverride(otherProject, () => connectToServer(name, config)) + expect(second.type).toBe('connected') + const secondTools = await runWithCwdOverride(otherProject, () => fetchToolsForClient(second)) + expect(secondTools).not.toBe(firstTools) + runWithCwdOverride(otherProject, () => setMcpServerEnabled(name, false)) + await runWithCwdOverride(otherProject, () => clearServerCache(name, config)) + const disabled = await runWithCwdOverride(otherProject, () => connectToServer(name, config)) + expect(await runWithCwdOverride(otherProject, () => fetchToolsForClient(disabled))).toEqual([]) + expect(await fetchToolsForClient(first)).toBe(firstTools) + await expect(invoke(secondTools[0]!)).rejects.toThrow('disabled') + await invoke(firstTools[0]!) + expect(await instances()).toHaveLength(2) + } finally { + await runWithCwdOverride(otherProject, () => clearServerCache(name, config)) + } + }) + + test('disable during slow initialization prevents publishing the connected client', async () => { + const ready = join(root, 'ready') + config = { ...config, args: [...(config as { args: string[] }).args, ready] } as ScopedMcpServerConfig + const pending = inProject(() => connectToServer(name, config)) + await waitForInstance(1) + inProject(() => setMcpServerEnabled(name, false)) + await writeFile(ready, '') + expect((await pending).type).toBe('disabled') + expect(await instances()).toHaveLength(1) + }) + + test('clearing a disabled slow initialization cancels it without waiting for its handshake', async () => { + const ready = join(root, 'ready') + config = { ...config, args: [...(config as { args: string[] }).args, ready] } as ScopedMcpServerConfig + const pending = inProject(() => connectToServer(name, config)) + await waitForInstance(1) + inProject(() => setMcpServerEnabled(name, false)) + const clearing = inProject(() => clearServerCache(name, config)) + let timer: ReturnType | undefined + const result = await Promise.race([ + clearing.then(() => 'cleared'), + new Promise(resolve => { timer = setTimeout(() => resolve('blocked'), 1500) }), + ]) + clearTimeout(timer) + await writeFile(ready, '') + await clearing + await pending + expect(result).toBe('cleared') + expect(await instances()).toHaveLength(1) + }) + + test('an invalidated slow connection cannot win after rapid disable and re-enable', async () => { + const ready = join(root, 'ready') + config = { ...config, args: [...(config as { args: string[] }).args, ready] } as ScopedMcpServerConfig + const pending = inProject(() => connectToServer(name, config)) + await waitForInstance(1) + inProject(() => setMcpServerEnabled(name, false)) + const clearing = inProject(() => clearServerCache(name, config)) + inProject(() => setMcpServerEnabled(name, true)) + const newer = inProject(() => connectToServer(name, config)) + await waitForInstance(2) + await writeFile(ready, '') + expect((await pending).type).not.toBe('connected') + expect((await newer).type).toBe('connected') + await clearing + expect(await inProject(() => connectToServer(name, config))).toBe(await newer) + }) +}) diff --git a/src/services/mcp/client.lifecycle.test.ts b/src/services/mcp/client.lifecycle.test.ts index c8d77ead..cececcbf 100644 --- a/src/services/mcp/client.lifecycle.test.ts +++ b/src/services/mcp/client.lifecycle.test.ts @@ -6,6 +6,7 @@ import { connectToServer, fetchToolsForClient, getServerCacheKey, + getMcpClientCacheKey, setMcpConnectionClosedHandler, } from './client.js' @@ -85,7 +86,7 @@ describe('MCP connection ownership', () => { setMcpConnectionClosedHandler(closed) active.client.onclose?.() expect(connectToServer.cache.has(getServerCacheKey(name, config))).toBe(false) - expect(fetchToolsForClient.cache.has(name)).toBe(false) + expect(fetchToolsForClient.cache.has(getMcpClientCacheKey(active))).toBe(false) expect(closed).toHaveBeenCalledWith(name, active.client) }) diff --git a/src/services/mcp/client.ts b/src/services/mcp/client.ts index 08bd290d..6e725a55 100644 --- a/src/services/mcp/client.ts +++ b/src/services/mcp/client.ts @@ -132,7 +132,8 @@ import { wrapFetchWithStepUpDetection, } from './auth.js' import { markClaudeAiMcpConnected } from './claudeai.js' -import { getAllMcpConfigs, isMcpServerDisabled } from './config.js' +import { getAllMcpConfigs, isMcpServerDisabled, isMcpServerDisabledForExecution } from './config.js' +import { getCwd, runWithCwdOverride } from '../../utils/cwd.js' import { getMcpServerHeaders } from './headersHelper.js' import { SdkControlClientTransport } from './SdkControlTransport.js' import type { @@ -577,7 +578,12 @@ let onMcpConnectionClosed: ((name: string, client: Client) => void) | undefined // A delayed close belongs to its connection attempt, not whichever connection // currently has the same server name (including a changed configuration). -const connectionAttempts = new Map() +type ConnectionAttempt = { + key: string + cancel?: () => Promise + cleanup?: () => Promise +} +const connectionAttempts = new Map() export function setMcpConnectionClosedHandler( handler: ((name: string, client: Client) => void) | undefined, @@ -589,10 +595,11 @@ export function notifyMcpConnectionClosed(name: string, client: Client): void { onMcpConnectionClosed?.(name, client) } -function clearServerFetchCaches(name: string): void { - fetchToolsForClient.cache.delete(name) - fetchResourcesForClient.cache.delete(name) - fetchCommandsForClient.cache.delete(name) +function clearServerFetchCaches(name: string, connectionKey: string): void { + const key = `${connectionKey}-connected` + fetchToolsForClient.cache.delete(key) + fetchResourcesForClient.cache.delete(key) + fetchCommandsForClient.cache.delete(key) if (feature('MCP_SKILLS')) { fetchMcpSkillsForClient!.cache.delete(name) } @@ -608,7 +615,7 @@ export function getServerCacheKey( name: string, serverRef: ScopedMcpServerConfig, ): string { - return `${name}-${jsonStringify(serverRef)}` + return jsonStringify([getCwd(), name, serverRef]) } /** @@ -618,7 +625,16 @@ export function getServerCacheKey( * @param serverRef Scoped server configuration * @returns A wrapped client (either connected or failed) */ -export const connectToServer = memoize( +const connectionProjects = new WeakMap() + +export function getMcpClientCacheKey(client: MCPServerConnection): string { + const cwd = client.type === 'connected' ? connectionProjects.get(client.client) : undefined + return runWithCwdOverride(cwd ?? getCwd(), () => + `${getServerCacheKey(client.name, client.config)}-${client.type}`, + ) +} + +const connectToServerMemoized = memoize( async ( name: string, serverRef: ScopedMcpServerConfig, @@ -632,8 +648,8 @@ export const connectToServer = memoize( }, ): Promise => { const connectStartTime = Date.now() - const attempt = { key: getServerCacheKey(name, serverRef) } - connectionAttempts.set(name, attempt) + const attempt: ConnectionAttempt = { key: getServerCacheKey(name, serverRef) } + connectionAttempts.set(attempt.key, attempt) let inProcessServer: | { connect(t: Transport): Promise; close(): Promise } | undefined @@ -1070,6 +1086,20 @@ export const connectToServer = memoize( } } + if (isMcpServerDisabledForExecution(name) || connectionAttempts.get(attempt.key) !== attempt) { + await transport.close().catch(() => {}) + await inProcessServer?.close().catch(() => {}) + return isMcpServerDisabledForExecution(name) + ? { name, type: 'disabled', config: serverRef } + : { name, type: 'failed', config: serverRef, error: 'MCP connection superseded' } + } + attempt.cancel = async () => { + if (transport instanceof StdioClientTransport && transport.pid) { + try { process.kill(transport.pid, 'SIGTERM') } catch { /* already exited */ } + } + await client.close().catch(() => {}) + await inProcessServer?.close().catch(() => {}) + } const connectPromise = client.connect(transport) const timeoutPromise = new Promise((_, reject) => { const timeoutId = setTimeout(() => { @@ -1179,6 +1209,15 @@ export const connectToServer = memoize( throw error } + if (isMcpServerDisabledForExecution(name) || connectionAttempts.get(attempt.key) !== attempt) { + await client.close().catch(() => {}) + await inProcessServer?.close().catch(() => {}) + return isMcpServerDisabledForExecution(name) + ? { name, type: 'disabled', config: serverRef } + : { name, type: 'failed', config: serverRef, error: 'MCP connection superseded' } + } + connectionProjects.set(client, getCwd()) + const capabilities = client.getServerCapabilities() const serverVersion = client.getServerVersion() const rawInstructions = client.getInstructions() @@ -1408,9 +1447,9 @@ export const connectToServer = memoize( ) originalOnclose?.() - if (connectionAttempts.get(name) !== attempt) return - connectionAttempts.delete(name) - clearServerFetchCaches(name) + if (connectionAttempts.get(attempt.key) !== attempt) return + connectionAttempts.delete(attempt.key) + clearServerFetchCaches(name, attempt.key) connectToServer.cache.delete(attempt.key) logMCPDebug(name, `Cleared connection cache for reconnection`) @@ -1596,6 +1635,7 @@ export const connectToServer = memoize( await cleanup() } + attempt.cleanup = wrappedCleanup const connectionDurationMs = Date.now() - connectStartTime logEvent('tengu_mcp_server_connection_succeeded', { connectionDurationMs, @@ -1657,6 +1697,35 @@ export const connectToServer = memoize( getServerCacheKey, ) +// Keep the policy check outside memoization: cached connections and retained +// tool closures must obey a disable even when session control delivery failed. +export const connectToServer = Object.assign( + async (...args: Parameters): Promise => { + const [name, config] = args + if (isMcpServerDisabledForExecution(name)) return { name, config, type: 'disabled' } + const pending = connectToServerMemoized(...args) + let result = await pending + if (isMcpServerDisabledForExecution(name)) { + if (result.type === 'connected') await result.cleanup() + result = { name, config, type: 'disabled' } + } + if (result.type === 'disabled' && connectToServerMemoized.cache.get(getServerCacheKey(name, config)) === pending) { + connectToServerMemoized.cache.delete(getServerCacheKey(name, config)) + } + return result + }, + { cache: connectToServerMemoized.cache }, +) + +function assertMcpServerEnabled(name: string, client: object): void { + if (isMcpServerDisabledForExecution(name, connectionProjects.get(client))) { + throw new TelemetrySafeError_I_VERIFIED_THIS_IS_NOT_CODE_OR_FILEPATHS( + `MCP server "${name}" is disabled`, + 'MCP server disabled', + ) + } +} + /** * Clears the memoize cache for a specific server * @param name Server name @@ -1671,7 +1740,23 @@ export async function clearServerCache( // Detach before awaiting cleanup: a concurrent enable owns its own cache // entry, and clearing an empty cache must never start a new server. connectToServer.cache.delete(key) - if (connectionAttempts.get(name)?.key === key) clearServerFetchCaches(name) + const attempt = connectionAttempts.get(key) + if (attempt) { + connectionAttempts.delete(key) + clearServerFetchCaches(name, key) + if (attempt.cleanup) { + await attempt.cleanup() + return + } + // A disabled slow initializer must not delay control delivery. Close its + // transport now; any setup still awaiting environment/auth work sees the + // invalidated attempt before it can spawn or publish a connection. + await attempt.cancel?.() + void cached?.then(async client => { + if (client.type === 'connected') await client.cleanup() + }).catch(() => {}) + return + } try { const wrappedClient = await cached @@ -1697,12 +1782,17 @@ export async function clearServerCache( export async function ensureConnectedClient( client: ConnectedMCPServer, ): Promise { + assertMcpServerEnabled(client.name, client.client) // SDK MCP servers run in-process and are handled separately via setupSdkMcpClients if (client.config.type === 'sdk') { return client } - const connectedClient = await connectToServer(client.name, client.config) + const connectedClient = await runWithCwdOverride( + connectionProjects.get(client.client) ?? getCwd(), + () => connectToServer(client.name, client.config), + ) + assertMcpServerEnabled(client.name, client.client) if (connectedClient.type !== 'connected') { throw new TelemetrySafeError_I_VERIFIED_THIS_IS_NOT_CODE_OR_FILEPATHS( `MCP server "${client.name}" is not connected`, @@ -2001,7 +2091,7 @@ export const fetchToolsForClient = memoizeWithLRU( return [] } }, - (client: MCPServerConnection) => client.name, + getMcpClientCacheKey, MCP_FETCH_CACHE_SIZE, ) @@ -2034,7 +2124,7 @@ export const fetchResourcesForClient = memoizeWithLRU( return [] } }, - (client: MCPServerConnection) => client.name, + getMcpClientCacheKey, MCP_FETCH_CACHE_SIZE, ) @@ -2110,7 +2200,7 @@ export const fetchCommandsForClient = memoizeWithLRU( return [] } }, - (client: MCPServerConnection) => client.name, + getMcpClientCacheKey, MCP_FETCH_CACHE_SIZE, ) @@ -3226,6 +3316,7 @@ async function callMCPTool({ _meta?: Record structuredContent?: Record }> { + assertMcpServerEnabled(name, client) const toolStartTime = Date.now() let progressInterval: NodeJS.Timeout | undefined @@ -3475,6 +3566,7 @@ export async function setupSdkMcpClients( // Connect the client await client.connect(transport) + connectionProjects.set(client, getCwd()) // Get capabilities from the server const capabilities = client.getServerCapabilities() diff --git a/src/services/mcp/config.execution.test.ts b/src/services/mcp/config.execution.test.ts new file mode 100644 index 00000000..7ad18bf9 --- /dev/null +++ b/src/services/mcp/config.execution.test.ts @@ -0,0 +1,49 @@ +import '../../../preload.ts' +import { afterEach, beforeEach, describe, expect, test } from 'bun:test' +import { mkdtemp, rm, writeFile } from 'node:fs/promises' +import { join } from 'node:path' +import { tmpdir } from 'node:os' +import { isMcpServerDisabledForExecution } from './config.js' +import { getGlobalClaudeFile } from '../../utils/env.js' +import { _setGlobalConfigCacheForTesting, getProjectPathForConfig } from '../../utils/config.js' + +let root: string +let previousConfigDir: string | undefined +beforeEach(async () => { + root = await mkdtemp(join(tmpdir(), 'qa005-config-')) + previousConfigDir = process.env.CLAUDE_CONFIG_DIR + process.env.CLAUDE_CONFIG_DIR = root + getGlobalClaudeFile.cache.clear?.() +}) +afterEach(async () => { + if (previousConfigDir === undefined) delete process.env.CLAUDE_CONFIG_DIR + else process.env.CLAUDE_CONFIG_DIR = previousConfigDir + getGlobalClaudeFile.cache.clear?.() + _setGlobalConfigCacheForTesting(null) + await rm(root, { recursive: true, force: true }) +}) + +describe('MCP execution policy freshness', () => { + test('reads another process write immediately even when the global cache still permits it', async () => { + const key = getProjectPathForConfig(root) + _setGlobalConfigCacheForTesting({ projects: { [key]: { disabledMcpServers: [] } } } as never) + const file = getGlobalClaudeFile() + expect(isMcpServerDisabledForExecution('echo', root)).toBe(false) + const proc = Bun.spawn([process.execPath, '-e', 'require("fs").writeFileSync(process.argv[1], process.argv[2])', file, + JSON.stringify({ projects: { [key]: { disabledMcpServers: ['echo'] } } })], { stdout: 'pipe', stderr: 'pipe' }) + expect(await proc.exited).toBe(0) + expect(isMcpServerDisabledForExecution('echo', root)).toBe(true) + expect(isMcpServerDisabledForExecution('echo', join(root, 'other'))).toBe(false) + await writeFile(file, JSON.stringify({ projects: { [key]: { disabledMcpServers: [] } } })) + expect(isMcpServerDisabledForExecution('echo', root)).toBe(false) + }) + + test('does not authorize execution from unreadable policy shapes', async () => { + const key = getProjectPathForConfig(root) + for (const contents of ['{', 'null', '[]', '{"projects":null}', JSON.stringify({ projects: { [key]: null } }), + JSON.stringify({ projects: { [key]: { disabledMcpServers: 'echo' } } })]) { + await writeFile(getGlobalClaudeFile(), contents) + expect(() => isMcpServerDisabledForExecution('echo', root)).toThrow('Cannot read MCP enablement state') + } + }) +}) diff --git a/src/services/mcp/config.ts b/src/services/mcp/config.ts index 04855f32..a0b6001e 100644 --- a/src/services/mcp/config.ts +++ b/src/services/mcp/config.ts @@ -14,6 +14,7 @@ import { saveGlobalConfig, } from '../../utils/config.js' import { getCwd } from '../../utils/cwd.js' +import { getGlobalClaudeFile } from '../../utils/env.js' import { logForDebugging } from '../../utils/debug.js' import { getErrnoCode } from '../../utils/errors.js' import { getFsImplementation } from '../../utils/fsOperations.js' @@ -1688,6 +1689,42 @@ export function isMcpServerDisabled(name: string): boolean { return disabledServers.includes(name) } +/** Read the execution policy from disk: another process may have just disabled + * this server and the normal config cache refreshes only once per second. + * An unreadable or malformed policy must not authorize a new tool execution. + */ +export function isMcpServerDisabledForExecution(name: string, cwd = getCwd()): boolean { + let contents: string + try { + contents = getFsImplementation().readFileSync(getGlobalClaudeFile(), { encoding: 'utf8' }) + } catch (error) { + if (getErrnoCode(error) === 'ENOENT') return isDefaultDisabledBuiltin(name) + throw new Error(`Cannot read MCP enablement state for "${name}"`) + } + const config = safeParseJSONWithoutCache(contents.replace(/^\uFEFF/, '')) + if (!config || typeof config !== 'object' || Array.isArray(config)) { + throw new Error(`Cannot read MCP enablement state for "${name}"`) + } + const projects = (config as { projects?: Record }).projects + if (projects !== undefined && (!projects || typeof projects !== 'object' || Array.isArray(projects))) { + throw new Error(`Cannot read MCP enablement state for "${name}"`) + } + const project = projects?.[getProjectPathForConfig(cwd)] as { + enabledMcpServers?: unknown + disabledMcpServers?: unknown + } | undefined + if (project !== undefined && (!project || typeof project !== 'object' || Array.isArray(project))) { + throw new Error(`Cannot read MCP enablement state for "${name}"`) + } + const servers = isDefaultDisabledBuiltin(name) ? project?.enabledMcpServers : project?.disabledMcpServers + if (servers !== undefined && (!Array.isArray(servers) || !servers.every(value => typeof value === 'string'))) { + throw new Error(`Cannot read MCP enablement state for "${name}"`) + } + return isDefaultDisabledBuiltin(name) + ? !(servers as string[] | undefined)?.includes(name) + : (servers as string[] | undefined)?.includes(name) ?? false +} + function toggleMembership( list: string[], name: string, diff --git a/src/services/mcp/useManageMCPConnections.lifecycle.test.tsx b/src/services/mcp/useManageMCPConnections.lifecycle.test.tsx index b1c41a6e..8fcafdbc 100644 --- a/src/services/mcp/useManageMCPConnections.lifecycle.test.tsx +++ b/src/services/mcp/useManageMCPConnections.lifecycle.test.tsx @@ -1,10 +1,20 @@ -import { afterAll, afterEach, beforeEach, describe, expect, mock, test } from 'bun:test' +import '../../../preload.ts' +import { afterAll, afterEach, beforeEach, describe, expect, mock, spyOn, test } from 'bun:test' import React from 'react' import { render } from 'ink' import { PassThrough } from 'node:stream' +import { mkdtempSync, mkdirSync, rmSync } from 'node:fs' +import { join } from 'node:path' +import { tmpdir } from 'node:os' +import { + ToolListChangedNotificationSchema, + PromptListChangedNotificationSchema, + ResourceListChangedNotificationSchema, +} from '@modelcontextprotocol/sdk/types.js' +import { runWithCwdOverride } from '../../utils/cwd.js' import type { AppState } from '../../state/AppState.js' import { getDefaultAppState } from '../../state/AppStateStore.js' -import type { ConnectedMCPServer } from './types.js' +import type { ConnectedMCPServer, ScopedMcpServerConfig } from './types.js' const appStateModule = { ...await import('../../state/AppState.js') } const clientModule = { ...await import('./client.js') } @@ -58,8 +68,8 @@ const { useManageMCPConnections } = await import('./useManageMCPConnections.js') let actions: ReturnType -function Harness() { - actions = useManageMCPConnections(undefined) +function Harness({ configs }: { configs?: Record }) { + actions = useManageMCPConnections(configs) return null } @@ -293,3 +303,120 @@ test('ignores a previous connection close after an explicit reconnect', async () app.unmount() } }) + + +describe('MCP list change notification cache isolation', () => { + test.each(['tools', 'prompts', 'resources'] as const)( + '%s notification refreshes only the originating project for same-name servers', + async (kind) => { + isDisabled = false + const root = mkdtempSync(join(tmpdir(), 'qa005-notification-')) + const firstProject = join(root, 'first') + const secondProject = join(root, 'second') + mkdirSync(firstProject) + mkdirSync(secondProject) + const connections: ConnectedMCPServer[] = [] + const fetchers = { + tools: clientModule.fetchToolsForClient, + prompts: clientModule.fetchCommandsForClient, + resources: clientModule.fetchResourcesForClient, + } + const schemas = { + tools: ToolListChangedNotificationSchema, + prompts: PromptListChangedNotificationSchema, + resources: ResourceListChangedNotificationSchema, + } + let app: ReturnType | undefined + let notificationSpy: ReturnType | undefined + try { + async function createFixture(project: string, label: string) { + let version = 1 + const requests: string[] = [] + const result = await runWithCwdOverride(project, () => clientModule.setupSdkMcpClients( + { 'test-server': { type: 'sdk', name: 'test-server' } }, + async (_name, message) => { + if (!('method' in message) || !('id' in message)) return message + const method = message.method + requests.push(method) + const itemName = `${label}-v${version}` + const result = method === 'initialize' + ? { + protocolVersion: '2024-11-05', + capabilities: { + tools: { listChanged: true }, + prompts: { listChanged: true }, + resources: { listChanged: true }, + }, + serverInfo: { name: 'notification-fixture', version: '1' }, + } + : method === 'tools/list' + ? { tools: [{ name: itemName, inputSchema: { type: 'object' } }] } + : method === 'prompts/list' + ? { prompts: [{ name: itemName }] } + : { resources: [{ name: itemName, uri: `fixture://${itemName}` }] } + return { jsonrpc: '2.0', id: message.id, result } + }, + )) + const client = result.clients[0] + if (!client || client.type !== 'connected') throw new Error(`Fixture must connect: ${JSON.stringify(client)}`) + connections.push(client) + return { client, requests, advance: () => { version++ } } + } + + const first = await createFixture(firstProject, 'first') + const second = await createFixture(secondProject, 'second') + const fetchList = fetchers[kind] + const firstBefore = await fetchList(first.client) + const secondBefore = await fetchList(second.client) + const firstRequestsBefore = first.requests.filter(method => method === `${kind}/list`).length + const secondRequestsBefore = second.requests.filter(method => method === `${kind}/list`).length + expect(firstBefore).not.toBe(secondBefore) + + let notify: (() => Promise) | undefined + const setNotificationHandler = first.client.client.setNotificationHandler.bind(first.client.client) + notificationSpy = spyOn(first.client.client, 'setNotificationHandler').mockImplementation((schema, handler) => { + if (schema === schemas[kind]) notify = handler as () => Promise + setNotificationHandler(schema, handler) + }) + state = { ...state, mcp: { ...state.mcp, clients: [first.client] } } + reconnectMcpServerImpl.mockResolvedValue({ + name: 'test-server', client: first.client, tools: [], commands: [], resources: [], + }) + app = render(, { + stdout: new PassThrough(), stderr: new PassThrough(), stdin: new PassThrough(), + exitOnCtrlC: false, patchConsole: false, + }) + await Bun.sleep(0) + await actions.reconnectMcpServer('test-server') + expect(notify).toBeDefined() + + first.advance() + second.advance() + // Delivery can occur while a different project is active. The client + // that registered the handler still owns this notification's cache. + await runWithCwdOverride(secondProject, () => notify!()) + await Bun.sleep(20) + const firstAfter = await fetchList(first.client) + const secondAfter = await fetchList(second.client) + expect(firstAfter).not.toBe(firstBefore) + expect(firstAfter[0]?.name).toContain('first-v2') + expect(secondAfter).toBe(secondBefore) + expect(secondAfter[0]?.name).toContain('second-v1') + expect(first.requests.filter(method => method === `${kind}/list`)).toHaveLength(firstRequestsBefore + 1) + expect(second.requests.filter(method => method === `${kind}/list`)).toHaveLength(secondRequestsBefore) + const published = kind === 'tools' + ? state.mcp.tools + : kind === 'prompts' + ? state.mcp.commands + : state.mcp.resources['test-server'] + expect(published?.[0]?.name).toContain('first-v2') + } finally { + app?.unmount() + notificationSpy?.mockRestore() + await Promise.all(connections.map(connection => connection.cleanup())) + for (const fetchList of Object.values(fetchers)) fetchList.cache.clear() + rmSync(root, { recursive: true, force: true }) + } + }, + ) +}) diff --git a/src/services/mcp/useManageMCPConnections.ts b/src/services/mcp/useManageMCPConnections.ts index 61599c07..75e7b1b8 100644 --- a/src/services/mcp/useManageMCPConnections.ts +++ b/src/services/mcp/useManageMCPConnections.ts @@ -10,6 +10,7 @@ import { fetchResourcesForClient, fetchToolsForClient, getMcpToolsCommandsAndResources, + getMcpClientCacheKey, reconnectMcpServerImpl, setMcpConnectionClosedHandler, } from './client.js' @@ -530,9 +531,9 @@ export function useManageMCPConnections( try { // Grab cached promise before invalidating to log previous count const previousToolsPromise = fetchToolsForClient.cache.get( - client.name, + getMcpClientCacheKey(client), ) - fetchToolsForClient.cache.delete(client.name) + fetchToolsForClient.cache.delete(getMcpClientCacheKey(client)) const newTools = await fetchToolsForClient(client) const newCount = newTools.length if (previousToolsPromise) { @@ -582,7 +583,7 @@ export function useManageMCPConnections( try { // Skills come from resources, not prompts — don't invalidate their // cache here. fetchMcpSkillsForClient returns the cached result. - fetchCommandsForClient.cache.delete(client.name) + fetchCommandsForClient.cache.delete(getMcpClientCacheKey(client)) const [mcpPrompts, mcpSkills] = await Promise.all([ fetchCommandsForClient(client), feature('MCP_SKILLS') @@ -618,14 +619,14 @@ export function useManageMCPConnections( type: 'resources' as AnalyticsMetadata_I_VERIFIED_THIS_IS_NOT_CODE_OR_FILEPATHS, }) try { - fetchResourcesForClient.cache.delete(client.name) + fetchResourcesForClient.cache.delete(getMcpClientCacheKey(client)) if (feature('MCP_SKILLS')) { // Skills are discovered from resources, so refresh them too. // Invalidate prompts cache as well: we write commands here, // and a concurrent prompts/list_changed could otherwise have // us stomp its fresh result with our cached stale one. fetchMcpSkillsForClient!.cache.delete(client.name) - fetchCommandsForClient.cache.delete(client.name) + fetchCommandsForClient.cache.delete(getMcpClientCacheKey(client)) const [newResources, mcpPrompts, mcpSkills] = await Promise.all([ fetchResourcesForClient(client),