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() {