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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user