From a09512c3ad5df93bbfd52297a5a9901a76bacdf5 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:14:37 -0700 Subject: [PATCH] refactor(mcp): move loop start/stop + reconnect signalling into mcp_tool_loop.py (writes via origin module) --- tools/mcp_tool.py | 179 ++-------------------------------------- tools/mcp_tool_loop.py | 180 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 186 insertions(+), 173 deletions(-) diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 2adce4a8d4..e675abcb7d 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -185,6 +185,12 @@ from tools.mcp_tool_loop import ( # noqa: F401 -- re-exported for callers and t _wrap_with_home_override, _wrap_with_dashboard_oauth_flow, _run_on_mcp_loop, + _signal_reconnect, + reconnect_mcp_server, + _wait_for_server_session_ready, + _signal_reconnect_and_wait, + _ensure_mcp_loop, + _stop_mcp_loop, ) from tools.mcp_tool_discovery import ( # noqa: F401 -- re-exported for callers and test patches _record_connect_failure, @@ -706,106 +712,6 @@ def _reset_server_error(server_name: str) -> None: _server_breaker_opened_at.pop(server_name, None) -def _signal_reconnect(server: Any) -> bool: - """Ask a server task to rebuild its transport, thread-safely. - - Handlers run on caller threads while the event lives on the MCP loop, so - it is set via ``call_soon_threadsafe`` when the loop runs (direct - ``.set()`` otherwise). False when the server has no reconnect machinery. - """ - event = getattr(server, "_reconnect_event", None) - if event is None: - return False - loop = _mcp_loop - if ( - isinstance(event, asyncio.Event) - and loop is not None - and loop.is_running() - ): - loop.call_soon_threadsafe(event.set) - else: - event.set() - return True - - -def reconnect_mcp_server(server_name: str) -> bool: - """Ask a currently-live MCP server to rebuild after external re-auth.""" - with _lock: - server = _servers.get(server_name) - if server is None: - return False - return _signal_reconnect(server) - - -def _wait_for_server_session_ready( - srv: "MCPServerTask", - *, - old_session: Any = None, - timeout: float = 15.0, -) -> bool: - """Poll until the server exposes a usable, ready session. - - During a reconnect ``srv.session`` is briefly None or still the stale - object; retrying blindly there burns breaker strikes. With - ``old_session`` the observed session must differ from it. Iteration- - bounded, not deadline-bounded: tests freeze ``time.monotonic``. - """ - poll_interval = 0.25 - iterations = max(1, int(max(float(timeout), 0.0) / poll_interval)) - for i in range(iterations): - session = getattr(srv, "session", None) - ready = getattr(srv, "_ready", None) - is_ready = True - if ready is not None and hasattr(ready, "is_set"): - try: - is_ready = bool(ready.is_set()) - except Exception: - is_ready = True - if session is not None and session is not old_session and is_ready: - return True - if i < iterations - 1: - time.sleep(poll_interval) - return False - - -def _signal_reconnect_and_wait( - server_name: str, - srv: "MCPServerTask", - *, - op_description: str, - timeout: float = 15.0, -) -> bool: - """Request a transport rebuild and wait for the fresh session. - - ``_ready`` is cleared on the loop BEFORE ``_reconnect_event`` is set; - otherwise the readiness poll returns immediately and retries against the - same dead session. - """ - loop = _mcp_loop - if loop is None or not loop.is_running(): - return False - - old_session = getattr(srv, "session", None) - - def _request_reconnect() -> None: - ready = getattr(srv, "_ready", None) - if ready is not None and hasattr(ready, "clear"): - ready.clear() - reconnect_event = getattr(srv, "_reconnect_event", None) - if reconnect_event is not None and hasattr(reconnect_event, "set"): - reconnect_event.set() - - logger.info( - "MCP server '%s': %s requesting transport reconnect", - server_name, op_description, - ) - loop.call_soon_threadsafe(_request_reconnect) - return _wait_for_server_session_ready( - srv, - old_session=old_session, - timeout=timeout, - ) - # Raw server names opted into parallel tool calls. Raw identity matters: # ``foo-bar`` and ``foo_bar`` both sanitize to ``foo_bar`` but must not share # policy. @@ -858,22 +764,6 @@ _MCP_DISCOVERY_LOCK_MAX_RETRIES: int = 240 _MCP_DISCOVERY_LOCK_RETRY_DELAY_S: float = 0.5 -def _ensure_mcp_loop(): - """Start the background event loop thread if not already running.""" - global _mcp_loop, _mcp_thread - with _lock: - if _mcp_loop is not None and _mcp_loop.is_running(): - return - _mcp_loop = asyncio.new_event_loop() - _mcp_loop.set_exception_handler(_mcp_loop_exception_handler) - _mcp_thread = threading.Thread( - target=_mcp_loop.run_forever, - name="mcp-event-loop", - daemon=True, - ) - _mcp_thread.start() - - # --------------------------------------------------------------------------- # Connecting, lazy start, discovery # --------------------------------------------------------------------------- @@ -884,60 +774,3 @@ def _ensure_mcp_loop(): # --------------------------------------------------------------------------- -def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: - """Stop the background event loop and join its thread.""" - global _mcp_loop, _mcp_thread - with _lock: - if only_if_idle and (_servers or _server_connecting): - logger.debug("Leaving MCP event loop running; active servers are registered or connecting") - return False - loop = _mcp_loop - thread = _mcp_thread - _mcp_loop = None - _mcp_thread = None - if loop is not None: - # Drain before stopping: tasks still suspended when the loop closes - # get resumed by the GC against a closed loop. shutdown_mcp_servers - # only reaps servers held in _servers; everything else ends up here. - stop_owned_by_loop = False - if loop.is_running(): - from agent.async_utils import safe_schedule_threadsafe - - future = safe_schedule_threadsafe( - _drain_and_stop_mcp_loop(), loop, - logger=logger, - log_message="MCP loop drain: failed to schedule", - log_level=logging.WARNING, - ) - if future is not None: - stop_owned_by_loop = True - try: - future.result(timeout=_MCP_LOOP_DRAIN_TIMEOUT + 1) - except TimeoutError: - logger.warning( - "Timed out waiting for MCP loop drain after %.1fs", - _MCP_LOOP_DRAIN_TIMEOUT + 1, - ) - except BaseException as exc: - logger.warning("Error draining MCP loop tasks: %s", exc) - elif not loop.is_closed(): - try: - loop.run_until_complete( - _drain_mcp_loop_tasks(timeout=_MCP_LOOP_DRAIN_TIMEOUT) - ) - except BaseException as exc: - logger.warning("Error draining stopped MCP loop tasks: %s", exc) - - if not stop_owned_by_loop and loop.is_running(): - loop.call_soon_threadsafe(loop.stop) - if thread is not None: - thread.join(timeout=5) - if thread.is_alive(): - logger.warning("MCP event loop thread did not stop within 5.0s") - try: - loop.close() - except Exception as exc: - logger.warning("Unable to close MCP event loop cleanly: %s", exc) - # The loop is gone, so no session can be in flight: reap active too. - _kill_orphaned_mcp_children(include_active=True) - return True diff --git a/tools/mcp_tool_loop.py b/tools/mcp_tool_loop.py index 8876d6bb9c..eb7b68113e 100644 --- a/tools/mcp_tool_loop.py +++ b/tools/mcp_tool_loop.py @@ -11,6 +11,7 @@ import concurrent.futures import errno import logging import os +import threading import time from typing import Any, Coroutine from tools.mcp_tool_common import _core @@ -227,3 +228,182 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): if future.done(): return future.result() continue + +def _signal_reconnect(server: Any) -> bool: + """Ask a server task to rebuild its transport, thread-safely. + + Handlers run on caller threads while the event lives on the MCP loop, so + it is set via ``call_soon_threadsafe`` when the loop runs (direct + ``.set()`` otherwise). False when the server has no reconnect machinery. + """ + event = getattr(server, "_reconnect_event", None) + if event is None: + return False + loop = _core._mcp_loop + if ( + isinstance(event, asyncio.Event) + and loop is not None + and loop.is_running() + ): + loop.call_soon_threadsafe(event.set) + else: + event.set() + return True + + +def reconnect_mcp_server(server_name: str) -> bool: + """Ask a currently-live MCP server to rebuild after external re-auth.""" + with _core._lock: + server = _core._servers.get(server_name) + if server is None: + return False + return _core._signal_reconnect(server) + + +def _wait_for_server_session_ready( + srv: Any, + *, + old_session: Any = None, + timeout: float = 15.0, +) -> bool: + """Poll until the server exposes a usable, ready session. + + During a reconnect ``srv.session`` is briefly None or still the stale + object; retrying blindly there burns breaker strikes. With + ``old_session`` the observed session must differ from it. Iteration- + bounded, not deadline-bounded: tests freeze ``time.monotonic``. + """ + poll_interval = 0.25 + iterations = max(1, int(max(float(timeout), 0.0) / poll_interval)) + for i in range(iterations): + session = getattr(srv, "session", None) + ready = getattr(srv, "_ready", None) + is_ready = True + if ready is not None and hasattr(ready, "is_set"): + try: + is_ready = bool(ready.is_set()) + except Exception: + is_ready = True + if session is not None and session is not old_session and is_ready: + return True + if i < iterations - 1: + time.sleep(poll_interval) + return False + + +def _signal_reconnect_and_wait( + server_name: str, + srv: Any, + *, + op_description: str, + timeout: float = 15.0, +) -> bool: + """Request a transport rebuild and wait for the fresh session. + + ``_ready`` is cleared on the loop BEFORE ``_reconnect_event`` is set; + otherwise the readiness poll returns immediately and retries against the + same dead session. + """ + loop = _core._mcp_loop + if loop is None or not loop.is_running(): + return False + + old_session = getattr(srv, "session", None) + + def _request_reconnect() -> None: + ready = getattr(srv, "_ready", None) + if ready is not None and hasattr(ready, "clear"): + ready.clear() + reconnect_event = getattr(srv, "_reconnect_event", None) + if reconnect_event is not None and hasattr(reconnect_event, "set"): + reconnect_event.set() + + logger.info( + "MCP server '%s': %s requesting transport reconnect", + server_name, op_description, + ) + loop.call_soon_threadsafe(_request_reconnect) + return _core._wait_for_server_session_ready( + srv, + old_session=old_session, + timeout=timeout, + ) + + +def _ensure_mcp_loop(): + """Start the background event loop thread if not already running. + + The loop/thread handles live on the ORIGIN module (tests read and reset + ``tools.mcp_tool._mcp_loop``), so they are written there, never here. + """ + from tools import mcp_tool as _origin + with _core._lock: + if _origin._mcp_loop is not None and _origin._mcp_loop.is_running(): + return + _origin._mcp_loop = asyncio.new_event_loop() + _origin._mcp_loop.set_exception_handler(_core._mcp_loop_exception_handler) + _origin._mcp_thread = threading.Thread( + target=_origin._mcp_loop.run_forever, + name="mcp-event-loop", + daemon=True, + ) + _origin._mcp_thread.start() + + +def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: + """Stop the background event loop and join its thread.""" + from tools import mcp_tool as _origin + with _core._lock: + if only_if_idle and (_core._servers or _core._server_connecting): + logger.debug("Leaving MCP event loop running; active servers are registered or connecting") + return False + loop = _origin._mcp_loop + thread = _origin._mcp_thread + _origin._mcp_loop = None + _origin._mcp_thread = None + if loop is not None: + # Drain before stopping: tasks still suspended when the loop closes + # get resumed by the GC against a closed loop. shutdown_mcp_servers + # only reaps servers held in _servers; everything else ends up here. + stop_owned_by_loop = False + if loop.is_running(): + from agent.async_utils import safe_schedule_threadsafe + + future = safe_schedule_threadsafe( + _core._drain_and_stop_mcp_loop(), loop, + logger=logger, + log_message="MCP loop drain: failed to schedule", + log_level=logging.WARNING, + ) + if future is not None: + stop_owned_by_loop = True + try: + future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1) + except TimeoutError: + logger.warning( + "Timed out waiting for MCP loop drain after %.1fs", + _core._MCP_LOOP_DRAIN_TIMEOUT + 1, + ) + except BaseException as exc: + logger.warning("Error draining MCP loop tasks: %s", exc) + elif not loop.is_closed(): + try: + loop.run_until_complete( + _core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT) + ) + except BaseException as exc: + logger.warning("Error draining stopped MCP loop tasks: %s", exc) + + if not stop_owned_by_loop and loop.is_running(): + loop.call_soon_threadsafe(loop.stop) + if thread is not None: + thread.join(timeout=5) + if thread.is_alive(): + logger.warning("MCP event loop thread did not stop within 5.0s") + try: + loop.close() + except Exception as exc: + logger.warning("Unable to close MCP event loop cleanly: %s", exc) + # The loop is gone, so no session can be in flight: reap active too. + _core._kill_orphaned_mcp_children(include_active=True) + return True