refactor(tools): MCP _dispatch absorbs _invoke_with_recovery; oauth 401 recovery flattened; small predicate folds

This commit is contained in:
Teknium
2026-09-03 01:37:03 -07:00
parent bccfd1de26
commit 35b3888fc5
4 changed files with 25 additions and 32 deletions
+5 -6
View File
@@ -342,22 +342,21 @@ class MCPOAuthManager:
async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None:
"""Single recovery attempt behind *pending*; always clears the dedup slot."""
can_refresh = False
try:
# Disk changed (external refresh)? Else: if the SDK can refresh in place, let the caller retry.
if await self.invalidate_if_disk_changed(server_name):
can_refresh = True
else:
can_refresh = await self.invalidate_if_disk_changed(server_name)
if not can_refresh:
try:
can_refresh = bool(entry.provider.context.can_refresh_token())
except Exception: # no context / not callable / probe failed
can_refresh = False
except Exception as exc: # pragma: no cover — defensive
logger.warning("MCP OAuth '%s': 401 handler failed: %s", server_name, exc)
can_refresh = False
finally:
if not pending.done():
pending.set_result(can_refresh)
entry.pending_401.pop(key, None)
if not pending.done():
pending.set_result(can_refresh)
async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool:
"""Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect and retry. False: no
+2 -2
View File
@@ -307,6 +307,6 @@ def _is_session_expired_error(exc: BaseException) -> bool:
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids false positives.
msg = str(current).lower()
found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS)
stack.extend(getattr(current, "exceptions", ()))
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))
stack.extend((*getattr(current, "exceptions", ()), getattr(current, "__cause__", None),
getattr(current, "__context__", None)))
return found
+17 -22
View File
@@ -36,8 +36,8 @@ _STDIO_DIED_AGAIN_MSG = (
def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]:
"""Approval gate for write-capable tools on ``trust: untrusted`` servers. None to proceed,
else a ``tool_error``. Fail-closed: approval-system errors block."""
trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL)
if trust != _core._TRUST_UNTRUSTED or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True:
if (_core._server_trust_levels.get(server_name, _core._TRUST_FULL) != _core._TRUST_UNTRUSTED
or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True):
return None
try: # lazy: tools.approval routes the prompt to whichever surface owns the session
from tools.approval import request_elicitation_consent
@@ -210,13 +210,18 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
f"{type(retry_exc).__name__}: {_exc_str(retry_exc)}"))
def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: str,
recoverers, on_final_failure: Callable[[BaseException], None],
record_outcome: bool = False) -> str:
"""Run ``call_once``, walking ``recoverers`` (``(server_name, exc, retry_call, op) ->
Optional[str]``, None = not its kind; order matters) on failure. Unrecovered exceptions go
through ``on_final_failure`` and become the generic call-failed error. ``record_outcome``
applies breaker bookkeeping to the FIRST attempt only; retries own theirs."""
def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers,
on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str:
"""Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine function *call*
on the MCP loop and, on failure, walk ``recoverers`` (``(server_name, exc, retry_call, op) -> Optional[str]``,
None = not its kind; order matters). Unrecovered exceptions go through ``on_final_failure`` and become the
generic call-failed error. ``record_outcome`` applies breaker bookkeeping to the FIRST attempt only."""
if callable(getattr(server, "mark_tool_call", None)):
server.mark_tool_call()
def call_once():
return _core._run_on_mcp_loop(call, timeout=tool_timeout)
try:
result = call_once()
return _record_call_outcome(server_name, result) if record_outcome else result
@@ -231,16 +236,6 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st
return tool_error(_sanitize_error(f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}"))
def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers,
on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str:
"""Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine
function *call* on the MCP loop and walk the recovery ladder (:func:`_invoke_with_recovery`)."""
if callable(getattr(server, "mark_tool_call", None)):
server.mark_tool_call()
return _invoke_with_recovery(server_name, lambda: _core._run_on_mcp_loop(call, timeout=tool_timeout), op,
recoverers, on_final_failure, record_outcome=record_outcome)
@asynccontextmanager
async def _track_inflight_rpc(server: Any, server_name: str, op: str):
"""Register the running RPC so teardown can fail it fast. A deliberate teardown
@@ -254,8 +249,8 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str):
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 RuntimeError(f"MCP {op} on '{server_name}' was aborted by a reconnect teardown; retry the "
f"request on the rebuilt session") from None
raise
finally:
if tracked:
@@ -272,7 +267,7 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str
raise _StdioChildExited(f"MCP stdio subprocess for '{server_name}' had already exited when the call was dispatched")
_call_coro = server.session.call_tool(tool_name, arguments=args)
_watch_children = getattr(server, "_watch_stdio_children", None)
if not (_watch_children is not None and inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)):
if not (inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)):
# Stubbed sessions return a non-awaitable, or there is no child-watcher to race: plain await.
return await _call_coro if asyncio.iscoroutine(_call_coro) else _call_coro
rpc_task = asyncio.ensure_future(_call_coro)
+1 -2
View File
@@ -51,8 +51,7 @@ class MCPServerHealthMixin:
return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None)
def _next_stdio_recycle_deadline(self) -> Optional[float]:
deadlines = self._stdio_recycle_deadlines()
return min(d for d, _ in deadlines) if deadlines else None
return min((d for d, _ in self._stdio_recycle_deadlines()), default=None)
def _mark_stdio_recycled(self, reason: str) -> None:
"""Mark a stdio session dormant before its transport finishes closing."""