From 35b3888fc538a96cb1900a3d66d7ba7b44aa7c4c Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:37:03 -0700 Subject: [PATCH] refactor(tools): MCP _dispatch absorbs _invoke_with_recovery; oauth 401 recovery flattened; small predicate folds --- tools/mcp_oauth_manager.py | 11 +++++------ tools/mcp_tool_errors.py | 4 ++-- tools/mcp_tool_handlers.py | 39 +++++++++++++++++--------------------- tools/mcp_tool_health.py | 3 +-- 4 files changed, 25 insertions(+), 32 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index c5cdda47ca..9f749de74e 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -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 diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 5209782883..eb0aa5f3c2 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -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 diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 22b02df029..f1c519bca3 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -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) diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index f5ed602d1e..4a8af5753a 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -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."""