diff --git a/src/vs/base/node/processes.ts b/src/vs/base/node/processes.ts index 3a9fd987e6d..c2028fc5409 100644 --- a/src/vs/base/node/processes.ts +++ b/src/vs/base/node/processes.ts @@ -8,10 +8,12 @@ import { Stats, promises } from 'fs'; import { getCaseInsensitive } from '../common/objects.js'; import * as path from '../common/path.js'; import * as Platform from '../common/platform.js'; -import * as process from '../common/process.js'; +import * as processCommon from '../common/process.js'; import { CommandOptions, ForkOptions, Source, SuccessData, TerminateResponse, TerminateResponseCode } from '../common/processes.js'; import * as Types from '../common/types.js'; import * as pfs from './pfs.js'; +import { FileAccess } from '../common/network.js'; +import Stream from 'stream'; export { Source, TerminateResponseCode, type CommandOptions, type ForkOptions, type SuccessData, type TerminateResponse }; export type ValueCallback = (value: T | Promise) => void; @@ -19,7 +21,7 @@ export type ErrorCallback = (error?: any) => void; export type ProgressCallback = (progress: T) => void; -export function getWindowsShell(env = process.env as Platform.IProcessEnvironment): string { +export function getWindowsShell(env = processCommon.env as Platform.IProcessEnvironment): string { return env['comspec'] || 'cmd.exe'; } @@ -81,17 +83,17 @@ async function fileExistsDefault(path: string): Promise { return false; } -export function getWindowPathExtensions(env = process.env) { +export function getWindowPathExtensions(env = processCommon.env) { return (getCaseInsensitive(env, 'PATHEXT') as string || '.COM;.EXE;.BAT;.CMD').split(';'); } -export async function findExecutable(command: string, cwd?: string, paths?: string[], env: Platform.IProcessEnvironment = process.env as Platform.IProcessEnvironment, fileExists: (path: string) => Promise = fileExistsDefault): Promise { +export async function findExecutable(command: string, cwd?: string, paths?: string[], env: Platform.IProcessEnvironment = processCommon.env as Platform.IProcessEnvironment, fileExists: (path: string) => Promise = fileExistsDefault): Promise { // If we have an absolute path then we take it. if (path.isAbsolute(command)) { return await fileExists(command) ? command : undefined; } if (cwd === undefined) { - cwd = process.cwd(); + cwd = processCommon.cwd(); } const dir = path.dirname(command); if (dir !== '.') { @@ -140,3 +142,42 @@ export async function findExecutable(command: string, cwd?: string, paths?: stri const fullPath = path.join(cwd, command); return await fileExists(fullPath) ? fullPath : undefined; } + +/** + * Kills a process and all its children. + * @param pid the process id to kill + * @param forceful whether to forcefully kill the process (default: false). Note + * that on Windows, terminal processes can _only_ be killed forcefully and this + * will throw when not forceful. + */ +export async function killTree(pid: number, forceful = false) { + let child: cp.ChildProcessByStdio; + if (Platform.isWindows) { + const windir = process.env['WINDIR'] || 'C:\\Windows'; + const taskKill = path.join(windir, 'System32', 'taskkill.exe'); + + const args = ['/T']; + if (forceful) { + args.push('/F'); + } + args.push('/PID', String(pid)); + child = cp.spawn(taskKill, args, { stdio: ['ignore', 'pipe', 'pipe'] }); + } else { + const killScript = FileAccess.asFileUri('vs/base/node/terminateProcess.sh').fsPath; + child = cp.spawn('/bin/sh', [killScript, String(pid), forceful ? '9' : '15'], { stdio: ['ignore', 'pipe', 'pipe'] }); + } + + return new Promise((resolve, reject) => { + const stdout: Buffer[] = []; + child.stdout.on('data', (data) => stdout.push(data)); + child.stderr.on('data', (data) => stdout.push(data)); + child.on('error', reject); + child.on('exit', (code) => { + if (code === 0) { + resolve(); + } else { + reject(new Error(`taskkill exited with code ${code}: ${Buffer.concat(stdout).toString()}`)); + } + }); + }); +} diff --git a/src/vs/base/node/terminateProcess.sh b/src/vs/base/node/terminateProcess.sh index acdcbf8ed42..a8e8738fdb4 100755 --- a/src/vs/base/node/terminateProcess.sh +++ b/src/vs/base/node/terminateProcess.sh @@ -1,12 +1,13 @@ -#!/bin/bash +#!/bin/sh + +ROOT_PID=$1 +SIGNAL=$2 terminateTree() { - for cpid in $(/usr/bin/pgrep -P $1); do - terminateTree $cpid - done - kill -9 $1 > /dev/null 2>&1 + for cpid in $(pgrep -P $1); do + terminateTree $cpid + done + kill -$SIGNAL $1 > /dev/null 2>&1 } -for pid in $*; do - terminateTree $pid -done \ No newline at end of file +terminateTree $ROOT_PID diff --git a/src/vs/workbench/api/node/extHostMcpNode.ts b/src/vs/workbench/api/node/extHostMcpNode.ts index 0e65834ab19..ee90057f2e3 100644 --- a/src/vs/workbench/api/node/extHostMcpNode.ts +++ b/src/vs/workbench/api/node/extHostMcpNode.ts @@ -8,14 +8,16 @@ import { readFile } from 'fs/promises'; import { homedir } from 'os'; import { parseEnvFile } from '../../../base/common/envfile.js'; import { untildify } from '../../../base/common/labels.js'; +import { DisposableMap } from '../../../base/common/lifecycle.js'; +import * as path from '../../../base/common/path.js'; import { StreamSplitter } from '../../../base/node/nodeStreams.js'; import { findExecutable } from '../../../base/node/processes.js'; import { ILogService, LogLevel } from '../../../platform/log/common/log.js'; import { McpConnectionState, McpServerLaunch, McpServerTransportStdio, McpServerTransportType } from '../../contrib/mcp/common/mcpTypes.js'; +import { McpStdioStateHandler } from '../../contrib/mcp/node/mcpStdioStateHandler.js'; +import { IExtHostInitDataService } from '../common/extHostInitDataService.js'; import { ExtHostMcpService } from '../common/extHostMcp.js'; import { IExtHostRpcService } from '../common/extHostRpcService.js'; -import * as path from '../../../base/common/path.js'; -import { IExtHostInitDataService } from '../common/extHostInitDataService.js'; export class NodeExtHostMpcService extends ExtHostMcpService { constructor( @@ -26,10 +28,7 @@ export class NodeExtHostMpcService extends ExtHostMcpService { super(extHostRpc, logService, initDataService); } - private nodeServers = new Map(); + private nodeServers = this._register(new DisposableMap()); protected override _startMcp(id: number, launch: McpServerLaunch): void { if (launch.type === McpServerTransportType.Stdio) { @@ -42,8 +41,7 @@ export class NodeExtHostMpcService extends ExtHostMcpService { override $stopMcp(id: number): void { const nodeServer = this.nodeServers.get(id); if (nodeServer) { - nodeServer.abortCtrl.abort(); - this.nodeServers.delete(id); + nodeServer.stop(); // will get removed from map when process is fully stopped } else { super.$stopMcp(id); } @@ -52,7 +50,7 @@ export class NodeExtHostMpcService extends ExtHostMcpService { override $sendMessage(id: number, message: string): void { const nodeServer = this.nodeServers.get(id); if (nodeServer) { - nodeServer.child.stdin.write(message + '\n'); + nodeServer.write(message); } else { super.$sendMessage(id, message); } @@ -82,7 +80,6 @@ export class NodeExtHostMpcService extends ExtHostMcpService { env[key] = value === null ? undefined : String(value); } - const abortCtrl = new AbortController(); let child: ChildProcessWithoutNullStreams; try { const home = homedir(); @@ -102,16 +99,17 @@ export class NodeExtHostMpcService extends ExtHostMcpService { child = spawn(executable, args, { stdio: 'pipe', cwd, - signal: abortCtrl.signal, env, shell, }); } catch (e) { onError(e); - abortCtrl.abort(); return; } + // Create the connection manager for graceful shutdown + const connectionManager = new McpStdioStateHandler(child); + this._proxy.$onDidChangeState(id, { state: McpConnectionState.Kind.Starting }); child.stdout.pipe(new StreamSplitter('\n')).on('data', line => this._proxy.$onDidReceiveMessage(id, line.toString())); @@ -126,22 +124,22 @@ export class NodeExtHostMpcService extends ExtHostMcpService { child.on('spawn', () => this._proxy.$onDidChangeState(id, { state: McpConnectionState.Kind.Running })); child.on('error', e => { - if (abortCtrl.signal.aborted) { + onError(e); + }); + child.on('exit', code => { + this.nodeServers.deleteAndDispose(id); + + if (code === 0 || connectionManager.stopped) { this._proxy.$onDidChangeState(id, { state: McpConnectionState.Kind.Stopped }); } else { - onError(e); - } - }); - child.on('exit', code => - code === 0 || abortCtrl.signal.aborted - ? this._proxy.$onDidChangeState(id, { state: McpConnectionState.Kind.Stopped }) - : this._proxy.$onDidChangeState(id, { + this._proxy.$onDidChangeState(id, { state: McpConnectionState.Kind.Error, message: `Process exited with code ${code}`, - }) - ); + }); + } + }); - this.nodeServers.set(id, { abortCtrl, child }); + this.nodeServers.set(id, connectionManager); } } diff --git a/src/vs/workbench/contrib/debug/node/debugAdapter.ts b/src/vs/workbench/contrib/debug/node/debugAdapter.ts index 4892338a1c1..1698883b62c 100644 --- a/src/vs/workbench/contrib/debug/node/debugAdapter.ts +++ b/src/vs/workbench/contrib/debug/node/debugAdapter.ts @@ -15,6 +15,7 @@ import * as nls from '../../../../nls.js'; import { IExtensionDescription } from '../../../../platform/extensions/common/extensions.js'; import { IDebugAdapterExecutable, IDebugAdapterNamedPipeServer, IDebugAdapterServer, IDebuggerContribution, IPlatformSpecificAdapterContribution } from '../common/debug.js'; import { AbstractDebugAdapter } from '../common/abstractDebugAdapter.js'; +import { killTree } from '../../../../base/node/processes.js'; /** * An implementation that communicates via two streams with the debug adapter. @@ -288,15 +289,7 @@ export class ExecutableDebugAdapter extends StreamDebugAdapter { // processes. Therefore we use TASKKILL.EXE await this.cancelPendingRequests(); if (platform.isWindows) { - return new Promise((c, e) => { - const killer = cp.exec(`taskkill /F /T /PID ${this.serverProcess!.pid}`, function (err, stdout, stderr) { - if (err) { - return e(err); - } - }); - killer.on('exit', c); - killer.on('error', e); - }); + return killTree(this.serverProcess!.pid!, true); } else { this.serverProcess.kill('SIGTERM'); return Promise.resolve(undefined); diff --git a/src/vs/workbench/contrib/mcp/node/mcpStdioStateHandler.ts b/src/vs/workbench/contrib/mcp/node/mcpStdioStateHandler.ts new file mode 100644 index 00000000000..fba931a0280 --- /dev/null +++ b/src/vs/workbench/contrib/mcp/node/mcpStdioStateHandler.ts @@ -0,0 +1,94 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + * Licensed under the MIT License. See License.txt in the project root for license information. + *--------------------------------------------------------------------------------------------*/ + +import { ChildProcessWithoutNullStreams } from 'child_process'; +import { TimeoutTimer } from '../../../../base/common/async.js'; +import { IDisposable } from '../../../../base/common/lifecycle.js'; +import { killTree } from '../../../../base/node/processes.js'; +import { isWindows } from '../../../../base/common/platform.js'; + +const enum McpProcessState { + Running, + StdinEnded, + KilledPolite, + KilledForceful, +} + +/** + * Manages graceful shutdown of MCP stdio connections following the MCP specification. + * + * Per spec, shutdown should: + * 1. Close the input stream to the child process + * 2. Wait for the server to exit, or send SIGTERM if it doesn't exit within 10 seconds + * 3. Send SIGKILL if the server doesn't exit within 10 seconds after SIGTERM + * 4. Allow forceful killing if called twice + */ +export class McpStdioStateHandler implements IDisposable { + private static readonly GRACE_TIME_MS = 10_000; + + private _procState = McpProcessState.Running; + private _nextTimeout?: IDisposable; + + public get stopped() { + return this._procState !== McpProcessState.Running; + } + + constructor( + private readonly _child: ChildProcessWithoutNullStreams, + private readonly _graceTimeMs: number = McpStdioStateHandler.GRACE_TIME_MS + ) { } + + /** + * Initiates graceful shutdown. If called while shutdown is already in progress, + * forces immediate termination. + */ + public stop(): void { + if (this._procState === McpProcessState.Running) { + try { + this._child.stdin.end(); + } catch (error) { + // If stdin.end() fails, continue with termination sequence + // This can happen if the stream is already in an error state + } + this._nextTimeout = new TimeoutTimer(() => this.killPolite(), this._graceTimeMs); + } else { + this._nextTimeout?.dispose(); + this.killForceful(); + } + } + + private async killPolite() { + this._procState = McpProcessState.KilledPolite; + this._nextTimeout = new TimeoutTimer(() => this.killForceful(), this._graceTimeMs); + + if (this._child.pid) { + if (!isWindows) { + await killTree(this._child.pid, false); + } + } else { + this._child.kill('SIGTERM'); + } + } + + private async killForceful() { + this._procState = McpProcessState.KilledForceful; + + if (this._child.pid) { + await killTree(this._child.pid, true); + } else { + this._child.kill(); + } + } + + public write(message: string): void { + if (!this.stopped) { + this._child.stdin.write(message + '\n'); + } + } + + public dispose() { + this._nextTimeout?.dispose(); + } +} diff --git a/src/vs/workbench/contrib/mcp/test/node/mcpStdioStateHandler.test.ts b/src/vs/workbench/contrib/mcp/test/node/mcpStdioStateHandler.test.ts new file mode 100644 index 00000000000..9bee4963b01 --- /dev/null +++ b/src/vs/workbench/contrib/mcp/test/node/mcpStdioStateHandler.test.ts @@ -0,0 +1,96 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + * Licensed under the MIT License. See License.txt in the project root for license information. + *--------------------------------------------------------------------------------------------*/ + +import { spawn } from 'child_process'; +import { ensureNoDisposablesAreLeakedInTestSuite } from '../../../../../base/test/common/utils.js'; +import * as assert from 'assert'; +import { McpStdioStateHandler } from '../../node/mcpStdioStateHandler.js'; +import { isWindows } from '../../../../../base/common/platform.js'; + +const GRACE_TIME = 100; + +suite('McpStdioStateHandler', () => { + const store = ensureNoDisposablesAreLeakedInTestSuite(); + + function run(code: string) { + const child = spawn('node', ['-e', code], { + stdio: 'pipe', + env: { ...process.env, ELECTRON_RUN_AS_NODE: '1' }, + }); + + return { + child, + handler: store.add(new McpStdioStateHandler(child, GRACE_TIME)), + processId: new Promise((resolve) => { + child.on('spawn', () => resolve(child.pid!)); + }), + output: new Promise((resolve) => { + let output = ''; + child.stderr.setEncoding('utf-8').on('data', (data) => { + output += data.toString(); + }); + child.stdout.setEncoding('utf-8').on('data', (data) => { + output += data.toString(); + }); + child.on('close', () => resolve(output)); + }), + }; + } + + test('stdin ends process', async () => { + const { child, handler, output } = run(` + const data = require('fs').readFileSync(0, 'utf-8'); + process.stdout.write('Data received: ' + data); + process.on('SIGTERM', () => process.stdout.write('SIGTERM received')); + `); + + child.stdin.write('Hello MCP!'); + handler.stop(); + const result = await output; + assert.strictEqual(result.trim(), 'Data received: Hello MCP!'); + }); + + if (!isWindows) { + test('sigterm after grace', async () => { + const { handler, output } = run(` + setInterval(() => {}, 1000); + process.stdin.on('end', () => process.stdout.write('stdin ended\\n')); + process.stdin.resume(); + process.on('SIGTERM', () => { + process.stdout.write('SIGTERM received', () => process.exit(0)); + }); + `); + + const before = Date.now(); + handler.stop(); + const result = await output; + const delay = Date.now() - before; + assert.strictEqual(result.trim(), 'stdin ended\nSIGTERM received'); + assert.ok(delay >= GRACE_TIME, `Expected at least ${GRACE_TIME}ms delay, got ${delay}ms`); + }); + } + + test('sigkill after grace', async () => { + const { handler, output } = run(` + setInterval(() => {}, 1000); + process.stdin.on('end', () => process.stdout.write('stdin ended\\n')); + process.stdin.resume(); + process.on('SIGTERM', () => { + process.stdout.write('SIGTERM received'); + }); + `); + + const before = Date.now(); + handler.stop(); + const result = await output; + const delay = Date.now() - before; + if (!isWindows) { + assert.strictEqual(result.trim(), 'stdin ended\nSIGTERM received'); + } else { + assert.strictEqual(result.trim(), 'stdin ended'); + } + assert.ok(delay >= GRACE_TIME * 2, `Expected at least ${GRACE_TIME * 2}ms delay, got ${delay}ms`); + }); +});