mirror of
https://github.com/usestrix/strix.git
synced 2026-10-08 03:08:08 +00:00
Harden MCP client transport resilience
This commit is contained in:
parent
717ffc8f4c
commit
ad41517613
6 changed files with 703 additions and 78 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
145
strix/tools/mcp/failures.py
Normal 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
|
||||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
297
tests/test_mcp_resilience.py
Normal file
297
tests/test_mcp_resilience.py
Normal 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()
|
||||
Loading…
Add table
Reference in a new issue