refactor(tools/mcp): split mcp_tool.py into transport/lifecycle/schema/handlers/... sibling modules; compact watchdog and schema cache
This commit is contained in:
@@ -0,0 +1,342 @@
|
||||
"""MCP process lifecycle: stdio child PID tracking and orphan cleanup, graceful
|
||||
server shutdown and draining of the background MCP loop."""
|
||||
|
||||
import logging
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from typing import Dict, Optional
|
||||
from tools.mcp_tool_common import _core
|
||||
|
||||
logger = logging.getLogger("tools.mcp_tool")
|
||||
|
||||
|
||||
# Live stdio MCP children (pid -> server_name), added after connection and
|
||||
# removed on normal shutdown, so they can be force-killed if SDK teardown fails.
|
||||
_stdio_pids: Dict[int, str] = {}
|
||||
|
||||
|
||||
# PIDs that survived their session context exit (SDK teardown failed to kill
|
||||
# them); detected in _run_stdio's finally, reaped by _kill_orphaned_mcp_children().
|
||||
# Kept separate from _stdio_pids so cleanup sweeps never race active sessions.
|
||||
_orphan_stdio_pids: set = set()
|
||||
|
||||
|
||||
_orphan_stdio_pid_servers: Dict[int, str] = {}
|
||||
|
||||
|
||||
# pid -> pgid captured at spawn. The SDK spawns children with
|
||||
# start_new_session=True (PGID == PID); grandchildren inherit that PGID and
|
||||
# keep it after the direct child exits, so killpg still reaches them. Tracked
|
||||
# separately from _stdio_pids so the PGID survives the child's removal.
|
||||
# Empty on Windows (os.getpgid is POSIX-only).
|
||||
_stdio_pgids: Dict[int, int] = {}
|
||||
|
||||
|
||||
def _snapshot_child_pids() -> set:
|
||||
"""Current direct-child PIDs: /proc on Linux, else psutil, else empty set."""
|
||||
my_pid = os.getpid()
|
||||
|
||||
# /proc/<pid>/task/<tid>/children is per-THREAD, and stdio_client() spawns
|
||||
# from the MCP loop thread, so union every task's children — reading only
|
||||
# the main thread's file returns an empty set on every Linux install.
|
||||
try:
|
||||
task_dir = f"/proc/{my_pid}/task"
|
||||
tids = os.listdir(task_dir)
|
||||
found: set = set()
|
||||
for tid in tids:
|
||||
try:
|
||||
with open(f"{task_dir}/{tid}/children", encoding="utf-8") as f:
|
||||
found.update(int(p) for p in f.read().split() if p.strip())
|
||||
except (FileNotFoundError, OSError, ValueError):
|
||||
continue # thread exited between listdir and open
|
||||
return found
|
||||
except (FileNotFoundError, OSError, ValueError):
|
||||
pass
|
||||
|
||||
try:
|
||||
import psutil
|
||||
return {c.pid for c in psutil.Process(my_pid).children()}
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return set()
|
||||
|
||||
|
||||
# argv markers of non-MCP gateway children that can race into the snapshot
|
||||
# delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use
|
||||
# start_new_session). Matched against argv[1:] because Python/Java children
|
||||
# start with the interpreter path.
|
||||
_NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = (
|
||||
"tui_gateway.slash_worker",
|
||||
"tui_gateway.entry",
|
||||
"-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style)
|
||||
"eclipse.jdt.ls",
|
||||
"org.eclipse.equinox.launcher_",
|
||||
)
|
||||
|
||||
|
||||
def _filter_mcp_children(pids: set) -> set:
|
||||
"""Drop non-MCP children from a PID snapshot delta.
|
||||
|
||||
Tracking a stray child in _stdio_pgids is catastrophic if it lacks
|
||||
start_new_session: its pgid can be the TUI parent's, so the shutdown
|
||||
killpg() would kill the TUI itself.
|
||||
"""
|
||||
if not pids:
|
||||
return pids
|
||||
try:
|
||||
import psutil
|
||||
except ImportError:
|
||||
return pids # keep all PIDs (prior behavior)
|
||||
filtered: set = set()
|
||||
for pid in pids:
|
||||
try:
|
||||
argv = psutil.Process(pid).cmdline()
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied, OSError):
|
||||
# Raced away or zombie — cannot be our fresh server, unsafe to track.
|
||||
continue
|
||||
if any(
|
||||
marker in arg
|
||||
for arg in argv[1:]
|
||||
for marker in _NON_MCP_CHILD_CMDLINE_MARKERS
|
||||
):
|
||||
continue
|
||||
filtered.add(pid)
|
||||
return filtered
|
||||
|
||||
|
||||
def shutdown_mcp_servers(*, scope: Optional[str] = None):
|
||||
"""Close MCP server connections (in parallel) and stop the background loop.
|
||||
|
||||
Each server Task is signalled to exit its own ``async with`` so the anyio
|
||||
cancel-scope cleanup runs in the Task that opened it. ``scope`` restricts
|
||||
teardown to one multiplexed profile's servers (its ``/reload-mcp`` must not
|
||||
kill other profiles') and leaves the shared loop running if anything else
|
||||
is still connected.
|
||||
"""
|
||||
with _core._lock:
|
||||
selected = [
|
||||
name for name in _core._servers
|
||||
if scope is None or _core._server_scope_keys.get(name) == scope
|
||||
]
|
||||
servers_snapshot = [_core._servers[name] for name in selected]
|
||||
|
||||
# Fast path: nothing to shut down. Still clear the connect-cooldown maps —
|
||||
# a server that failed to connect is never in ``_servers``, so this is the
|
||||
# most likely state for stale backoff entries; a restart must retry at once.
|
||||
if not servers_snapshot:
|
||||
with _core._lock:
|
||||
_core._server_connect_retry_after.clear()
|
||||
_core._server_connect_failures.clear()
|
||||
_core._stop_mcp_loop(only_if_idle=scope is not None)
|
||||
return
|
||||
|
||||
async def _shutdown():
|
||||
results = await asyncio.gather(
|
||||
*(server.shutdown() for server in servers_snapshot),
|
||||
return_exceptions=True,
|
||||
)
|
||||
for server, result in zip(servers_snapshot, results):
|
||||
if isinstance(result, Exception):
|
||||
logger.debug(
|
||||
"Error closing MCP server '%s': %s", server.name, result,
|
||||
)
|
||||
with _core._lock:
|
||||
for name in selected:
|
||||
_core._servers.pop(name, None)
|
||||
_core._server_scope_keys.pop(name, None)
|
||||
# Drop connect-retry cooldowns too: a restart must re-attempt every
|
||||
# server immediately, not honour a stale per-server backoff.
|
||||
_core._server_connect_retry_after.clear()
|
||||
_core._server_connect_failures.clear()
|
||||
|
||||
with _core._lock:
|
||||
loop = _core._mcp_loop
|
||||
if loop is not None and loop.is_running():
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
future = safe_schedule_threadsafe(
|
||||
_shutdown(), loop,
|
||||
logger=logger,
|
||||
log_message="MCP shutdown: failed to schedule",
|
||||
)
|
||||
if future is not None:
|
||||
try:
|
||||
future.result(timeout=15)
|
||||
except BaseException as exc:
|
||||
logger.debug("Error during MCP shutdown: %s", exc)
|
||||
|
||||
# Unconditional final sweep: whether ``_shutdown`` ran, timed out, or was
|
||||
# never scheduled, no stale connect-cooldown state may survive shutdown.
|
||||
with _core._lock:
|
||||
_core._server_connect_retry_after.clear()
|
||||
_core._server_connect_failures.clear()
|
||||
|
||||
_core._stop_mcp_loop(only_if_idle=scope is not None)
|
||||
|
||||
|
||||
def _kill_orphaned_mcp_children(
|
||||
include_active: bool = False,
|
||||
server_name: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors.
|
||||
|
||||
By default only ``_orphan_stdio_pids`` (PIDs that outlived their session
|
||||
context) are reaped so concurrent cron jobs / live sessions are untouched;
|
||||
``include_active=True`` also kills every ``_stdio_pids`` entry and is only
|
||||
for final shutdown after the MCP loop has stopped. ``server_name`` limits
|
||||
the sweep to one server (stdio reconnects cleaning up their old transport).
|
||||
|
||||
On POSIX signals go via ``os.killpg`` to the spawn-time pgid when tracked,
|
||||
so reparented grandchildren are reaped too; falls back to ``os.kill``.
|
||||
"""
|
||||
import signal as _signal
|
||||
|
||||
with _core._lock:
|
||||
pids: Dict[int, str] = {}
|
||||
for opid in _orphan_stdio_pids:
|
||||
owner = _orphan_stdio_pid_servers.get(opid, "orphan")
|
||||
if server_name is not None and owner != server_name:
|
||||
continue
|
||||
pids[opid] = owner
|
||||
for opid in pids:
|
||||
_orphan_stdio_pids.discard(opid)
|
||||
_orphan_stdio_pid_servers.pop(opid, None)
|
||||
if include_active:
|
||||
active = dict(_stdio_pids)
|
||||
if server_name is not None:
|
||||
active = {
|
||||
pid: owner
|
||||
for pid, owner in active.items()
|
||||
if owner == server_name
|
||||
}
|
||||
pids.update(active)
|
||||
for pid in active:
|
||||
_stdio_pids.pop(pid, None)
|
||||
# Snapshot pgids for the pids we're about to kill, then drop them so a
|
||||
# future spawn can't collide with stale state.
|
||||
pgids: Dict[int, int] = {pid: _stdio_pgids[pid] for pid in pids if pid in _stdio_pgids}
|
||||
for pid in pgids:
|
||||
_stdio_pgids.pop(pid, None)
|
||||
|
||||
# Fast path: nothing to reap — skip the 2s sleep every MCP-free shutdown
|
||||
# would otherwise pay.
|
||||
if not pids:
|
||||
return
|
||||
|
||||
# Our own pgid, so _send_signal never killpg()s the gateway itself.
|
||||
try:
|
||||
_my_pgid = os.getpgrp()
|
||||
except (AttributeError, OSError):
|
||||
_my_pgid = None # Windows or restricted environment
|
||||
|
||||
def _send_signal(pid: int, sig: int, server_name: str) -> None:
|
||||
"""SIGTERM/SIGKILL via pgroup on POSIX, fall back to pid signal."""
|
||||
pgid = pgids.get(pid)
|
||||
killpg = getattr(os, "killpg", None)
|
||||
if pgid is not None and killpg is not None:
|
||||
if _my_pgid is not None and pgid == _my_pgid:
|
||||
# Child shares the gateway's pgroup: killpg would kill the
|
||||
# gateway too, so use per-pid kill. Warn because per-pid kill
|
||||
# can't reach grandchildren in this group (inherent trade-off).
|
||||
logger.warning(
|
||||
"MCP server '%s' pgid %d matches gateway pgid; skipping "
|
||||
"killpg to avoid self-kill and using per-pid kill — any "
|
||||
"grandchildren in this group may not be reaped",
|
||||
server_name, pgid,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
killpg(pgid, sig)
|
||||
return
|
||||
except (ProcessLookupError, PermissionError, OSError) as exc:
|
||||
# Pgroup gone or refused — still try the direct child.
|
||||
logger.debug(
|
||||
"killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)",
|
||||
pgid, sig, server_name, exc,
|
||||
)
|
||||
try:
|
||||
os.kill(pid, sig)
|
||||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pass
|
||||
|
||||
for pid, server_name in pids.items():
|
||||
_send_signal(pid, _signal.SIGTERM, server_name)
|
||||
logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, server_name)
|
||||
|
||||
time.sleep(2)
|
||||
|
||||
_sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM)
|
||||
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the portable check.
|
||||
from gateway.status import _pid_exists
|
||||
for pid, server_name in pids.items():
|
||||
if not _pid_exists(pid):
|
||||
continue # exited after SIGTERM
|
||||
_send_signal(pid, _sigkill, server_name)
|
||||
logger.warning(
|
||||
"Force-killed MCP process %d (%s) after SIGTERM timeout",
|
||||
pid, server_name,
|
||||
)
|
||||
|
||||
|
||||
def _stop_mcp_loop_if_idle() -> bool:
|
||||
"""Stop the MCP loop only when no registered server still owns it.
|
||||
|
||||
Probe paths create temporary MCPServerTasks not placed in ``_servers``;
|
||||
they may clean up an idle loop but must not tear down the process-global
|
||||
loop under live agent tools, or later calls fail with
|
||||
``MCP event loop is not running``.
|
||||
"""
|
||||
return _core._stop_mcp_loop(only_if_idle=True)
|
||||
|
||||
|
||||
async def _drain_mcp_loop_tasks(
|
||||
*,
|
||||
timeout: Optional[float] = None,
|
||||
) -> None:
|
||||
"""Cancel every task still pending on the MCP loop and reap it.
|
||||
|
||||
``Task.cancel()`` only schedules the throw, so tasks need a cancellation
|
||||
cycle before the loop goes away; wait for them here, on their owning loop,
|
||||
bounded so a task that suppresses cancellation cannot hang process exit.
|
||||
"""
|
||||
if timeout is None:
|
||||
timeout = _core._MCP_LOOP_DRAIN_TIMEOUT
|
||||
current = asyncio.current_task()
|
||||
pending = [t for t in asyncio.all_tasks() if t is not current and not t.done()]
|
||||
if not pending:
|
||||
return
|
||||
logger.debug("Draining %d pending task(s) from the MCP loop", len(pending))
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
|
||||
done, still_pending = await asyncio.wait(pending, timeout=timeout)
|
||||
for task in done:
|
||||
if task.cancelled():
|
||||
continue
|
||||
try:
|
||||
task.exception()
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
except Exception as exc:
|
||||
logger.debug("Pending MCP loop task ended during shutdown: %s", exc)
|
||||
|
||||
if still_pending:
|
||||
logger.warning(
|
||||
"%d MCP loop task(s) still pending after %.1fs drain",
|
||||
len(still_pending), timeout,
|
||||
)
|
||||
|
||||
|
||||
async def _drain_and_stop_mcp_loop() -> None:
|
||||
"""Drain pending tasks, then stop the loop from its owning thread.
|
||||
|
||||
Both must run as one loop-owned sequence: a ``loop.stop`` queued separately
|
||||
by a timed-out caller can overtake the scheduled drain, leaving the drain
|
||||
coroutine itself pending when the loop is closed.
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await _drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)
|
||||
finally:
|
||||
loop.call_soon(loop.stop)
|
||||
Reference in New Issue
Block a user