From ad41517613e0d487a81e533f9fa2d4c9aa7fdd72 Mon Sep 17 00:00:00 2001 From: yoni Date: Fri, 28 Aug 2026 03:14:39 +0000 Subject: [PATCH] Harden MCP client transport resilience --- strix/tools/mcp/__init__.py | 4 + strix/tools/mcp/client.py | 24 ++- strix/tools/mcp/config.py | 12 ++ strix/tools/mcp/failures.py | 145 +++++++++++++++++ strix/tools/mcp/session.py | 299 ++++++++++++++++++++++++++--------- tests/test_mcp_resilience.py | 297 ++++++++++++++++++++++++++++++++++ 6 files changed, 703 insertions(+), 78 deletions(-) create mode 100644 strix/tools/mcp/failures.py create mode 100644 tests/test_mcp_resilience.py diff --git a/strix/tools/mcp/__init__.py b/strix/tools/mcp/__init__.py index 869e34aa..eb1de2f8 100644 --- a/strix/tools/mcp/__init__.py +++ b/strix/tools/mcp/__init__.py @@ -13,6 +13,7 @@ from strix.tools.mcp.config import ( McpAuth, McpConnectionConfig, ) +from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify from strix.tools.mcp.loader import load_user_mcp_configs from strix.tools.mcp.naming import namespaced_tool_name from strix.tools.mcp.registry import ( @@ -38,6 +39,8 @@ __all__ = [ "MCP_REGISTRY_CONTEXT_KEY", "BearerAuth", "ConnectedMcpServer", + "FailureInfo", + "HttpStatusRecorder", "McpAuth", "McpCallInfo", "McpConnectionConfig", @@ -50,6 +53,7 @@ __all__ = [ "SupervisedMcpSession", "attach_mcp_requests", "call_mcp", + "classify", "connect_mcp_servers", "describe_mcp", "list_mcps", diff --git a/strix/tools/mcp/client.py b/strix/tools/mcp/client.py index f36eadab..db5c3d3f 100644 --- a/strix/tools/mcp/client.py +++ b/strix/tools/mcp/client.py @@ -30,13 +30,17 @@ from agents.mcp import ( create_static_tool_filter, ) from mcp.client.stdio import stdio_client +from mcp.shared._httpx_utils import create_mcp_http_client +from strix.tools.mcp.failures import HttpStatusRecorder from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcpSession if TYPE_CHECKING: from collections.abc import Callable + import httpx + from strix.tools.mcp.config import McpConnectionConfig from strix.tools.mcp.registry import McpConnectionRequest, McpRegistry @@ -135,16 +139,34 @@ def _build_server(config: McpConnectionConfig) -> MCPServer: cache_tools_list=True, ) + recorder = HttpStatusRecorder() + + def httpx_client_factory( + headers: dict[str, str] | None = None, + timeout: httpx.Timeout | None = None, + auth: httpx.Auth | None = None, + ) -> httpx.AsyncClient: + client = create_mcp_http_client(headers=headers, timeout=timeout, auth=auth) + client.event_hooks.setdefault("response", []).append(recorder) + return client + http_params: MCPServerStreamableHttpParams = { "url": cast("str", config.url), "headers": _auth_headers(config), + "timeout": config.http_timeout_seconds, + "sse_read_timeout": config.sse_read_timeout_seconds, + "httpx_client_factory": httpx_client_factory, } - return MCPServerStreamableHttp( + server = MCPServerStreamableHttp( params=http_params, name=config.name, tool_filter=tool_filter, cache_tools_list=True, + client_session_timeout_seconds=config.session_timeout_seconds, ) + # The recorder is intentionally private and contains only status metadata. + server._strix_http_status_recorder = recorder # type: ignore[attr-defined] + return server def _mcp_result_to_tool_output(server: MCPServer, result: Any) -> Any: diff --git a/strix/tools/mcp/config.py b/strix/tools/mcp/config.py index b0cfffd1..e822b598 100644 --- a/strix/tools/mcp/config.py +++ b/strix/tools/mcp/config.py @@ -65,6 +65,18 @@ class McpConnectionConfig(BaseModel): MCP inventory every agent renders in its prompt, so it describes the connection once rather than being repeated onto each of its tools.""" + http_timeout_seconds: float = Field(default=30.0, gt=0) + """HTTP request timeout; the SDK's 5-second default is below tool p95s.""" + + sse_read_timeout_seconds: float = Field(default=300.0, gt=0) + """Stream read timeout; the SDK's 5-second default is below tool p95s.""" + + session_timeout_seconds: float = Field(default=60.0, gt=0) + """MCP operation timeout for SQL queries and cloud describe fan-outs.""" + + max_concurrent_calls: int = Field(default=4, ge=1) + """Maximum concurrent calls for this connection name across sessions.""" + @model_validator(mode="after") def _check_transport_fields(self) -> McpConnectionConfig: if self.transport == "http" and not self.url: diff --git a/strix/tools/mcp/failures.py b/strix/tools/mcp/failures.py new file mode 100644 index 00000000..a9a0c886 --- /dev/null +++ b/strix/tools/mcp/failures.py @@ -0,0 +1,145 @@ +"""Classify MCP connection failures without retaining sensitive request data.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass +from datetime import UTC, datetime +from email.utils import parsedate_to_datetime +from typing import Literal, cast + +import httpx +from agents.exceptions import UserError +from mcp.shared.exceptions import McpError + + +FailureKind = Literal["auth", "rate_limit", "server", "transport", "timeout", "protocol", "unknown"] + +_PRIORITY: dict[FailureKind, int] = { + "auth": 0, + "rate_limit": 1, + "server": 2, + "protocol": 3, + "timeout": 4, + "transport": 5, + "unknown": 6, +} +_HTTP_ERROR_RE = re.compile(r"\bHTTP error\s+(\d{3})\b", re.IGNORECASE) + + +@dataclass(frozen=True) +class FailureInfo: + """A non-sensitive description of one connection failure.""" + + kind: FailureKind + status: int | None = None + reason: str | None = None + retry_after: float | None = None + request_method: str | None = None + request_path: str | None = None + + @property + def retryable(self) -> bool: + return self.kind != "auth" + + +def _retry_after(value: str | None) -> float | None: + if not value: + return None + try: + return max(0.0, float(value)) + except ValueError: + pass + try: + date = parsedate_to_datetime(value) + if date.tzinfo is None: + date = date.replace(tzinfo=UTC) + return max(0.0, (date - datetime.now(UTC)).total_seconds()) + except (TypeError, ValueError, OverflowError): + return None + + +def _from_status( + status: int, + reason: str | None = None, + retry_after: float | None = None, + *, + request_method: str | None = None, + request_path: str | None = None, +) -> FailureInfo: + if status in (401, 403): + kind: FailureKind = "auth" + elif status == 429: + kind = "rate_limit" + elif 500 <= status <= 599: + kind = "server" + elif 400 <= status <= 499: + kind = "protocol" + else: + kind = "unknown" + return FailureInfo( + kind, + status, + reason, + retry_after, + request_method, + request_path, + ) + + +def _direct(exc: BaseException) -> FailureInfo | None: + if isinstance(exc, httpx.HTTPStatusError): + response = exc.response + request = response.request + return _from_status( + response.status_code, + response.reason_phrase, + _retry_after(response.headers.get("Retry-After")), + request_method=request.method, + request_path=request.url.path, + ) + if isinstance(exc, httpx.TimeoutException): + return FailureInfo("timeout", reason="request timed out") + if isinstance(exc, httpx.TransportError): + return FailureInfo("transport", reason="transport error") + if isinstance(exc, McpError): + return FailureInfo("protocol", reason="MCP protocol error") + if isinstance(exc, UserError): + match = _HTTP_ERROR_RE.search(str(exc)) + if match: + return _from_status(int(match.group(1))) + return None + + +def classify(exc: BaseException) -> FailureInfo: + """Return the most specific non-sensitive classification in an exception tree.""" + direct = _direct(exc) + matches: list[FailureInfo] = [direct] if direct is not None else [] + if isinstance(exc, BaseExceptionGroup): + group = cast("BaseExceptionGroup[BaseException]", exc) + matches.extend(classify(child) for child in group.exceptions) + if matches: + return min(matches, key=lambda info: _PRIORITY[info.kind]) + return FailureInfo("unknown", reason="unknown failure") + + +class HttpStatusRecorder: + """Capture the last non-success response from one HTTP connection.""" + + def __init__(self) -> None: + self._failure: FailureInfo | None = None + + def __call__(self, response: httpx.Response) -> None: + if not 200 <= response.status_code < 300: + request = response.request + self._failure = _from_status( + response.status_code, + response.reason_phrase, + _retry_after(response.headers.get("Retry-After")), + request_method=request.method, + request_path=request.url.path, + ) + + def take(self) -> FailureInfo | None: + failure, self._failure = self._failure, None + return failure diff --git a/strix/tools/mcp/session.py b/strix/tools/mcp/session.py index d3062b45..07c73f32 100644 --- a/strix/tools/mcp/session.py +++ b/strix/tools/mcp/session.py @@ -28,10 +28,9 @@ lifetime, and ``cleanup()``. Three consequences: "connection unavailable" value instead of a cancellation propagating into the agent loop. -When a call fails the supervisor rebuilds and reconnects the session once (reusing -the same config, so the same bearer token, never re-fetching credentials) and -re-runs the one failed call once. If that still fails, the connection is marked -dead: every later call returns the standard failed-tool output. +When a call fails the supervisor classifies it, retries boundedly, and temporarily +quarantines transient failures before lazily reviving the session. Authentication +failures and repeated transient exhaustion permanently retire a connection. Security: the connection's :class:`~strix.tools.mcp.config.McpConnectionConfig` holds a live bearer credential and is kept here in memory only, on the same @@ -46,8 +45,12 @@ import asyncio import contextlib import dataclasses import logging +import secrets +import time from typing import TYPE_CHECKING, Any, cast +from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify + if TYPE_CHECKING: from collections.abc import Awaitable, Callable @@ -70,6 +73,17 @@ logger = logging.getLogger(__name__) # before the supervising task is cancelled instead. Bounds teardown so a slow or # hung in-flight call cannot stall it forever. _SHUTDOWN_TIMEOUT = 10.0 +_MAX_ATTEMPTS = 3 +_SETTLE_DELAY = 0.05 +_SEMAPHORES: dict[str, asyncio.Semaphore] = {} +_JITTER = secrets.SystemRandom() + + +def _retry_delay(attempt: int, retry_after: float | None) -> float: + if retry_after is not None: + return retry_after + base = min(8.0, 0.5 * (2 ** (attempt - 1))) + return base + _JITTER.uniform(0.0, base * 0.1) # type: ignore[no-any-return] class McpConnectionUnavailableError(RuntimeError): @@ -129,6 +143,14 @@ class SupervisedMcpSession: self._dead = False self._closing = False self._on_dead: Callable[[], None] | None = None + self._recorder: HttpStatusRecorder | None = None + self._unavailable_until: float | None = None + self._quarantine_count = 0 + self._last_failure = FailureInfo("unknown", reason="connection unavailable") + self._reconnect_lock = asyncio.Lock() + self._call_semaphore = _SEMAPHORES.setdefault( + self._name, asyncio.Semaphore(config.max_concurrent_calls) + ) # Guards the idle-death self-heal against a flapping server: set after an # idle reconnect, cleared once a real call runs. If the session dies idle # again before serving anything, we give up instead of reconnecting in a @@ -161,6 +183,14 @@ class SupervisedMcpSession: self._dead = False self._closing = False self._on_dead = None + self._recorder = None + self._unavailable_until = None + self._quarantine_count = 0 + self._last_failure = FailureInfo("unknown", reason="connection unavailable") + self._reconnect_lock = asyncio.Lock() + self._call_semaphore = _SEMAPHORES.setdefault( + name, asyncio.Semaphore(config.max_concurrent_calls if config else 4) + ) self._healed_without_progress = False return self @@ -185,6 +215,15 @@ class SupervisedMcpSession: def is_dead(self) -> bool: return self._dead + @property + def is_unavailable(self) -> bool: + """Whether the connection is temporarily quarantined.""" + return ( + not self._dead + and self._unavailable_until is not None + and time.monotonic() < self._unavailable_until + ) + def set_on_dead(self, callback: Callable[[], None] | None) -> None: """Register a one-shot callback fired when the connection transitions to dead. @@ -198,11 +237,22 @@ class SupervisedMcpSession: """ self._on_dead = callback - def _mark_dead(self) -> None: + def _mark_dead(self, failure: FailureInfo | None = None, *, attempt: int = 1) -> None: """Flip the connection to dead and fire ``on_dead`` once on the transition.""" if self._dead: return + failure = failure or self._last_failure self._dead = True + self._unavailable_until = None + logger.error( + "MCP connection %r permanently unavailable kind=%s status=%s reason=%s " + "attempt=%d delay=0", + self._name, + failure.kind, + failure.status, + failure.reason, + attempt, + ) callback = self._on_dead if callback is None: return @@ -280,7 +330,7 @@ class SupervisedMcpSession: # -- caller-facing operations -------------------------------------------- async def list_tools(self) -> list[MCPTool]: - """List the connection's tools, reconnecting once if the session died. + """List the connection's tools, retrying transient session failures. Raises :class:`McpConnectionUnavailableError` when the connection is dead. """ @@ -297,7 +347,7 @@ class SupervisedMcpSession: label: str, result_transform: ResultTransform | None = None, ) -> Any: - """Run one tool call, reconnecting once and retrying once on session death. + """Run one tool call with bounded retries for transient session failures. Returns the tool output on success, or the standard failed-tool output (``success: False``) with a "connection unavailable" message when the @@ -361,8 +411,14 @@ class SupervisedMcpSession: await self._safe_cleanup() self._fail_pending() return - except Exception: - logger.exception("Skipping MCP connection %r", self._name) + except BaseException as exc: # noqa: BLE001 - classify cancellation groups + failure = classify(exc) + logger.warning( + "Skipping MCP connection %r kind=%s status=%s attempt=1 delay=0", + self._name, + failure.kind, + failure.status, + ) self._report_ready(value=False) await self._safe_cleanup() self._fail_pending() @@ -391,24 +447,40 @@ class SupervisedMcpSession: # calls then short-circuit to the dead output without this task. if self._closing: raise + failure = self._recorder.take() if self._recorder is not None else None + failure = failure or FailureInfo("transport", reason="session cancelled") + self._last_failure = failure if not self._healed_without_progress: logger.warning( - "MCP connection %r session died while idle; reconnecting once", + "MCP connection %r session died while idle kind=%s status=%s " + "attempt=1 delay=0; reconnecting", self._name, + failure.kind, + failure.status, ) - if await self._reconnect(): + self._healed_without_progress = True + reconnected, reconnect_failure = await self._reconnect() + if reconnected: logger.info( - "MCP connection %r reconnected after an idle death", self._name + "MCP connection %r revived kind=%s status=%s attempt=1 delay=%.2f", + self._name, + failure.kind, + failure.status, + _SETTLE_DELAY, ) - self._healed_without_progress = True continue + if reconnect_failure is not None: + failure = reconnect_failure else: logger.warning( "MCP connection %r died again before serving a call; " - "marking it unavailable", + "marking it permanently unavailable kind=%s status=%s " + "attempt=1 delay=0", self._name, + failure.kind, + failure.status, ) - self._mark_dead() + self._mark_dead(failure) await self._safe_cleanup() return if request is None: # shutdown sentinel @@ -420,73 +492,138 @@ class SupervisedMcpSession: if not request.future.done(): request.future.set_result(outcome) self._pending.discard(request.future) + if self._dead: + return - # -- run one job with reconnect-once + retry-once ------------------------- + # -- run one job with bounded classified retries -------------------------- - async def _execute(self, job: Job) -> _Outcome: - """Run one job; on a session failure reconnect once and retry it once.""" - if self._dead or self._server is None: + 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) - try: - return _Outcome(value=await job(self._server)) - except asyncio.CancelledError: - # For a supervised session a cancellation here is the transport scope - # dying under an in-flight call: a session death, not a real cancel - # (shutdown never cancels the task, it uses the sentinel). For an - # adopted session there is no such scope, so a cancel is real. - if not self._supervised or self._closing: - raise - logger.warning( - "MCP connection %r was cancelled mid-call (session died); reconnecting once", + if self._unavailable_until is not None: + remaining = self._unavailable_until - time.monotonic() + if remaining > 0: + return _Outcome(dead=True) + self._unavailable_until = None + logger.info( + "MCP connection %r revive started kind=%s status=%s attempt=1 delay=%.2f", self._name, - ) - except Exception: # noqa: BLE001 - any call failure is treated as a session death - logger.warning( - "MCP connection %r failed mid-call; reconnecting once", self._name + self._last_failure.kind, + self._last_failure.status, + remaining, ) - if not await self._reconnect(): - self._mark_dead() - return _Outcome(dead=True) + failure: FailureInfo | None = None + for attempt in range(1, _MAX_ATTEMPTS + 1): + 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: + 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) + continue + assert self._server is not None + try: + async with self._call_semaphore: + return _Outcome(value=await self._run_job_call(job)) + except asyncio.CancelledError: + if not self._supervised or self._closing: + raise + failure = ( + self._recorder.take() if self._recorder is not None else None + ) or FailureInfo("transport", reason="session cancelled") + except BaseException as exc: # noqa: BLE001 - classify SDK groups + failure = classify(exc) + if failure.kind == "unknown" and self._recorder is not None: + failure = self._recorder.take() or failure - try: - return _Outcome(value=await job(self._server)) - except asyncio.CancelledError: - if not self._supervised or self._closing: - raise - logger.warning( - "MCP connection %r was cancelled again after reconnect; marking it unavailable", - self._name, - ) - self._mark_dead() - return _Outcome(dead=True) - except Exception: # noqa: BLE001 - any retry failure means the connection is dead - logger.warning( - "MCP connection %r failed again after reconnect; marking it unavailable", - self._name, - ) - self._mark_dead() - return _Outcome(dead=True) + self._last_failure = failure + if failure.kind == "auth": + self._mark_dead(failure, attempt=attempt) + return _Outcome(dead=True) + if attempt == _MAX_ATTEMPTS: + 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) + reconnected, reconnect_failure = await self._reconnect() + if not reconnected and reconnect_failure is not None: + self._mark_dead(reconnect_failure, attempt=attempt) + return _Outcome(dead=True) + return _Outcome(dead=True) - async def _reconnect(self) -> bool: - """Rebuild and reconnect the session once, reusing the stored config/token.""" - await self._safe_cleanup() - if self._config is None: - return False - try: - self._server = await self._open() - except asyncio.CancelledError: - if self._closing: - raise - logger.warning("MCP reconnect for %r was cancelled; giving up", self._name) - self._server = None - return False - except Exception: - logger.exception("MCP reconnect for %r failed", self._name) - self._server = None - return False - logger.info("MCP connection %r reconnected", self._name) - return True + async def _run_job_call(self, job: Job) -> Any: + return await job(self._server) # type: ignore[arg-type] + + 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", + self._name, + failure.kind, + failure.status, + attempt, + delay, + ) + + def _quarantine(self, failure: FailureInfo, *, attempt: int) -> None: + self._quarantine_count += 1 + if self._quarantine_count >= 3: + self._mark_dead(failure, attempt=attempt) + return + cooldown = 30.0 * (2 ** (self._quarantine_count - 1)) + self._unavailable_until = time.monotonic() + cooldown + logger.warning( + "MCP connection %r quarantined kind=%s status=%s attempt=%d delay=%.2f", + self._name, + failure.kind, + failure.status, + attempt, + cooldown, + ) + + async def _reconnect(self) -> tuple[bool, FailureInfo | None]: + """Rebuild and reconnect, reusing the stored config and settling briefly.""" + async with self._reconnect_lock: + await self._safe_cleanup() + if self._config is None: + return False, FailureInfo("transport", reason="no reconnect config") + try: + server = await self._open() + except asyncio.CancelledError: + if self._closing: + raise + self._server = None + return False, FailureInfo("transport", reason="reconnect cancelled") + except BaseException as exc: # noqa: BLE001 - classify SDK groups + 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 + # connect() is the only readiness surface exposed by the SDK. + self._server = server + try: + await asyncio.sleep(_SETTLE_DELAY) + except asyncio.CancelledError: + self._server = None + with contextlib.suppress(BaseException): + await server.cleanup() # type: ignore[no-untyped-call] + if self._closing: + raise + return False, FailureInfo("transport", reason="reconnect cancelled") + return True, None async def _open(self) -> MCPServer: """Build and connect the SDK server, reusing the existing setup steps. @@ -500,6 +637,7 @@ class SupervisedMcpSession: if self._config is None: raise RuntimeError(f"MCP connection {self._name!r} has no config to connect") server = _build_server(self._config) + self._recorder = getattr(server, "_strix_http_status_recorder", None) try: await server.connect() # type: ignore[no-untyped-call] except BaseException: @@ -529,8 +667,15 @@ class SupervisedMcpSession: self._pending.clear() def _unavailable_message(self) -> str: + if self._unavailable_until is not None: + remaining = max(0.0, self._unavailable_until - time.monotonic()) + return ( + f"MCP connection {self._name!r} is temporarily unavailable " + f"(kind={self._last_failure.kind}, status={self._last_failure.status}); " + f"retrying in about {remaining:.0f} seconds." + ) return ( - f"MCP connection {self._name!r} is unavailable: its live session could " - "not be reached and a reconnect attempt failed. It is marked unavailable " - "for the rest of this run." + f"MCP connection {self._name!r} is unavailable " + f"(kind={self._last_failure.kind}, status={self._last_failure.status}); " + "it will not be retried." ) diff --git a/tests/test_mcp_resilience.py b/tests/test_mcp_resilience.py new file mode 100644 index 00000000..bd91bf83 --- /dev/null +++ b/tests/test_mcp_resilience.py @@ -0,0 +1,297 @@ +"""Fast regression tests for MCP failure handling and lifecycle resilience.""" + +from __future__ import annotations + +import asyncio +import importlib +from datetime import UTC, datetime, timedelta +from typing import Any + +import httpx +import pytest +from agents.exceptions import UserError +from mcp.shared.exceptions import McpError +from mcp.types import ErrorData + +from strix.tools.mcp import BearerAuth, McpConnectionConfig +from strix.tools.mcp import client as mcp_client +from strix.tools.mcp import session as mcp_session +from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify + + +_test_mcp_client = importlib.import_module("tests.test_mcp_client") +FakeMCPServer: Any = _test_mcp_client.FakeMCPServer +_mcp_tool: Any = _test_mcp_client._mcp_tool + + +def _http_error(status: int, *, retry_after: str | None = None) -> httpx.HTTPStatusError: + request = httpx.Request( + "POST", + "https://provider.example/tools?token=secret-query", + headers={"Authorization": "Bearer secret-header"}, + content=b"secret-body", + ) + response = httpx.Response( + status, + request=request, + headers={"Retry-After": retry_after} if retry_after else None, + ) + return httpx.HTTPStatusError("provider failure", request=request, response=response) + + +@pytest.mark.parametrize( + ("exc", "kind"), + [ + (_http_error(401), "auth"), + (_http_error(429), "rate_limit"), + (_http_error(503), "server"), + (_http_error(404), "protocol"), + (httpx.ReadTimeout("timed out"), "timeout"), + (httpx.ConnectError("disconnected"), "transport"), + (McpError(ErrorData(code=-1, message="bad response")), "protocol"), + (UserError("Failed to call tool: HTTP error 403"), "auth"), + ], +) +def test_classifies_failures(exc: BaseException, kind: str) -> None: + assert classify(exc).kind == kind + + +def test_classifies_nested_exception_groups_by_specificity() -> None: + error = ExceptionGroup( + "outer", + [ExceptionGroup("inner", [httpx.ConnectError("down"), _http_error(401)])], + ) + info = classify(error) + assert info.kind == "auth" + assert info.status == 401 + assert info.retryable is False + + +def test_retry_after_parses_seconds_and_http_date() -> None: + seconds = HttpStatusRecorder() + seconds(_http_error(429, retry_after="12").response) + assert seconds.take() is not None + assert seconds.take() is None + + date = (datetime.now(UTC) + timedelta(seconds=20)).strftime("%a, %d %b %Y %H:%M:%S GMT") + recorder = HttpStatusRecorder() + recorder(_http_error(429, retry_after=date).response) + info = recorder.take() + assert info is not None + retry_after = info.retry_after + assert retry_after is not None + assert 0 <= retry_after <= 20 + + +def test_recorder_only_keeps_non_sensitive_request_metadata() -> None: + recorder = HttpStatusRecorder() + response = _http_error(500, retry_after="3").response + recorder(response) + info = recorder.take() + assert info == FailureInfo( + "server", + 500, + "Internal Server Error", + 3, + "POST", + "/tools", + ) + assert "secret" not in repr(info) + assert recorder.take() is None + + +def _config(name: str, **kwargs: Any) -> McpConnectionConfig: + return McpConnectionConfig( + name=name, + url="https://provider.example/mcp", + auth=BearerAuth(token="secret-token"), # noqa: S106 # nosec B106 + **kwargs, + ) + + +async def _no_sleep(_delay: float) -> None: + return None + + +def _zero_delay(_attempt: int, _retry_after: float | None) -> float: + return 0 + + +def _sequence_server(name: str, error: BaseException | None = None) -> Any: + server = FakeMCPServer(name, [_mcp_tool("read")]) + original_call_tool = server.call_tool + + async def call_tool(tool_name: str, arguments: dict[str, Any] | None, meta: Any = None) -> Any: + if error is not None: + raise error + return await original_call_tool(tool_name, arguments, meta) + + server.call_tool = call_tool + return server + + +@pytest.mark.asyncio +async def test_rate_limit_retries_and_succeeds( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(mcp_session, "_retry_delay", _zero_delay) + monkeypatch.setattr(asyncio, "sleep", _no_sleep) + builds = iter( + [ + _sequence_server("rate", _http_error(429, retry_after="0")), + _sequence_server("rate"), + ] + ) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(builds)) + session = mcp_session.SupervisedMcpSession(_config("rate")) + assert await session.start() + result = await session.dispatch("read", {}, label="rate_read") + assert result == {"type": "text", "text": "routed:read"} + assert session.is_dead is False + await session.aclose() + + +@pytest.mark.asyncio +async def test_server_exhaustion_quarantines_then_revives( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(mcp_session, "_retry_delay", _zero_delay) + monkeypatch.setattr(asyncio, "sleep", _no_sleep) + clock = [100.0] + monkeypatch.setattr("strix.tools.mcp.session.time.monotonic", lambda: clock[0]) + builds = iter( + [ + _sequence_server("quarantine", _http_error(500)), + _sequence_server("quarantine", _http_error(500)), + _sequence_server("quarantine", _http_error(500)), + _sequence_server("quarantine"), + ] + ) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(builds)) + session = mcp_session.SupervisedMcpSession(_config("quarantine")) + assert await session.start() + result = await session.dispatch("read", {}, label="quarantine_read") + assert result["success"] is False + assert session.is_dead is False + assert session.is_unavailable is True + clock[0] += 31 + result = await session.dispatch("read", {}, label="quarantine_read") + assert result == {"type": "text", "text": "routed:read"} + await session.aclose() + + +@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()) + session = mcp_session.SupervisedMcpSession(_config("auth")) + assert await session.start() + result = await session.dispatch("read", {}, label="auth_read") + assert result["success"] is False + assert session.is_dead is True + assert builds == [] + await session.aclose() + + +@pytest.mark.asyncio +async def test_cancelled_call_uses_recorded_status( + monkeypatch: pytest.MonkeyPatch, +) -> None: + recorder = HttpStatusRecorder() + first = _sequence_server("cancelled", asyncio.CancelledError()) + first._strix_http_status_recorder = recorder + second = _sequence_server("cancelled") + builds = iter([first, second]) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(builds)) + monkeypatch.setattr(mcp_session, "_retry_delay", _zero_delay) + monkeypatch.setattr(asyncio, "sleep", _no_sleep) + + session = mcp_session.SupervisedMcpSession(_config("cancelled")) + assert await session.start() + recorder(_http_error(503).response) + result = await session.dispatch("read", {}, label="cancelled_read") + assert result == {"type": "text", "text": "routed:read"} + await session.aclose() + + +@pytest.mark.asyncio +async def test_build_server_passes_explicit_http_values( + monkeypatch: pytest.MonkeyPatch, +) -> None: + captured: dict[str, Any] = {} + + class Server: + def __init__(self, **kwargs: Any) -> None: + captured.update(kwargs) + + monkeypatch.setattr(mcp_client, "MCPServerStreamableHttp", Server) + config = _config( + "values", + http_timeout_seconds=11, + sse_read_timeout_seconds=22, + session_timeout_seconds=33, + ) + mcp_client._build_server(config) + assert captured["params"]["timeout"] == 11 + assert captured["params"]["sse_read_timeout"] == 22 + assert captured["client_session_timeout_seconds"] == 33 + factory = captured["params"]["httpx_client_factory"] + client = factory(headers={}, timeout=httpx.Timeout(1), auth=None) + assert client.event_hooks["response"] + await client.aclose() + + +@pytest.mark.asyncio +async def test_same_name_sessions_share_concurrency_cap() -> None: + active = 0 + peak = 0 + + def slow_server() -> Any: + server = FakeMCPServer("cap", [_mcp_tool("read")]) + original_call_tool = server.call_tool + + async def call_tool( + tool_name: str, arguments: dict[str, Any] | None, meta: Any = None + ) -> Any: + nonlocal active, peak + active += 1 + peak = max(peak, active) + await asyncio.sleep(0.01) + active -= 1 + return await original_call_tool(tool_name, arguments, meta) + + server.call_tool = call_tool + return server + + first = slow_server() + second = slow_server() + config = _config("cap", max_concurrent_calls=1) + left = mcp_session.SupervisedMcpSession.adopt(first, name="cap", config=config) + right = mcp_session.SupervisedMcpSession.adopt(second, name="cap", config=config) + await asyncio.gather( + left.dispatch("read", {}, label="cap_read"), + right.dispatch("read", {}, label="cap_read"), + ) + assert peak == 1 + await left.aclose() + await right.aclose() + + +@pytest.mark.asyncio +async def test_resilience_logs_do_not_include_request_secrets( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + 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) + session = mcp_session.SupervisedMcpSession(_config("redaction")) + assert await session.start() + with caplog.at_level("WARNING"): + await session.dispatch("read", {}, label="redaction_read") + assert "secret-token" not in caplog.text + assert "secret-query" not in caplog.text + assert "secret-header" not in caplog.text + assert "secret-body" not in caplog.text + await session.aclose()