Merge pull request #94184 from kshitijk4poor/fix/85125-3b-mcp-recovery

fix(mcp): recover poisoned connections + fail fast on dead stdio transports (#85125 3b)
This commit is contained in:
kshitij
2026-08-25 03:35:38 +05:30
committed by GitHub
2 changed files with 364 additions and 11 deletions

View File

@@ -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

View File

@@ -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