mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-09 03:17:58 +00:00
fix: prevent MCP server restarts during active tool executions
- Add tracking of active tool executions in McpHub - Prevent server restarts when tools are running - Update toggleToolAlwaysAllow to skip restart during tool execution - Update toggleToolEnabledForPrompt to skip restart during tool execution - Prevent toggleServerDisabled when tools are running - Add comprehensive tests for the new behavior Fixes #7189
This commit is contained in:
parent
4222036c58
commit
6155283ff0
2 changed files with 376 additions and 14 deletions
|
|
@ -151,6 +151,7 @@ export class McpHub {
|
||||||
isConnecting: boolean = false
|
isConnecting: boolean = false
|
||||||
private refCount: number = 0 // Reference counter for active clients
|
private refCount: number = 0 // Reference counter for active clients
|
||||||
private configChangeDebounceTimers: Map<string, NodeJS.Timeout> = new Map()
|
private configChangeDebounceTimers: Map<string, NodeJS.Timeout> = new Map()
|
||||||
|
private activeToolExecutions: Map<string, Set<string>> = new Map() // Track active tool executions per server
|
||||||
|
|
||||||
constructor(provider: ClineProvider) {
|
constructor(provider: ClineProvider) {
|
||||||
this.providerRef = new WeakRef(provider)
|
this.providerRef = new WeakRef(provider)
|
||||||
|
|
@ -1169,6 +1170,16 @@ export class McpHub {
|
||||||
}
|
}
|
||||||
|
|
||||||
async restartConnection(serverName: string, source?: "global" | "project"): Promise<void> {
|
async restartConnection(serverName: string, source?: "global" | "project"): Promise<void> {
|
||||||
|
// 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
|
this.isConnecting = true
|
||||||
|
|
||||||
// Check if MCP is globally enabled
|
// Check if MCP is globally enabled
|
||||||
|
|
@ -1357,6 +1368,20 @@ export class McpHub {
|
||||||
}
|
}
|
||||||
|
|
||||||
const serverSource = connection.server.source || "global"
|
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
|
// Update the server config in the appropriate file
|
||||||
await this.updateServerConfig(serverName, { disabled }, serverSource)
|
await this.updateServerConfig(serverName, { disabled }, serverSource)
|
||||||
|
|
||||||
|
|
@ -1603,19 +1628,69 @@ export class McpHub {
|
||||||
timeout = 60 * 1000
|
timeout = 60 * 1000
|
||||||
}
|
}
|
||||||
|
|
||||||
return await connection.client.request(
|
// Track this tool execution as active
|
||||||
{
|
const serverKey = `${serverName}:${source || "global"}`
|
||||||
method: "tools/call",
|
if (!this.activeToolExecutions.has(serverKey)) {
|
||||||
params: {
|
this.activeToolExecutions.set(serverKey, new Set())
|
||||||
name: toolName,
|
}
|
||||||
arguments: toolArguments,
|
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,
|
||||||
CallToolResultSchema,
|
{
|
||||||
{
|
timeout,
|
||||||
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,
|
shouldAllow: boolean,
|
||||||
): Promise<void> {
|
): Promise<void> {
|
||||||
try {
|
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) {
|
} catch (error) {
|
||||||
this.showErrorMessage(
|
this.showErrorMessage(
|
||||||
`Failed to toggle always allow for tool "${toolName}" on server "${serverName}" with source "${source}"`,
|
`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<void> {
|
||||||
|
// 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(
|
async toggleToolEnabledForPrompt(
|
||||||
serverName: string,
|
serverName: string,
|
||||||
source: "global" | "project",
|
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 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.
|
// When isEnabled is false, we want to add the tool to the disabledTools list.
|
||||||
const addToolToDisabledList = !isEnabled
|
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) {
|
} catch (error) {
|
||||||
this.showErrorMessage(`Failed to update settings for tool ${toolName}`, error)
|
this.showErrorMessage(`Failed to update settings for tool ${toolName}`, error)
|
||||||
throw error // Re-throw to ensure the error is properly handled
|
throw error // Re-throw to ensure the error is properly handled
|
||||||
|
|
|
||||||
|
|
@ -908,6 +908,162 @@ describe("McpHub", () => {
|
||||||
expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toBeDefined()
|
expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toBeDefined()
|
||||||
expect(writtenConfig.mcpServers["test-server"].alwaysAllow).toContain("new-tool")
|
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", () => {
|
describe("toggleToolEnabledForPrompt", () => {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue