diff --git a/agent/deadline.py b/agent/deadline.py index 5aa58e6c06..fa6bf2a7da 100644 --- a/agent/deadline.py +++ b/agent/deadline.py @@ -542,3 +542,32 @@ def kill_process_tree(pid: int, *, sig: Optional[int] = None) -> bool: except Exception: continue return signalled + + +class SuspectableBackend: + """Protocol for backends whose connection state can be *poisoned* by a + race (teardown-vs-keepalive, auth-lock corruption) without the backend + itself being dead. + + The contract is **cheap-mark, lazy-verify**: noticing a poisoned state + must never do I/O — ``mark_suspect`` just latches a reason string. The + NEXT caller pays for verification once, via ``ensure_healthy``: a cheap + health probe that either clears the suspicion (backend was fine) or + forces a reconnect/recycle before the call proceeds. This is what keeps + a single race from permanently parking a connection (#81051/#77765/ + #84132): instead of parking on the ambiguous event, the backend is + marked suspect and recycled exactly once on next use. + """ + + def mark_suspect(self, reason: str) -> None: + """Latch a suspicion about this backend. Must be cheap (no I/O).""" + raise NotImplementedError + + async def ensure_healthy(self, timeout: float = 5.0) -> bool: + """Verify a suspect backend before reuse. + + Returns True when the backend is healthy (clearing the suspicion); + returns False after forcing a reconnect/recycle so the caller's + normal no-session path handles the rebuild. Must not raise. + """ + raise NotImplementedError diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 4952c11a6e..8393538247 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -110,6 +110,7 @@ import shutil import sys import threading import time +from contextlib import asynccontextmanager from types import SimpleNamespace from typing import Callable from datetime import datetime @@ -2393,6 +2394,8 @@ class MCPServerTask: "_idle_timeout_seconds", "_max_lifetime_seconds", "_recycled_reason", "initialize_result", "_ping_unsupported", "_list_cache_meta", "_reconnect_retries", "_session_proven", "_was_parked", + "_inflight_tasks", "_reconnecting", "_suspect_reason", + "_teardown_race", "_permanent_grace_used", "_stdio_child_pids", ) def __init__(self, name: str): @@ -2427,6 +2430,30 @@ class MCPServerTask: # until the session proves healthy again — used to log the # parked→revived transition exactly once. self._was_parked: bool = False + # In-flight RPC bookkeeping (#48069 salvage): user-visible requests + # registered while running so a reconnect/shutdown teardown can fail + # them fast instead of orphaning them on a dying transport. + self._inflight_tasks: set = set() + # True while a deliberate teardown is failing in-flight calls — lets + # _track_inflight_rpc convert the cancel into a retryable error. + self._reconnecting: bool = False + # SuspectableBackend state (#81051/#77765/#84132): latched by races + # (teardown-vs-keepalive, auth-lock corruption); verified lazily by + # ensure_healthy() before the next call reuses the connection. + self._suspect_reason: Optional[str] = None + # Set when a teardown failed >=1 in-flight call: the following + # reconnect is a RACE RECOVERY, not a transport failure, and must not + # charge the rapid-drop budget (a single race must never reach park). + self._teardown_race: bool = False + # One-time grace: an auth/permanent-classified failure on a previously + # PROVEN session gets one suspect+reconnect cycle before the park + # ladder applies (single auth-lock corruption must not park). + self._permanent_grace_used: bool = False + # PIDs of the stdio subprocess spawned for the current transport + # (captured in _run_stdio). Used to fail in-flight calls FAST when + # the child dies instead of waiting out the full tool timeout + # (#81995). + self._stdio_child_pids: Set[int] = set() self._auth_type: str = "" self._refresh_lock = asyncio.Lock() # MCP stdio sessions are a single JSON-RPC stream. Some servers emit @@ -2869,6 +2896,118 @@ class MCPServerTask: "parking (state: parked → connected)", self.name, ) + # A session that just proved healthy on a fresh transport clears + # the one-time permanent-failure grace and any race bookkeeping. + self._permanent_grace_used = False + self._teardown_race = False + + # -- SuspectableBackend contract (agent.deadline) ----------------------- + + def mark_suspect(self, reason: str) -> None: + """Latch a suspicion about this connection. Cheap — no I/O. + + The NEXT call verifies via :meth:`ensure_healthy` and recycles the + transport if the probe fails, instead of the connection silently + staying poisoned until process restart (#81051/#77765/#84132). + """ + if self._suspect_reason is None and reason: + logger.warning( + "MCP server '%s': connection marked suspect (%s); next call " + "will health-check it", + self.name, reason, + ) + self._suspect_reason = reason or None + + async def ensure_healthy(self, timeout: float = 5.0) -> bool: + """Verify a suspect connection before reuse; recycle if dead. + + Returns True when healthy (suspicion cleared). On failure, requests a + reconnect, drops the stale session reference so the caller's normal + no-session path takes over, and returns False. Never raises. + """ + reason = self._suspect_reason + if not reason: + return True + if self.session is None: + # Nothing to verify — the reconnect path owns recovery now. + self._suspect_reason = None + self._reconnect_event.set() + return False + try: + await asyncio.wait_for(self._keepalive_probe(), timeout=timeout) + except Exception as exc: + root = _unwrap_exception_group(exc) + logger.warning( + "MCP server '%s': suspect connection (%s) failed health " + "check (%s: %s) — requesting reconnect (state: suspect → " + "degraded)", + self.name, reason, type(root).__name__, root, + ) + self._suspect_reason = None + self.mark_suspect(f"health check failed after {reason}") + self.session = None + self._ready.clear() + self._reconnect_event.set() + return False + logger.info( + "MCP server '%s': suspect connection passed health check " + "(%s) — clearing suspicion", + self.name, reason, + ) + self._suspect_reason = None + self._mark_session_proven() + return True + + def _fail_inflight_calls(self, reason: str) -> None: + """Cancel every in-flight RPC attached to this connection. + + Called from the lifecycle exits (reconnect/shutdown/recycle) BEFORE + the transport unwinds: the MCP SDK does not always fail pending + requests when its streams close, so without this an in-flight call + would wait out the full tool timeout on a dying transport. Cancelling + at least one task flags the cycle as a teardown race + (``_teardown_race``) so run() treats the following reconnect as + recovery rather than charging the rapid-drop budget. + """ + victims = [t for t in self._inflight_tasks if not t.done()] + if not victims: + return + self._reconnecting = True + self._teardown_race = True + self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)") + for task in victims: + task.cancel() + + def _stdio_children_dead(self) -> bool: + """True when every stdio child we spawned has exited. + + Best-effort: only meaningful for stdio transports with captured PIDs; + returns False (unknown → don't fail fast) otherwise. + """ + pids = getattr(self, "_stdio_child_pids", None) + if not pids or self._is_http(): + return False + for pid in pids: + # windows-footgun: ok — psutil.pid_exists handles Windows; the + # os.kill probe below only runs when psutil is unavailable. + import psutil + + if not psutil.pid_exists(pid): + continue # this one is dead + return True # alive (signal permission irrelevant for liveness) + return False # at least one child alive + return True + + async def _watch_stdio_children(self) -> None: + """Poll child liveness while a stdio RPC is in flight (#81995). + + Resolves when a tracked child dies; the caller then cancels the RPC + immediately instead of letting it hang for the full tool timeout. + """ + while True: + if self._stdio_children_dead(): + return + await asyncio.sleep(0.25) async def _wait_for_lifecycle_event(self) -> str: """Block until either _shutdown_event or _reconnect_event fires. @@ -2938,16 +3077,21 @@ class MCPServerTask: return "recycle" # Timeout — no lifecycle event fired. Probe the connection - # to detect stale/expired sessions. Prefer ``ping`` (MCP base - # protocol liveness): it works uniformly and stays a few bytes - # regardless of tool count, unlike ``list_tools`` (~1 MB on an - # 830-tool server). ``ping`` is an OPTIONAL utility, so a - # tool-capable server that doesn't implement it answers -32601; - # in that case fall back to the pre-ping ``list_tools`` probe - # for the rest of this connection rather than reconnect-looping. + # to detect stale/expired sessions — but NEVER while an RPC + # is in flight (#48069): the stdio session is a single + # JSON-RPC stream and a concurrent ping/list_tools can wedge + # the in-flight request. A busy server is provably alive. if self.session: + if self._rpc_lock.locked() or any( + not t.done() for t in self._inflight_tasks + ): + continue try: - await self._keepalive_probe() + async def _probe_under_lock(): + async with self._rpc_lock: + await self._keepalive_probe() + + await _probe_under_lock() except Exception as exc: root = _unwrap_exception_group(exc) logger.warning( @@ -2955,6 +3099,9 @@ class MCPServerTask: "reconnect (state: connected → degraded): %s: %s", self.name, type(root).__name__, root, ) + self.mark_suspect( + f"keepalive failed: {type(root).__name__}: {root}" + ) self._reconnect_event.set() break # Keepalive succeeded — the session survived a full @@ -2971,7 +3118,11 @@ class MCPServerTask: pass if self._shutdown_event.is_set(): + self._fail_inflight_calls("shutdown") return "shutdown" + # Deliberate teardown: fail any in-flight RPC NOW so it doesn't ride + # the dying transport to the full tool timeout (#48069/#81995). + self._fail_inflight_calls("reconnect") self._reconnect_event.clear() return "reconnect" @@ -3156,6 +3307,10 @@ class MCPServerTask: for _pid in new_pids: _stdio_pids[_pid] = self.name _stdio_pgids.update(new_pgids) + # Track the spawned children on the connection object for + # fast-fail of in-flight calls when the subprocess dies + # (#81995). + self._stdio_child_pids = set(new_pids) async with ClientSession( read_stream, write_stream, **sampling_kwargs ) as session: @@ -3873,7 +4028,22 @@ class MCPServerTask: # Only clear the consecutive-failure budget once the session # PROVED healthy — survived >=1 full keepalive interval or # served >=1 successful tool call (_mark_session_proven). - if self._session_proven: + if self._teardown_race and not self._session_proven: + # The previous cycle ended because a teardown cancelled + # in-flight calls (keepalive/refresh race, auth recovery) + # — that is RECOVERY, not a transport failure. Do NOT + # charge the rapid-drop budget: a single race must never + # reach the park (#81051/#77765/#84132). Only genuinely + # repeated unproven drops still exhaust the budget below. + logger.info( + "MCP server '%s': reconnect after teardown race " + "(in-flight calls were failed); not charging the " + "rapid-drop budget", + self.name, + ) + self._teardown_race = False + backoff = 1.0 + elif self._session_proven: self._reconnect_retries = 0 backoff = 1.0 else: @@ -4060,6 +4230,36 @@ class MCPServerTask: return if failure_class == "permanent": + # Auth-lock corruption guard (#81051/#77765/#84132): an + # auth-classified permanent failure on a previously + # PROVEN session is often a transient/ambiguous state + # (OAuth flow lock left corrupt by a raced teardown), + # not truly revoked credentials. Grant ONE + # suspect+reconnect cycle before the park ladder: mark + # the connection suspect so the next call health-checks + # it, and rebuild the transport instead of parking. + if ( + _is_auth_error(root) + and self._session_proven + and not self._permanent_grace_used + ): + self._permanent_grace_used = True + self.mark_suspect( + f"auth error on proven session: {root}" + ) + logger.warning( + "MCP server '%s': auth error on a previously " + "healthy session — marking suspect and forcing " + "one reconnect instead of parking (state: " + "connected → suspect): %s: %s", + self.name, type(root).__name__, root, + ) + self._reconnect_retries = 0 + backoff = 1.0 + await asyncio.sleep(_jittered(1.0)) + if self._shutdown_event.is_set(): + return + continue # A previously-working server now fails deterministically # (revoked credentials, URL now serving a web page, stdio # binary uninstalled). Retrying can't help — park @@ -4148,6 +4348,9 @@ class MCPServerTask: return finally: self.session = None + # Children of this transport are gone (or about to be); + # stale PIDs must never fast-fail the NEXT transport's calls. + self._stdio_child_pids = set() async def start(self, config: dict): """Create the background Task and wait until ready (or failed).""" @@ -5768,6 +5971,64 @@ def _mark_server_call_started(server: Any) -> None: mark_tool_call() +@asynccontextmanager +async def _track_inflight_rpc(server: Any, server_name: str, op: str): + """Register the running RPC on the server so teardown can fail it fast. + + Every user-visible request family wraps its RPC in this context + (#48069 salvage). If a deliberate reconnect/shutdown teardown cancels + the task (``_fail_inflight_calls`` sets ``_reconnecting`` first), the + cancel is converted into a clean retryable RuntimeError instead of a raw + CancelledError; external cancels (caller timeout, user interrupt) + propagate unchanged. + """ + inflight = getattr(server, "_inflight_tasks", None) + task = asyncio.current_task() + if task is not None and inflight is not None: + # Test doubles may pass a bare SimpleNamespace; tracking is then + # simply skipped (fast-fail teardown is a production-connection + # feature, not something a fake needs). + inflight.add(task) + try: + yield + except asyncio.CancelledError: + if getattr(server, "_reconnecting", False): + raise RuntimeError( + f"MCP {op} on '{server_name}' was aborted by a reconnect " + f"teardown; retry the request on the rebuilt session" + ) from None + raise + finally: + if task is not None and inflight is not None: + inflight.discard(task) + + +def _ensure_healthy_or_recycle(server: Any, server_name: str) -> None: + """Health-check a suspect connection before its next call (#85125 3b). + + Implements the SuspectableBackend cheap-mark/lazy-verify contract at the + dispatch boundary: a connection latched as suspect by a race or an auth + error is probed once; a failed probe recycles it so the call below hits + the normal reconnect path. A HEALTHY connection is never recycled here. + """ + if not getattr(server, "_suspect_reason", None): + return + with _lock: + loop = _mcp_loop + if loop is None or not loop.is_running(): + return # no background loop — nothing to verify against + try: + healthy = bool(_run_on_mcp_loop(server.ensure_healthy, timeout=15.0)) + except Exception as exc: # never let the probe break dispatch + logger.debug( + "MCP server '%s': suspect health check errored: %s", + server_name, exc, + ) + healthy = False + if not healthy: + _signal_reconnect(server) + + def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): """Return a sync handler that calls an MCP tool via the background loop. @@ -5845,14 +6106,77 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): async def _call(): _mark_server_call_started(server) - async with server._rpc_lock: + async with server._rpc_lock, _track_inflight_rpc( + server, server_name, f"tools/call {tool_name}" + ): # Snapshot the agent's context so an elicitation callback # triggered during this call (fired on the MCP recv loop # task, which doesn't inherit our contextvars) can replay # it and detect the gateway platform / session for routing. server._pending_call_context = contextvars.copy_context() try: - result = await server.session.call_tool(tool_name, arguments=args) + # Fast-fail (#81995): a stdio subprocess that is already + # dead must not own this call slot — fail immediately + # instead of waiting out the full tool timeout on a + # transport nobody will ever answer. + _stdio_dead = getattr(server, "_stdio_children_dead", None) + # callable() + real-bool result: MagicMock attributes return + # truthy Mocks, which would spuriously trip the fast-fail. + if ( + callable(_stdio_dead) + and isinstance(_stdio_dead_result := _stdio_dead(), bool) + and _stdio_dead_result + ): + raise TimeoutError( + f"MCP stdio subprocess for '{server_name}' has " + f"exited; failing the call fast instead of " + f"waiting {float(tool_timeout):.0f}s" + ) + _call_coro = server.session.call_tool(tool_name, arguments=args) + _watch_children = getattr(server, "_watch_stdio_children", None) + _watch_ok = ( + _watch_children is not None + and inspect.isawaitable(_watch_children()) + and asyncio.iscoroutine(_call_coro) + ) + if not _watch_ok: + # Stubbed sessions (MagicMock in tests) return a + # non-awaitable, or there is no child-watcher to race + # against: plain await is exactly the pre-#81995 + # semantics. + result = ( + await _call_coro + if asyncio.iscoroutine(_call_coro) + else _call_coro + ) + else: + # Fast-fail machinery (#81995): the RPC races a + # stdio-children watcher so a dead subprocess fails + # the call immediately instead of riding out the full + # tool timeout. + rpc_task = asyncio.ensure_future(_call_coro) + watch_task = asyncio.ensure_future(_watch_children()) + try: + done, _pending = await asyncio.wait( + {rpc_task, watch_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + if watch_task in done and not rpc_task.done(): + rpc_task.cancel() + raise TimeoutError( + f"MCP stdio subprocess for '{server_name}' " + f"exited mid-call; failing the call fast " + f"instead of waiting " + f"{float(tool_timeout):.0f}s" + ) + result = await rpc_task + finally: + watch_task.cancel() + if not rpc_task.done(): + rpc_task.cancel() + await asyncio.gather( + rpc_task, watch_task, return_exceptions=True + ) finally: server._pending_call_context = None # The RPC round-trip completed — the session is demonstrably