From d0829c5c02ec7c26814d820b00c2b63cb013586b Mon Sep 17 00:00:00 2001 From: yoni Date: Fri, 28 Aug 2026 03:51:47 +0000 Subject: [PATCH] Polish MCP resilience follow-up --- strix/tools/mcp/session.py | 124 ++++++++++++----------------------- tests/test_mcp_client.py | 56 ++++++++++------ tests/test_mcp_resilience.py | 26 ++++++-- 3 files changed, 99 insertions(+), 107 deletions(-) diff --git a/strix/tools/mcp/session.py b/strix/tools/mcp/session.py index d0f4a8d7..c997fe67 100644 --- a/strix/tools/mcp/session.py +++ b/strix/tools/mcp/session.py @@ -39,6 +39,8 @@ serialized into the run's event stream, or written to disk; :meth:`__repr__` omits it and the token field's own ``repr`` is already suppressed. """ +# ruff: noqa: BLE001 + from __future__ import annotations import asyncio @@ -411,20 +413,7 @@ class SupervisedMcpSession: await self._safe_cleanup() self._fail_pending() return - except BaseExceptionGroup as exc: - failure = classify(exc) - logger.warning( - "Skipping MCP connection %r kind=%s status=%s attempt=1 delay=0", - self._name, - failure.kind, - failure.status, - exc_info=True, - ) - self._report_ready(value=False) - await self._safe_cleanup() - self._fail_pending() - return - except Exception as exc: # noqa: BLE001 - classify ordinary connect failures + except (BaseExceptionGroup, Exception) as exc: failure = classify(exc) logger.warning( "Skipping MCP connection %r kind=%s status=%s attempt=1 delay=0", @@ -479,7 +468,7 @@ class SupervisedMcpSession: # -- run one job with bounded classified retries -------------------------- - async def _execute(self, job: Job) -> _Outcome: # noqa: PLR0911, PLR0912, PLR0915 + async def _execute(self, job: Job) -> _Outcome: # noqa: PLR0911, PLR0912 """Run one job with classified retries and temporary quarantine.""" if self._dead: return _Outcome(dead=True) @@ -489,40 +478,32 @@ class SupervisedMcpSession: return _Outcome(dead=True) self._unavailable_until = None logger.info( - "MCP connection %r revive started kind=%s status=%s attempt=1 delay=%.2f", + "MCP connection %r revive started kind=%s status=%s attempt=1", self._name, self._last_failure.kind, self._last_failure.status, - 0.0, ) + if self._call_semaphore is None: + self._call_semaphore = _call_semaphore( + self._name, + ( + self._config.max_concurrent_calls + if self._config is not None + else DEFAULT_MAX_CONCURRENT_CALLS + ), + ) failure: FailureInfo | None = None for attempt in range(1, _MAX_ATTEMPTS + 1): - if self._call_semaphore is None: - self._call_semaphore = _call_semaphore( - self._name, - ( - self._config.max_concurrent_calls - if self._config is not None - else DEFAULT_MAX_CONCURRENT_CALLS - ), - ) if self._server is None: reconnected, reconnect_failure = await self._reconnect() if not reconnected: failure = reconnect_failure or FailureInfo( "transport", reason="reconnect failed" ) - self._last_failure = failure - if failure.kind == "auth": - self._mark_dead(failure, attempt=attempt) - return _Outcome(dead=True) - if attempt == _MAX_ATTEMPTS: - await self._quarantine(failure, attempt=attempt) - return _Outcome(dead=True) - delay = _retry_delay(attempt, failure.retry_after) - self._log_retry(failure, attempt, delay) - await asyncio.sleep(delay) + outcome = await self._handle_failure(failure, attempt) + if outcome is not None: + return outcome continue assert self._server is not None call_semaphore = self._call_semaphore @@ -536,40 +517,35 @@ class SupervisedMcpSession: failure = ( self._recorder.take() if self._recorder is not None else None ) or FailureInfo("transport", reason="session cancelled") - except BaseExceptionGroup as exc: - failure = classify(exc) - if failure.kind == "unknown" and self._recorder is not None: - failure = self._recorder.take() or failure - except Exception as exc: # noqa: BLE001 - classify ordinary call failures + except (BaseExceptionGroup, Exception) as exc: failure = classify(exc) if failure.kind == "unknown" and self._recorder is not None: failure = self._recorder.take() or failure - self._last_failure = failure - if failure.kind == "auth": - self._mark_dead(failure, attempt=attempt) - return _Outcome(dead=True) - if attempt == _MAX_ATTEMPTS: - await self._quarantine(failure, attempt=attempt) - return _Outcome(dead=True) - delay = _retry_delay(attempt, failure.retry_after) - self._log_retry(failure, attempt, delay) - await asyncio.sleep(delay) + outcome = await self._handle_failure(failure, attempt) + if outcome is not None: + return outcome reconnected, reconnect_failure = await self._reconnect() if not reconnected: failure = reconnect_failure or FailureInfo("transport", reason="reconnect failed") - self._last_failure = failure - if failure.kind == "auth": - self._mark_dead(failure, attempt=attempt) - return _Outcome(dead=True) - if attempt == _MAX_ATTEMPTS: - await self._quarantine(failure, attempt=attempt) - return _Outcome(dead=True) - delay = _retry_delay(attempt, failure.retry_after) - self._log_retry(failure, attempt, delay) - await asyncio.sleep(delay) + outcome = await self._handle_failure(failure, attempt) + if outcome is not None: + return outcome return _Outcome(dead=True) + async def _handle_failure(self, failure: FailureInfo, attempt: int) -> _Outcome | None: + self._last_failure = failure + if failure.kind == "auth": + self._mark_dead(failure, attempt=attempt) + return _Outcome(dead=True) + if attempt == _MAX_ATTEMPTS: + await self._quarantine(failure, attempt=attempt) + return _Outcome(dead=True) + delay = _retry_delay(attempt, failure.retry_after) + self._log_retry(failure, attempt, delay) + await asyncio.sleep(delay) + return None + def _log_retry(self, failure: FailureInfo, attempt: int, delay: float) -> None: logger.warning( "MCP connection %r retryable failure kind=%s status=%s attempt=%d delay=%.2f", @@ -610,13 +586,7 @@ class SupervisedMcpSession: raise self._server = None return False, FailureInfo("transport", reason="reconnect cancelled") - except BaseExceptionGroup as exc: - self._server = None - failure = classify(exc) - if failure.kind == "unknown" and self._recorder is not None: - failure = self._recorder.take() or failure - return False, failure - except Exception as exc: # noqa: BLE001 - classify ordinary reconnect failures + except (BaseExceptionGroup, Exception) as exc: self._server = None failure = classify(exc) if failure.kind == "unknown" and self._recorder is not None: @@ -642,30 +612,20 @@ class SupervisedMcpSession: same task before the error propagates, so a failed connect never orphans an MCP subprocess or half-open HTTP session. """ - from strix.tools.mcp.client import BuiltMcpServer, _build_server + from strix.tools.mcp.client import _build_server if self._config is None: raise RuntimeError(f"MCP connection {self._name!r} has no config to connect") built = _build_server(self._config) - if isinstance(built, BuiltMcpServer): - server = built.server - self._recorder = built.recorder - else: - # Test and private consumers may still provide a connected server - # directly when replacing the private builder. - server = cast("MCPServer", built) # type: ignore[unreachable] - self._recorder = None + server = built.server + self._recorder = built.recorder try: await server.connect() # type: ignore[no-untyped-call] except asyncio.CancelledError: with contextlib.suppress(Exception): await server.cleanup() # type: ignore[no-untyped-call] raise - except BaseExceptionGroup: - with contextlib.suppress(Exception): - await server.cleanup() # type: ignore[no-untyped-call] - raise - except Exception: + except (BaseExceptionGroup, Exception): with contextlib.suppress(Exception): await server.cleanup() # type: ignore[no-untyped-call] raise diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 0ce12780..82161ad5 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -167,6 +167,10 @@ def _config(name: str, allowed_tools: list[str] | None) -> McpConnectionConfig: ) +def _built_server(server: MCPServer) -> mcp_client.BuiltMcpServer: + return mcp_client.BuiltMcpServer(server, None) + + def _ctx(registry: McpRegistry | None) -> ToolContext[dict[str, Any]]: context: dict[str, Any] = {} if registry is None else {MCP_REGISTRY_CONTEXT_KEY: registry} return ToolContext( @@ -300,7 +304,9 @@ async def test_connect_returns_sessions_without_registering_agent_tools( "fs": FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]), "db": FakeMCPServer("db", [_mcp_tool("query")]), } - monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) + monkeypatch.setattr( + mcp_client, "_build_server", lambda config: _built_server(servers[config.name]) + ) connections = await mcp_client.connect_mcp_servers( [_config("fs", None), _config("db", ["query"])] @@ -317,7 +323,7 @@ async def test_connect_returns_sessions_without_registering_agent_tools( @pytest.mark.asyncio async def test_tool_count_honors_the_allowlist(monkeypatch: pytest.MonkeyPatch) -> None: server = FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) connections = await mcp_client.connect_mcp_servers([_config("fs", ["read_file"])]) @@ -331,7 +337,7 @@ async def test_connection_notes_ride_on_the_connection( monkeypatch: pytest.MonkeyPatch, ) -> None: server = FakeMCPServer("db", [_mcp_tool("query")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) config = McpConnectionConfig( name="db", url="https://mcp.example.com", @@ -849,7 +855,9 @@ async def test_connect_skips_a_connection_whose_connect_is_cancelled( cleaned.append(self._name) servers = {"good": _Tracking("good"), "bad": _Tracking("bad", cancel_connect=True)} - monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) + monkeypatch.setattr( + mcp_client, "_build_server", lambda config: _built_server(servers[config.name]) + ) configs = [_config("good", ["t"]), _config("bad", ["t"])] @@ -885,7 +893,9 @@ async def test_connect_cleans_up_started_sessions_when_attach_is_cancelled( cleaned.append(self._name) servers = {"good": _Tracking("good"), "slow": _Tracking("slow", block_connect=True)} - monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) + monkeypatch.setattr( + mcp_client, "_build_server", lambda config: _built_server(servers[config.name]) + ) async def _attach() -> list[Any]: # Connect "good" first, then hang forever connecting "slow". @@ -964,7 +974,7 @@ async def test_attach_populates_registry_with_provider_and_transform( monkeypatch: pytest.MonkeyPatch, ) -> None: server = FakeMCPServer("db", [_mcp_tool("query")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) def transform(_label: str, structured: Any) -> Any: return {"kept": structured} @@ -997,7 +1007,7 @@ async def test_attach_bare_request_matches_the_command_line_shape( # The command-line path wraps each config in a bare request (no provider or # transform); purpose then falls back to the connection's notes. server = FakeMCPServer("db", [_mcp_tool("query")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) config = McpConnectionConfig( name="db", url="https://mcp.example.com", @@ -1028,7 +1038,9 @@ async def test_attach_is_fail_open_and_skips_a_failed_connection( raise RuntimeError("cannot reach server") servers = {"good": good, "bad": _Failing("bad", [_mcp_tool("t")])} - monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) + monkeypatch.setattr( + mcp_client, "_build_server", lambda config: _built_server(servers[config.name]) + ) registry = McpRegistry() connections = await attach_mcp_requests( @@ -1236,7 +1248,7 @@ async def test_call_mcp_reconnects_and_retries_after_a_session_death( first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403")) second = FakeMCPServer("fs", [_mcp_tool("read_file")]) built = iter([first, second]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built)) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built))) session = await _started_session(_secret_config("fs")) registry = McpRegistry() @@ -1264,10 +1276,10 @@ async def test_call_mcp_marks_connection_dead_when_reconnect_keeps_failing( first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403")) built = {"n": 0} - def _build(_config: McpConnectionConfig) -> MCPServer: + def _build(_config: McpConnectionConfig) -> mcp_client.BuiltMcpServer: built["n"] += 1 if built["n"] == 1: - return first + return _built_server(first) raise ConnectionError("cannot reconnect") monkeypatch.setattr(mcp_client, "_build_server", _build) @@ -1328,7 +1340,7 @@ async def test_idle_session_death_self_heals_on_reconnect( first = FakeMCPServer("fs", [_mcp_tool("read_file")]) second = FakeMCPServer("fs", [_mcp_tool("read_file")]) built = iter([first, second]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built)) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built))) session = await _started_session(_secret_config("fs")) registry = McpRegistry() @@ -1358,7 +1370,7 @@ async def test_flapping_idle_session_is_marked_dead_without_looping( first = FakeMCPServer("fs", [_mcp_tool("read_file")]) second = FakeMCPServer("fs", [_mcp_tool("read_file")]) built = iter([first, second]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built)) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built))) session = await _started_session(_secret_config("fs")) registry = McpRegistry() @@ -1427,7 +1439,7 @@ async def test_aclose_is_bounded_when_an_in_flight_call_hangs( # aclose falls back to cancelling the supervising task, and cleanup still runs. monkeypatch.setattr(mcp_session_mod, "_SHUTDOWN_TIMEOUT", 0.2) server = _HangingCallServer("fs", [_mcp_tool("read_file")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) session = await _started_session(_secret_config("fs")) call = asyncio.create_task(session.dispatch("read_file", {}, label="fs_read_file")) @@ -1451,7 +1463,7 @@ async def test_aclose_cleans_up_when_connect_is_cancelled_mid_await( # is cancelled; aclose must not raise on it and must still cancel + clean up the # partially connected supervisor. server = _HangingConnectServer("fs", [_mcp_tool("read_file")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) session = SupervisedMcpSession(_secret_config("fs")) start = asyncio.create_task(session.start()) @@ -1476,14 +1488,14 @@ async def test_a_session_death_is_contained_and_other_connections_survive( healthy = FakeMCPServer("healthy", [_mcp_tool("read_file")]) dying_builds = {"n": 0} - def _build(config: McpConnectionConfig) -> MCPServer: + def _build(config: McpConnectionConfig) -> mcp_client.BuiltMcpServer: if config.name == "healthy": - return healthy + return _built_server(healthy) # The dying connection connects once, then its rebuild raises, so it ends # up marked dead rather than recovering. dying_builds["n"] += 1 if dying_builds["n"] == 1: - return dying + return _built_server(dying) raise ConnectionError("cannot reconnect") monkeypatch.setattr(mcp_client, "_build_server", _build) @@ -1520,11 +1532,13 @@ async def test_reconnect_reuses_the_stored_config_and_never_logs_the_token( # inventory list_mcps emits. seen_tokens: list[str | None] = [] - def _build(config: McpConnectionConfig) -> MCPServer: + def _build(config: McpConnectionConfig) -> mcp_client.BuiltMcpServer: seen_tokens.append(config.auth.token if config.auth else None) if len(seen_tokens) == 1: - return _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403")) - return FakeMCPServer("fs", [_mcp_tool("read_file")]) + return _built_server( + _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403")) + ) + return _built_server(FakeMCPServer("fs", [_mcp_tool("read_file")])) monkeypatch.setattr(mcp_client, "_build_server", _build) diff --git a/tests/test_mcp_resilience.py b/tests/test_mcp_resilience.py index 42e29dfb..e1b35104 100644 --- a/tests/test_mcp_resilience.py +++ b/tests/test_mcp_resilience.py @@ -24,6 +24,10 @@ FakeMCPServer: Any = _test_mcp_client.FakeMCPServer _mcp_tool: Any = _test_mcp_client._mcp_tool +def _built_server(server: Any) -> Any: + return mcp_client.BuiltMcpServer(server, None) + + def _http_error(status: int, *, retry_after: str | None = None) -> httpx.HTTPStatusError: request = httpx.Request( "POST", @@ -67,6 +71,20 @@ def test_classifies_nested_exception_groups_by_specificity() -> None: assert info.retryable is False +@pytest.mark.parametrize("control_flow", [SystemExit, KeyboardInterrupt]) +@pytest.mark.asyncio +async def test_control_flow_exceptions_propagate( + control_flow: type[BaseException], +) -> None: + server = _sequence_server("control-flow", control_flow("stop")) + session = mcp_session.SupervisedMcpSession.adopt(server, name="control-flow") + + with pytest.raises(control_flow): + await session.dispatch("read", {}, label="control_flow") + + await session.aclose() + + @pytest.mark.asyncio async def test_retry_after_parses_seconds_and_http_date() -> None: seconds = HttpStatusRecorder() @@ -144,7 +162,7 @@ async def test_rate_limit_retries_and_succeeds( _sequence_server("rate"), ] ) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(builds)) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(builds))) session = mcp_session.SupervisedMcpSession(_config("rate")) assert await session.start() result = await session.dispatch("read", {}, label="rate_read") @@ -169,7 +187,7 @@ async def test_server_exhaustion_quarantines_then_revives( _sequence_server("quarantine"), ] ) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(builds)) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(builds))) session = mcp_session.SupervisedMcpSession(_config("quarantine")) assert await session.start() result = await session.dispatch("read", {}, label="quarantine_read") @@ -186,7 +204,7 @@ async def test_server_exhaustion_quarantines_then_revives( @pytest.mark.asyncio async def test_auth_failure_dies_without_retry(monkeypatch: pytest.MonkeyPatch) -> None: builds = [_sequence_server("auth", _http_error(401))] - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: builds.pop()) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(builds.pop())) session = mcp_session.SupervisedMcpSession(_config("auth")) assert await session.start() result = await session.dispatch("read", {}, label="auth_read") @@ -333,7 +351,7 @@ async def test_resilience_logs_do_not_include_request_secrets( monkeypatch.setattr(mcp_session, "_retry_delay", _zero_delay) monkeypatch.setattr(asyncio, "sleep", _no_sleep) server = _sequence_server("redaction", _http_error(401)) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) session = mcp_session.SupervisedMcpSession(_config("redaction")) assert await session.start() with caplog.at_level("WARNING"):