Polish MCP resilience follow-up

This commit is contained in:
yoni 2026-08-28 03:51:47 +00:00
parent ad0dc4f5e9
commit d0829c5c02
3 changed files with 99 additions and 107 deletions

View file

@ -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

View file

@ -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)

View file

@ -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"):