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