diff --git a/src/services/mcp/McpHub.ts b/src/services/mcp/McpHub.ts index 271c6e1fb3..e9bbf75614 100644 --- a/src/services/mcp/McpHub.ts +++ b/src/services/mcp/McpHub.ts @@ -151,6 +151,7 @@ export class McpHub { isConnecting: boolean = false private refCount: number = 0 // Reference counter for active clients private configChangeDebounceTimers: Map = new Map() + private activeToolExecutions: Map> = new Map() // Track active tool executions per server constructor(provider: ClineProvider) { this.providerRef = new WeakRef(provider) @@ -1169,6 +1170,16 @@ export class McpHub { } async restartConnection(serverName: string, source?: "global" | "project"): Promise { + // Check if there are active tool executions for this server + if (this.hasActiveToolExecutions(serverName, source)) { + console.log(`Skipping restart for ${serverName} - tools are currently executing`) + vscode.window.showWarningMessage( + t("mcp:errors.cannot_restart_tools_running", { serverName }) || + `Cannot restart server "${serverName}" while tools are running. Please wait for tool execution to complete.`, + ) + return + } + this.isConnecting = true // Check if MCP is globally enabled @@ -1357,6 +1368,20 @@ export class McpHub { } const serverSource = connection.server.source || "global" + + // Check if there are active tool executions for this server + if (this.hasActiveToolExecutions(serverName, serverSource)) { + // Show a warning message and don't proceed with the toggle + vscode.window.showWarningMessage( + t("mcp:errors.cannot_toggle_server_tools_running", { + serverName, + action: disabled ? "disable" : "enable", + }) || + `Cannot ${disabled ? "disable" : "enable"} server "${serverName}" while tools are running. Please wait for tool execution to complete.`, + ) + return + } + // Update the server config in the appropriate file await this.updateServerConfig(serverName, { disabled }, serverSource) @@ -1603,19 +1628,69 @@ export class McpHub { timeout = 60 * 1000 } - return await connection.client.request( - { - method: "tools/call", - params: { - name: toolName, - arguments: toolArguments, + // Track this tool execution as active + const serverKey = `${serverName}:${source || "global"}` + if (!this.activeToolExecutions.has(serverKey)) { + this.activeToolExecutions.set(serverKey, new Set()) + } + const executionId = `${toolName}:${Date.now()}` + this.activeToolExecutions.get(serverKey)!.add(executionId) + + try { + const result = await connection.client.request( + { + method: "tools/call", + params: { + name: toolName, + arguments: toolArguments, + }, }, - }, - CallToolResultSchema, - { - timeout, - }, - ) + CallToolResultSchema, + { + timeout, + }, + ) + + // Remove from active executions on success + this.activeToolExecutions.get(serverKey)?.delete(executionId) + if (this.activeToolExecutions.get(serverKey)?.size === 0) { + this.activeToolExecutions.delete(serverKey) + } + + return result + } catch (error) { + // Remove from active executions on error + this.activeToolExecutions.get(serverKey)?.delete(executionId) + if (this.activeToolExecutions.get(serverKey)?.size === 0) { + this.activeToolExecutions.delete(serverKey) + } + throw error + } + } + + /** + * Check if any tools are currently executing for a specific server + * @param serverName The name of the server to check + * @param source The source of the server (global or project) + * @returns true if there are active tool executions, false otherwise + */ + private hasActiveToolExecutions(serverName: string, source?: "global" | "project"): boolean { + const serverKey = `${serverName}:${source || "global"}` + const activeTools = this.activeToolExecutions.get(serverKey) + return activeTools ? activeTools.size > 0 : false + } + + /** + * Check if any tools are currently executing across all servers + * @returns true if there are any active tool executions, false otherwise + */ + private hasAnyActiveToolExecutions(): boolean { + for (const [, tools] of this.activeToolExecutions) { + if (tools.size > 0) { + return true + } + } + return false } /** @@ -1703,7 +1778,15 @@ export class McpHub { shouldAllow: boolean, ): Promise { try { - await this.updateServerToolList(serverName, source, toolName, "alwaysAllow", shouldAllow) + // Check if there are active tool executions for this server + if (this.hasActiveToolExecutions(serverName, source)) { + console.log(`Skipping server restart for ${serverName} - tools are currently executing`) + // Update the config file without triggering a restart + await this.updateServerToolListWithoutRestart(serverName, source, toolName, "alwaysAllow", shouldAllow) + } else { + // Normal flow - update and allow restart if needed + await this.updateServerToolList(serverName, source, toolName, "alwaysAllow", shouldAllow) + } } catch (error) { this.showErrorMessage( `Failed to toggle always allow for tool "${toolName}" on server "${serverName}" with source "${source}"`, @@ -1713,6 +1796,114 @@ export class McpHub { } } + /** + * Update server tool list without triggering a restart + * This is used when tools are actively running to prevent interruption + */ + private async updateServerToolListWithoutRestart( + serverName: string, + source: "global" | "project", + toolName: string, + listName: "alwaysAllow" | "disabledTools", + addTool: boolean, + ): Promise { + // Find the connection with matching name and source + const connection = this.findConnection(serverName, source) + + if (!connection) { + throw new Error(`Server ${serverName} with source ${source} not found`) + } + + // Determine the correct config path based on the source + let configPath: string + if (source === "project") { + // Get project MCP config path + const projectMcpPath = await this.getProjectMcpPath() + if (!projectMcpPath) { + throw new Error("Project MCP configuration file not found") + } + configPath = projectMcpPath + } else { + // Get global MCP settings path + configPath = await this.getMcpSettingsFilePath() + } + + // Normalize path for cross-platform compatibility + const normalizedPath = process.platform === "win32" ? configPath.replace(/\\/g, "/") : configPath + + // Read the appropriate config file + const content = await fs.readFile(normalizedPath, "utf-8") + const config = JSON.parse(content) + + if (!config.mcpServers) { + config.mcpServers = {} + } + + if (!config.mcpServers[serverName]) { + config.mcpServers[serverName] = { + type: "stdio", + command: "node", + args: [], // Default to an empty array; can be set later if needed + } + } + + if (!config.mcpServers[serverName][listName]) { + config.mcpServers[serverName][listName] = [] + } + + const targetList = config.mcpServers[serverName][listName] + const toolIndex = targetList.indexOf(toolName) + + if (addTool && toolIndex === -1) { + targetList.push(toolName) + } else if (!addTool && toolIndex !== -1) { + targetList.splice(toolIndex, 1) + } + + // Write the config file directly without triggering file watcher + // We'll temporarily disable the watcher to prevent restart + const watcherKey = source === "project" ? this.projectMcpWatcher : this.settingsWatcher + const wasWatcherActive = !!watcherKey + + if (wasWatcherActive && source === "project" && this.projectMcpWatcher) { + this.projectMcpWatcher.dispose() + this.projectMcpWatcher = undefined + } else if (wasWatcherActive && source === "global" && this.settingsWatcher) { + this.settingsWatcher.dispose() + this.settingsWatcher = undefined + } + + await fs.writeFile(normalizedPath, JSON.stringify(config, null, 2)) + + // Update the in-memory tool state without restarting + if (connection) { + // Update the tool's alwaysAllow or enabledForPrompt status in memory + const tools = connection.server.tools + if (tools) { + const tool = tools.find((t) => t.name === toolName) + if (tool) { + if (listName === "alwaysAllow") { + tool.alwaysAllow = addTool + } else if (listName === "disabledTools") { + tool.enabledForPrompt = !addTool + } + } + } + await this.notifyWebviewOfServerChanges() + } + + // Re-enable the watcher after a short delay + if (wasWatcherActive) { + setTimeout(async () => { + if (source === "project") { + await this.watchProjectMcpFile() + } else { + await this.watchMcpSettingsFile() + } + }, 1000) + } + } + async toggleToolEnabledForPrompt( serverName: string, source: "global" | "project", @@ -1723,7 +1914,22 @@ export class McpHub { // When isEnabled is true, we want to remove the tool from the disabledTools list. // When isEnabled is false, we want to add the tool to the disabledTools list. const addToolToDisabledList = !isEnabled - await this.updateServerToolList(serverName, source, toolName, "disabledTools", addToolToDisabledList) + + // Check if there are active tool executions for this server + if (this.hasActiveToolExecutions(serverName, source)) { + console.log(`Skipping server restart for ${serverName} - tools are currently executing`) + // Update the config file without triggering a restart + await this.updateServerToolListWithoutRestart( + serverName, + source, + toolName, + "disabledTools", + addToolToDisabledList, + ) + } else { + // Normal flow - update and allow restart if needed + await this.updateServerToolList(serverName, source, toolName, "disabledTools", addToolToDisabledList) + } } catch (error) { this.showErrorMessage(`Failed to update settings for tool ${toolName}`, error) throw error // Re-throw to ensure the error is properly handled diff --git a/src/services/mcp/__tests__/McpHub.spec.ts b/src/services/mcp/__tests__/McpHub.spec.ts index ebce2d5b2a..9165851199 100644 --- a/src/services/mcp/__tests__/McpHub.spec.ts +++ b/src/services/mcp/__tests__/McpHub.spec.ts @@ -908,6 +908,162 @@ describe("McpHub", () => { expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toBeDefined() expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toContain("new-tool") }) + + it("should skip server restart when toggling alwaysAllow during active tool execution", async () => { + // Mock fs.readFile to return existing config + vi.mocked(fs.readFile).mockResolvedValue( + JSON.stringify({ + mcpServers: { + "test-server": { + type: "stdio", + command: "node", + args: ["test.js"], + alwaysAllow: [], + }, + }, + }), + ) + + // Create a connected server + const mockConnection: ConnectedMcpConnection = { + type: "connected", + server: { + name: "test-server", + config: JSON.stringify({ + type: "stdio", + command: "node", + args: ["test.js"], + alwaysAllow: [], + }), + status: "connected", + source: "global", + tools: [ + { + name: "test-tool", + description: "Test tool", + alwaysAllow: false, + enabledForPrompt: true, + }, + ], + }, + client: { + request: vi.fn().mockImplementation(() => { + // Simulate long-running tool + return new Promise((resolve) => { + setTimeout(() => resolve({ content: [] }), 500) + }) + }), + } as any, + transport: {} as any, + } + + mcpHub.connections = [mockConnection] + + // Spy on console.log to verify the skip message + const consoleLogSpy = vi.spyOn(console, "log") + + // Start a tool execution + const toolPromise = mcpHub.callTool("test-server", "test-tool", {}) + + // While tool is running, toggle alwaysAllow + await mcpHub.toggleToolAlwaysAllow("test-server", "global", "test-tool", true) + + // Verify that the skip message was logged + expect(consoleLogSpy).toHaveBeenCalledWith( + "Skipping server restart for test-server - tools are currently executing", + ) + + // Verify that the config was updated + const writeCalls = vi.mocked(fs.writeFile).mock.calls + const lastWriteCall = writeCalls[writeCalls.length - 1] + const writtenConfig = JSON.parse(lastWriteCall[1] as string) + expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toContain("test-tool") + + // Verify that the in-memory tool state was updated + const tool = mockConnection.server.tools?.find((t) => t.name === "test-tool") + expect(tool?.alwaysAllow).toBe(true) + + // Wait for tool to complete + await toolPromise + + // Verify that the server status is still connected + expect(mockConnection.server.status).toBe("connected") + + consoleLogSpy.mockRestore() + }) + + it("should prevent server restart when tools are running", async () => { + const mockConnection: ConnectedMcpConnection = { + type: "connected", + server: { + name: "test-server", + config: JSON.stringify({ type: "stdio", command: "node", args: ["test.js"] }), + status: "connected", + source: "global", + }, + client: { + request: vi.fn().mockImplementation(() => { + // Simulate long-running tool + return new Promise((resolve) => { + setTimeout(() => resolve({ content: [] }), 500) + }) + }), + } as any, + transport: {} as any, + } + + mcpHub.connections = [mockConnection] + + // Start a tool execution + const toolPromise = mcpHub.callTool("test-server", "some-tool", {}) + + // Try to restart the server while tool is running + await mcpHub.restartConnection("test-server", "global") + + // Verify that the server was not restarted + expect(mcpHub.connections[0]).toBe(mockConnection) + expect(mockConnection.server.status).toBe("connected") + + // Wait for tool to complete + await toolPromise + }) + + it("should prevent toggling server disabled state when tools are running", async () => { + const mockConnection: ConnectedMcpConnection = { + type: "connected", + server: { + name: "test-server", + config: JSON.stringify({ type: "stdio", command: "node", args: ["test.js"] }), + status: "connected", + source: "global", + disabled: false, + }, + client: { + request: vi.fn().mockImplementation(() => { + // Simulate long-running tool + return new Promise((resolve) => { + setTimeout(() => resolve({ content: [] }), 500) + }) + }), + } as any, + transport: {} as any, + } + + mcpHub.connections = [mockConnection] + + // Start a tool execution + const toolPromise = mcpHub.callTool("test-server", "some-tool", {}) + + // Try to disable the server while tool is running + await mcpHub.toggleServerDisabled("test-server", true, "global") + + // Verify that the server was not disabled + expect(mockConnection.server.disabled).toBe(false) + expect(mockConnection.server.status).toBe("connected") + + // Wait for tool to complete + await toolPromise + }) }) describe("toggleToolEnabledForPrompt", () => {