From 4633a0abbfe19e0d379eadd330346f10747ea39b Mon Sep 17 00:00:00 2001 From: Adarsh Choudhary Date: Sun, 27 Sep 2026 20:22:30 +0200 Subject: [PATCH] fix(client): wait for the whole server process tree to exit on stdio dispose KillTree killed the entire process tree but only waited on the root process. On Windows the root is the cmd.exe /c wrapper, so DisposeAsync could return while the actual server (and conhost.exe) were still terminating and holding the working directory and open files. Snapshot the root's descendants before killing the tree, then wait on each of them within the same ShutdownTimeout. Descendants are discovered with a Toolhelp32 snapshot on Windows, /proc on Linux and proc_listchildpids on macOS. Start times guard against stale parent IDs from reused PIDs. Fixes part 1 of #1894. --- .../ModelContextProtocol.Core.csproj | 3 +- .../ProcessHelper.cs | 289 ++++++++++++++++-- .../Transport/StdioClientTransportTests.cs | 73 +++++ 3 files changed, 339 insertions(+), 26 deletions(-) diff --git a/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj b/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj index 3fbef0377..59c8cfc1e 100644 --- a/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj +++ b/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj @@ -15,6 +15,8 @@ the implementation, so the obsolete-usage diagnostic is suppressed project-wide here while staying active for external consumers of the package. --> $(NoWarn);MCP9005 + + true @@ -24,7 +26,6 @@ $(NoWarn);CS0436 - true diff --git a/src/ModelContextProtocol.Core/ProcessHelper.cs b/src/ModelContextProtocol.Core/ProcessHelper.cs index 44736a308..4063d7c13 100644 --- a/src/ModelContextProtocol.Core/ProcessHelper.cs +++ b/src/ModelContextProtocol.Core/ProcessHelper.cs @@ -16,53 +16,292 @@ internal static class ProcessHelper /// /// On .NET Core 3.0+ this uses Process.Kill(entireProcessTree: true). /// On .NET Standard 2.0, it uses platform-specific commands (taskkill on Windows, pgrep/kill on Unix). - /// The method waits for the specified timeout for processes to exit before continuing. + /// The method waits for the specified timeout for the process and its descendants to exit before continuing. /// This is particularly useful for applications that spawn child processes (like Node.js) /// that wouldn't be terminated automatically when the parent process exits. /// public static void KillTree(this Process process, TimeSpan timeout) { + // Snapshot the descendants before killing anything, as the parent/child links needed to find + // them are gone once the tree has been torn down. On Windows the root is usually the cmd.exe /c + // wrapper, so the actual server is one of these descendants and waiting on the root alone isn't + // enough to know that the server has released its working directory and files. + List descendants = GetDescendantProcesses(process); + try + { #if NETSTANDARD2_0 - // Process.Kill(entireProcessTree) is not available on .NET Standard 2.0. - // Use platform-specific commands to kill the process tree. - var pid = process.Id; - if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) + // Process.Kill(entireProcessTree) is not available on .NET Standard 2.0. + // Use platform-specific commands to kill the process tree. + var pid = process.Id; + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) + { + RunProcessAndWaitForExit( + "taskkill", + $"/T /F /PID {pid}", + timeout, + out var _); + } + else + { + var children = new HashSet(); + GetAllChildIdsUnix(pid, children, timeout); + foreach (var childId in children) + { + KillProcessUnix(childId, timeout); + } + + KillProcessUnix(pid, timeout); + } +#else + try + { + process.Kill(entireProcessTree: true); + } + catch + { + // Process has already exited + return; + } +#endif + + // wait until the processes finish exiting/getting killed. + // We don't want to wait forever here because the tree is already supposed to be dying, we just want to give it long enough + // to try and flush what it can and stop. If it cannot do that in a reasonable time frame then we will just ignore it. + // The root and all of its descendants share the same timeout. + Stopwatch stopwatch = Stopwatch.StartNew(); + process.WaitForExit(GetRemainingMilliseconds(timeout, stopwatch)); + foreach (Process descendant in descendants) + { + try + { + descendant.WaitForExit(GetRemainingMilliseconds(timeout, stopwatch)); + } + catch + { + // The process can no longer be waited on, e.g. because it has already exited. + } + } + } + finally { - RunProcessAndWaitForExit( - "taskkill", - $"/T /F /PID {pid}", - timeout, - out var _); + foreach (Process descendant in descendants) + { + descendant.Dispose(); + } } - else + } + + private static int GetRemainingMilliseconds(TimeSpan timeout, Stopwatch stopwatch) => + timeout == Timeout.InfiniteTimeSpan ? Timeout.Infinite : + (int)Math.Max(0, (timeout - stopwatch.Elapsed).TotalMilliseconds); + + /// + /// Gets all currently running descendants of , or as many of them as can be found. + /// + private static List GetDescendantProcesses(Process root) + { + List descendants = []; + try { - var children = new HashSet(); - GetAllChildIdsUnix(pid, children, timeout); - foreach (var childId in children) + Func> getChildIds; + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows) || RuntimeInformation.IsOSPlatform(OSPlatform.Linux)) { - KillProcessUnix(childId, timeout); + // Both platforms list every process along with its parent ID in one pass. + ILookup childIds = + (RuntimeInformation.IsOSPlatform(OSPlatform.Windows) ? GetParentIdsWindows() : GetParentIdsLinux()) + .ToLookup(p => p.ParentId, p => p.Id); + getChildIds = id => childIds[id]; + } + else if (RuntimeInformation.IsOSPlatform(OSPlatform.OSX)) + { + getChildIds = GetChildIdsMacOS; + } + else + { + return descendants; } - KillProcessUnix(pid, timeout); + HashSet visited = [root.Id]; + Queue parents = new(); + parents.Enqueue(root); + while (parents.Count > 0) + { + Process parent = parents.Dequeue(); + foreach (int childId in getChildIds(parent.Id)) + { + if (visited.Add(childId) && TryGetChildProcess(parent, childId) is { } child) + { + descendants.Add(child); + parents.Enqueue(child); + } + } + } } -#else + catch + { + // Finding descendants is best effort. Whatever was found so far will still be waited on. + } + + return descendants; + } + + private static Process? TryGetChildProcess(Process parent, int childId) + { + Process? child = null; try { - process.Kill(entireProcessTree: true); + child = Process.GetProcessById(childId); + + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) + { + // Open and cache a handle to the process so that later operations, such as waiting for it to exit, + // are bound to this process even if it exits and its ID is reused. + try + { + _ = child.SafeHandle; + } + catch + { + } + } + + // A process that started before its supposed parent can't be its child. Its parent ID is stale, + // because the real parent has exited and the ID has since been reused. + if (child.StartTime >= parent.StartTime) + { + return child; + } } catch { - // Process has already exited - return; + // The process has already exited or can't be inspected. } -#endif - // wait until the process finishes exiting/getting killed. - // We don't want to wait forever here because the task is already supposed to be dying, we just want to give it long enough - // to try and flush what it can and stop. If it cannot do that in a reasonable time frame then we will just ignore it. - process.WaitForExit((int)timeout.TotalMilliseconds); + child?.Dispose(); + return null; } + private static IEnumerable<(int Id, int ParentId)> GetParentIdsWindows() + { + List<(int Id, int ParentId)> processes = []; + + nint snapshot = CreateToolhelp32Snapshot(TH32CS_SNAPPROCESS, 0); + if (snapshot == INVALID_HANDLE_VALUE) + { + return processes; + } + + try + { + PROCESSENTRY32W entry = default; + unsafe + { + entry.dwSize = (uint)sizeof(PROCESSENTRY32W); + } + + for (int found = Process32FirstW(snapshot, ref entry); found != 0; found = Process32NextW(snapshot, ref entry)) + { + processes.Add(((int)entry.th32ProcessID, (int)entry.th32ParentProcessID)); + } + } + finally + { + _ = CloseHandle(snapshot); + } + + return processes; + } + + private static IEnumerable<(int Id, int ParentId)> GetParentIdsLinux() + { + foreach (string directory in Directory.EnumerateDirectories("/proc")) + { + if (!int.TryParse(Path.GetFileName(directory), out int id)) + { + continue; + } + + string stat; + try + { + stat = File.ReadAllText(Path.Combine(directory, "stat")); + } + catch + { + // The process has exited since the directory was enumerated. + continue; + } + + // The format is "pid (comm) state ppid ...". comm can contain spaces and parentheses, so parse from the last ')'. + int commEnd = stat.LastIndexOf(')'); + if (commEnd < 0) + { + continue; + } + + string[] fields = stat.Substring(commEnd + 1).Split([' '], StringSplitOptions.RemoveEmptyEntries); + if (fields.Length > 1 && int.TryParse(fields[1], out int parentId)) + { + yield return (id, parentId); + } + } + } + + private static unsafe IEnumerable GetChildIdsMacOS(int parentId) + { + // Calling with no buffer returns an upper bound for the number of children. + int count = proc_listchildpids(parentId, null, 0); + if (count <= 0) + { + return []; + } + + int[] ids = new int[count]; + fixed (int* pIds = ids) + { + count = proc_listchildpids(parentId, pIds, ids.Length * sizeof(int)); + } + + return ids.Take(count); + } + + private const uint TH32CS_SNAPPROCESS = 0x00000002; + private const nint INVALID_HANDLE_VALUE = -1; + + [StructLayout(LayoutKind.Sequential)] + private unsafe struct PROCESSENTRY32W + { + public uint dwSize; + public uint cntUsage; + public uint th32ProcessID; + public nuint th32DefaultHeapID; + public uint th32ModuleID; + public uint cntThreads; + public uint th32ParentProcessID; + public int pcPriClassBase; + public uint dwFlags; + public fixed char szExeFile[260]; + } + + [DllImport("kernel32.dll")] + [DefaultDllImportSearchPaths(DllImportSearchPath.System32)] + private static extern nint CreateToolhelp32Snapshot(uint dwFlags, uint th32ProcessID); + + [DllImport("kernel32.dll")] + [DefaultDllImportSearchPaths(DllImportSearchPath.System32)] + private static extern int Process32FirstW(nint hSnapshot, ref PROCESSENTRY32W lppe); + + [DllImport("kernel32.dll")] + [DefaultDllImportSearchPaths(DllImportSearchPath.System32)] + private static extern int Process32NextW(nint hSnapshot, ref PROCESSENTRY32W lppe); + + [DllImport("kernel32.dll")] + [DefaultDllImportSearchPaths(DllImportSearchPath.System32)] + private static extern int CloseHandle(nint hObject); + + [DllImport("libproc")] + private static extern unsafe int proc_listchildpids(int ppid, int* buffer, int buffersize); + #if NETSTANDARD2_0 private static void GetAllChildIdsUnix(int parentId, ISet children, TimeSpan timeout) { diff --git a/tests/ModelContextProtocol.Tests/Transport/StdioClientTransportTests.cs b/tests/ModelContextProtocol.Tests/Transport/StdioClientTransportTests.cs index 60ce9cf5a..2fb691037 100644 --- a/tests/ModelContextProtocol.Tests/Transport/StdioClientTransportTests.cs +++ b/tests/ModelContextProtocol.Tests/Transport/StdioClientTransportTests.cs @@ -2,6 +2,7 @@ using ModelContextProtocol.Client; using ModelContextProtocol.Protocol; using ModelContextProtocol.Tests.Utils; +using System.Diagnostics; using System.IO.Pipelines; using System.Runtime.InteropServices; using System.Text; @@ -305,6 +306,78 @@ void CaptureLines(string line) Assert.Contains("EXPLICIT_IS_SET", allOutput); } + [Fact] + public async Task DisposeAsync_WaitsForDescendantProcessesToExit() + { + // The server spawns a long-running descendant and reports its PID. On Windows the descendant sits + // below the cmd.exe /c wrapper, and on Unix below the shell, so it's never the process the transport started. + string pidFile = Path.Combine(Path.GetTempPath(), $"mcp-test-{Guid.NewGuid():N}.pid"); + try + { + StdioClientTransport transport = RuntimeInformation.IsOSPlatform(OSPlatform.Windows) ? + new(new() + { + Command = "powershell", + Arguments = ["-NoProfile", "-NonInteractive", "-Command", $"Set-Content -LiteralPath '{pidFile}' -Value $PID; Start-Sleep -Seconds 60"], + ShutdownTimeout = TimeSpan.FromSeconds(2), + }, LoggerFactory) : + new(new() + { + Command = "sh", + Arguments = ["-c", $"sleep 60 & echo $! > '{pidFile}'; wait"], + ShutdownTimeout = TimeSpan.FromSeconds(2), + }, LoggerFactory); + + var sessionTransport = await transport.ConnectAsync(TestContext.Current.CancellationToken); + + int descendantId = await ReadPidFileAsync(pidFile); + Assert.True(IsProcessRunning(descendantId), "The descendant process should be running before dispose."); + + await sessionTransport.DisposeAsync(); + + Assert.False(IsProcessRunning(descendantId), "The descendant process should have exited by the time dispose returns."); + } + finally + { + try { File.Delete(pidFile); } catch { } + } + + static async Task ReadPidFileAsync(string path) + { + var deadline = DateTime.UtcNow + TestConstants.DefaultTimeout; + while (true) + { + try + { + if (int.TryParse(File.ReadAllText(path).Trim(), out int pid)) + { + return pid; + } + } + catch (IOException) when (DateTime.UtcNow < deadline) + { + } + + Assert.True(DateTime.UtcNow < deadline, "Timed out waiting for the descendant process to report its PID."); + await Task.Delay(50, TestContext.Current.CancellationToken); + } + } + + static bool IsProcessRunning(int pid) + { + try + { + using var process = Process.GetProcessById(pid); + return !process.HasExited; + } + catch (ArgumentException) + { + // No process with this ID exists. + return false; + } + } + } + [Fact] public void GetDefaultEnvironmentVariables_ReturnsFreshDictionaryEachCall() {