Harden MCP client transport resilience

This commit is contained in:
yoni 2026-08-28 03:14:39 +00:00
parent 717ffc8f4c
commit ad41517613
6 changed files with 703 additions and 78 deletions

View file

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

View file

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

View file

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

145
strix/tools/mcp/failures.py Normal file
View file

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

View file

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

View file

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