mirror of
https://github.com/usestrix/strix.git
synced 2026-08-28 05:25:00 +00:00
Merge b29b2241b2 into 717ffc8f4c
This commit is contained in:
commit
b9da3266fe
7 changed files with 900 additions and 169 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
|
||||
|
||||
|
|
@ -71,6 +75,13 @@ class ConnectedMcpServer(NamedTuple):
|
|||
notes: str | None = None
|
||||
|
||||
|
||||
class BuiltMcpServer(NamedTuple):
|
||||
"""A constructed SDK server and its optional HTTP failure recorder."""
|
||||
|
||||
server: MCPServer
|
||||
recorder: HttpStatusRecorder | None
|
||||
|
||||
|
||||
def _auth_headers(config: McpConnectionConfig) -> dict[str, str]:
|
||||
"""Build the per-server request headers from the connection's auth."""
|
||||
auth = config.auth
|
||||
|
|
@ -108,9 +119,12 @@ class _QuietMCPServerStdio(MCPServerStdio):
|
|||
return _quiet_stdio_streams(self.params)
|
||||
|
||||
|
||||
def _build_server(config: McpConnectionConfig) -> MCPServer:
|
||||
def _build_server(config: McpConnectionConfig) -> BuiltMcpServer:
|
||||
"""Construct (but do not connect) the SDK server for one connection.
|
||||
|
||||
The returned tuple carries the server and, for HTTP connections, a recorder
|
||||
that retains sanitized response metadata for the owning session.
|
||||
|
||||
When ``allowed_tools`` is a list the static filter means the server will not
|
||||
even list tools outside it, so it is the authoritative gate on what
|
||||
``describe_mcp`` and ``call_mcp`` can see. When it is ``None`` no filter is
|
||||
|
|
@ -128,22 +142,43 @@ def _build_server(config: McpConnectionConfig) -> MCPServer:
|
|||
"args": config.args,
|
||||
"env": config.env,
|
||||
}
|
||||
return _QuietMCPServerStdio(
|
||||
params=stdio_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
return BuiltMcpServer(
|
||||
_QuietMCPServerStdio(
|
||||
params=stdio_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
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(
|
||||
params=http_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
return BuiltMcpServer(
|
||||
MCPServerStreamableHttp(
|
||||
params=http_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
client_session_timeout_seconds=config.session_timeout_seconds,
|
||||
),
|
||||
recorder,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ from typing import Annotated, Literal
|
|||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
DEFAULT_MAX_CONCURRENT_CALLS = 4
|
||||
|
||||
|
||||
class BearerAuth(BaseModel):
|
||||
"""Header-token auth, sent as ``Authorization: Bearer <token>``."""
|
||||
|
||||
|
|
@ -65,6 +68,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=DEFAULT_MAX_CONCURRENT_CALLS, 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
|
||||
|
||||
async 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,14 @@ import asyncio
|
|||
import contextlib
|
||||
import dataclasses
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
import weakref
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from strix.tools.mcp.config import DEFAULT_MAX_CONCURRENT_CALLS
|
||||
from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
|
@ -70,6 +75,31 @@ 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: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, dict[str, asyncio.Semaphore]] = (
|
||||
weakref.WeakKeyDictionary()
|
||||
)
|
||||
_JITTER = secrets.SystemRandom()
|
||||
|
||||
# Everything the SDK can surface for a failed call: ordinary errors plus the
|
||||
# transport's task-group ``BaseExceptionGroup``. Caught wholesale and handed to
|
||||
# ``classify``; ``asyncio.CancelledError`` is always handled separately first,
|
||||
# so shutdown and genuine cancellation still propagate.
|
||||
_CLASSIFIABLE: tuple[type[BaseException], ...] = (BaseExceptionGroup, Exception)
|
||||
|
||||
|
||||
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]
|
||||
|
||||
|
||||
def _call_semaphore(name: str, limit: int) -> asyncio.Semaphore:
|
||||
loop = asyncio.get_running_loop()
|
||||
semaphores = _SEMAPHORES.setdefault(loop, {})
|
||||
return semaphores.setdefault(name, asyncio.Semaphore(limit))
|
||||
|
||||
|
||||
class McpConnectionUnavailableError(RuntimeError):
|
||||
|
|
@ -129,11 +159,12 @@ class SupervisedMcpSession:
|
|||
self._dead = False
|
||||
self._closing = False
|
||||
self._on_dead: Callable[[], None] | None = None
|
||||
# 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
|
||||
# tight loop.
|
||||
self._healed_without_progress = False
|
||||
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: asyncio.Semaphore | None = None
|
||||
|
||||
@classmethod
|
||||
def adopt(
|
||||
|
|
@ -147,7 +178,7 @@ class SupervisedMcpSession:
|
|||
|
||||
Calls run inline against ``server`` on the caller's task, matching the old
|
||||
direct-dispatch behavior. Reconnect is available only when ``config`` is
|
||||
given; otherwise a failed call marks the connection dead.
|
||||
given; otherwise a failed call can be quarantined but cannot be revived.
|
||||
"""
|
||||
self = cls.__new__(cls)
|
||||
self._name = name
|
||||
|
|
@ -161,7 +192,12 @@ class SupervisedMcpSession:
|
|||
self._dead = False
|
||||
self._closing = False
|
||||
self._on_dead = None
|
||||
self._healed_without_progress = False
|
||||
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 = None
|
||||
return self
|
||||
|
||||
# -- read-only accessors --------------------------------------------------
|
||||
|
|
@ -185,6 +221,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 +243,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 +336,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 +353,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 +417,15 @@ class SupervisedMcpSession:
|
|||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
return
|
||||
except Exception:
|
||||
logger.exception("Skipping MCP connection %r", self._name)
|
||||
except _CLASSIFIABLE as exc:
|
||||
failure = classify(exc)
|
||||
logger.warning(
|
||||
"Skipping MCP connection %r kind=%s status=%s attempt=1 delay=0",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
exc_info=True,
|
||||
)
|
||||
self._report_ready(value=False)
|
||||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
|
|
@ -384,109 +447,167 @@ class SupervisedMcpSession:
|
|||
# A cancellation while idle is the transport's task group cancelling
|
||||
# this supervising task because a background session task failed.
|
||||
# Contained here. If we are closing, this is an ordinary shutdown,
|
||||
# so let it propagate. Otherwise try to self-heal once: reconnect a
|
||||
# fresh session and keep serving. The flag stops a flapping server
|
||||
# (one that dies again before serving any call) from reconnecting in
|
||||
# a tight loop; there we give up and mark the connection dead. Later
|
||||
# calls then short-circuit to the dead output without this task.
|
||||
# so let it propagate. Otherwise quarantine the failed session and
|
||||
# keep serving requests so a later call can revive it.
|
||||
if self._closing:
|
||||
raise
|
||||
if not self._healed_without_progress:
|
||||
logger.warning(
|
||||
"MCP connection %r session died while idle; reconnecting once",
|
||||
self._name,
|
||||
)
|
||||
if await self._reconnect():
|
||||
logger.info(
|
||||
"MCP connection %r reconnected after an idle death", self._name
|
||||
)
|
||||
self._healed_without_progress = True
|
||||
continue
|
||||
else:
|
||||
logger.warning(
|
||||
"MCP connection %r died again before serving a call; "
|
||||
"marking it unavailable",
|
||||
self._name,
|
||||
)
|
||||
self._mark_dead()
|
||||
await self._safe_cleanup()
|
||||
return
|
||||
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 failure.kind == "auth":
|
||||
self._mark_dead(failure, attempt=1)
|
||||
return
|
||||
await self._quarantine(failure, attempt=1)
|
||||
if self._dead:
|
||||
return
|
||||
continue
|
||||
if request is None: # shutdown sentinel
|
||||
return
|
||||
outcome = await self._execute(request.job)
|
||||
# A served call is real progress: clear the idle-heal guard so a future
|
||||
# idle death is again allowed one reconnect.
|
||||
self._healed_without_progress = False
|
||||
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",
|
||||
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,
|
||||
)
|
||||
|
||||
if not await self._reconnect():
|
||||
self._mark_dead()
|
||||
return _Outcome(dead=True)
|
||||
|
||||
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",
|
||||
if self._call_semaphore is None:
|
||||
self._call_semaphore = _call_semaphore(
|
||||
self._name,
|
||||
(
|
||||
self._config.max_concurrent_calls
|
||||
if self._config is not None
|
||||
else DEFAULT_MAX_CONCURRENT_CALLS
|
||||
),
|
||||
)
|
||||
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)
|
||||
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"
|
||||
)
|
||||
outcome = await self._handle_failure(failure, attempt)
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
continue
|
||||
assert self._server is not None
|
||||
call_semaphore = self._call_semaphore
|
||||
assert call_semaphore is not None
|
||||
try:
|
||||
async with call_semaphore:
|
||||
return _Outcome(value=await job(self._server))
|
||||
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 _CLASSIFIABLE as exc:
|
||||
failure = classify(exc)
|
||||
if failure.kind == "unknown" and self._recorder is not None:
|
||||
failure = self._recorder.take() or failure
|
||||
|
||||
async def _reconnect(self) -> bool:
|
||||
"""Rebuild and reconnect the session once, reusing the stored config/token."""
|
||||
outcome = await self._handle_failure(failure, attempt)
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
reconnected, reconnect_failure = await self._reconnect()
|
||||
if not reconnected:
|
||||
failure = reconnect_failure or FailureInfo("transport", reason="reconnect failed")
|
||||
outcome = await self._handle_failure(failure, attempt)
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
return _Outcome(dead=True)
|
||||
|
||||
async def _handle_failure(self, failure: FailureInfo, attempt: int) -> _Outcome | None:
|
||||
self._last_failure = failure
|
||||
if failure.kind == "auth":
|
||||
self._mark_dead(failure, attempt=attempt)
|
||||
return _Outcome(dead=True)
|
||||
if attempt == _MAX_ATTEMPTS:
|
||||
await self._quarantine(failure, attempt=attempt)
|
||||
return _Outcome(dead=True)
|
||||
delay = _retry_delay(attempt, failure.retry_after)
|
||||
self._log_retry(failure, attempt, delay)
|
||||
await asyncio.sleep(delay)
|
||||
return None
|
||||
|
||||
def _log_retry(self, failure: FailureInfo, attempt: int, delay: float) -> None:
|
||||
logger.warning(
|
||||
"MCP connection %r retryable failure kind=%s status=%s attempt=%d delay=%.2f",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
attempt,
|
||||
delay,
|
||||
)
|
||||
|
||||
async def _quarantine(self, failure: FailureInfo, *, attempt: int) -> None:
|
||||
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
|
||||
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 _CLASSIFIABLE as exc:
|
||||
self._server = None
|
||||
failure = classify(exc)
|
||||
if failure.kind == "unknown" and self._recorder is not None:
|
||||
failure = self._recorder.take() or failure
|
||||
return False, failure
|
||||
# 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(Exception):
|
||||
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.
|
||||
|
|
@ -499,10 +620,16 @@ 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)
|
||||
built = _build_server(self._config)
|
||||
server = built.server
|
||||
self._recorder = built.recorder
|
||||
try:
|
||||
await server.connect() # type: ignore[no-untyped-call]
|
||||
except BaseException:
|
||||
except asyncio.CancelledError:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
raise
|
||||
except _CLASSIFIABLE:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
raise
|
||||
|
|
@ -529,8 +656,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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ import asyncio
|
|||
import contextlib
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
|
@ -160,11 +162,15 @@ def _config(name: str, allowed_tools: list[str] | None) -> McpConnectionConfig:
|
|||
return McpConnectionConfig(
|
||||
name=name,
|
||||
url="https://mcp.example.com",
|
||||
auth=BearerAuth(token="run-token"),
|
||||
auth=BearerAuth(token="run-token"), # nosec B106
|
||||
allowed_tools=allowed_tools,
|
||||
)
|
||||
|
||||
|
||||
def _built_server(server: MCPServer) -> mcp_client.BuiltMcpServer:
|
||||
return mcp_client.BuiltMcpServer(server, None)
|
||||
|
||||
|
||||
def _ctx(registry: McpRegistry | None) -> ToolContext[dict[str, Any]]:
|
||||
context: dict[str, Any] = {} if registry is None else {MCP_REGISTRY_CONTEXT_KEY: registry}
|
||||
return ToolContext(
|
||||
|
|
@ -204,7 +210,7 @@ def test_bearer_config_parses_from_dict() -> None:
|
|||
)
|
||||
|
||||
assert isinstance(config.auth, BearerAuth)
|
||||
assert config.auth.token == "abc"
|
||||
assert config.auth.token == "abc" # nosec B105
|
||||
assert config.allowed_tools == ["list_files"]
|
||||
|
||||
|
||||
|
|
@ -298,7 +304,9 @@ async def test_connect_returns_sessions_without_registering_agent_tools(
|
|||
"fs": FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]),
|
||||
"db": FakeMCPServer("db", [_mcp_tool("query")]),
|
||||
}
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_build_server", lambda config: _built_server(servers[config.name])
|
||||
)
|
||||
|
||||
connections = await mcp_client.connect_mcp_servers(
|
||||
[_config("fs", None), _config("db", ["query"])]
|
||||
|
|
@ -315,7 +323,7 @@ async def test_connect_returns_sessions_without_registering_agent_tools(
|
|||
@pytest.mark.asyncio
|
||||
async def test_tool_count_honors_the_allowlist(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
server = FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
|
||||
connections = await mcp_client.connect_mcp_servers([_config("fs", ["read_file"])])
|
||||
|
||||
|
|
@ -329,7 +337,7 @@ async def test_connection_notes_ride_on_the_connection(
|
|||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
server = FakeMCPServer("db", [_mcp_tool("query")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
config = McpConnectionConfig(
|
||||
name="db",
|
||||
url="https://mcp.example.com",
|
||||
|
|
@ -356,7 +364,7 @@ def test_build_server_stdio_branch() -> None:
|
|||
env={"TOKEN": "x"},
|
||||
)
|
||||
|
||||
server = mcp_client._build_server(config)
|
||||
server = mcp_client._build_server(config).server
|
||||
|
||||
assert isinstance(server, MCPServerStdio)
|
||||
assert server.name == "local_fs"
|
||||
|
|
@ -366,7 +374,7 @@ def test_build_server_stdio_branch() -> None:
|
|||
|
||||
|
||||
def test_build_server_http_branch() -> None:
|
||||
server = mcp_client._build_server(_config("files_main", ["list_files"]))
|
||||
server = mcp_client._build_server(_config("files_main", ["list_files"])).server
|
||||
|
||||
assert isinstance(server, MCPServerStreamableHttp)
|
||||
assert server.name == "files_main"
|
||||
|
|
@ -847,7 +855,9 @@ async def test_connect_skips_a_connection_whose_connect_is_cancelled(
|
|||
cleaned.append(self._name)
|
||||
|
||||
servers = {"good": _Tracking("good"), "bad": _Tracking("bad", cancel_connect=True)}
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_build_server", lambda config: _built_server(servers[config.name])
|
||||
)
|
||||
|
||||
configs = [_config("good", ["t"]), _config("bad", ["t"])]
|
||||
|
||||
|
|
@ -883,7 +893,9 @@ async def test_connect_cleans_up_started_sessions_when_attach_is_cancelled(
|
|||
cleaned.append(self._name)
|
||||
|
||||
servers = {"good": _Tracking("good"), "slow": _Tracking("slow", block_connect=True)}
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_build_server", lambda config: _built_server(servers[config.name])
|
||||
)
|
||||
|
||||
async def _attach() -> list[Any]:
|
||||
# Connect "good" first, then hang forever connecting "slow".
|
||||
|
|
@ -962,7 +974,7 @@ async def test_attach_populates_registry_with_provider_and_transform(
|
|||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
server = FakeMCPServer("db", [_mcp_tool("query")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
|
||||
def transform(_label: str, structured: Any) -> Any:
|
||||
return {"kept": structured}
|
||||
|
|
@ -995,7 +1007,7 @@ async def test_attach_bare_request_matches_the_command_line_shape(
|
|||
# The command-line path wraps each config in a bare request (no provider or
|
||||
# transform); purpose then falls back to the connection's notes.
|
||||
server = FakeMCPServer("db", [_mcp_tool("query")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
config = McpConnectionConfig(
|
||||
name="db",
|
||||
url="https://mcp.example.com",
|
||||
|
|
@ -1026,7 +1038,9 @@ async def test_attach_is_fail_open_and_skips_a_failed_connection(
|
|||
raise RuntimeError("cannot reach server")
|
||||
|
||||
servers = {"good": good, "bad": _Failing("bad", [_mcp_tool("t")])}
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_build_server", lambda config: _built_server(servers[config.name])
|
||||
)
|
||||
|
||||
registry = McpRegistry()
|
||||
connections = await attach_mcp_requests(
|
||||
|
|
@ -1213,7 +1227,7 @@ def _secret_config(name: str) -> McpConnectionConfig:
|
|||
return McpConnectionConfig(
|
||||
name=name,
|
||||
url="https://mcp.example.com",
|
||||
auth=BearerAuth(token="super-secret-bearer-token-42"),
|
||||
auth=BearerAuth(token="super-secret-bearer-token-42"), # nosec B106
|
||||
allowed_tools=["read_file"],
|
||||
)
|
||||
|
||||
|
|
@ -1234,7 +1248,7 @@ async def test_call_mcp_reconnects_and_retries_after_a_session_death(
|
|||
first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
||||
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
built = iter([first, second])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built)))
|
||||
|
||||
session = await _started_session(_secret_config("fs"))
|
||||
registry = McpRegistry()
|
||||
|
|
@ -1257,15 +1271,15 @@ async def test_call_mcp_reconnects_and_retries_after_a_session_death(
|
|||
async def test_call_mcp_marks_connection_dead_when_reconnect_keeps_failing(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# The session dies and the reconnect attempt also fails: the connection is
|
||||
# marked dead and the call returns the standard failed-tool output.
|
||||
# Reconnect failures are retried and then quarantine the connection rather
|
||||
# than permanently retiring it on the first failed reconnect.
|
||||
first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
||||
built = {"n": 0}
|
||||
|
||||
def _build(_config: McpConnectionConfig) -> MCPServer:
|
||||
def _build(_config: McpConnectionConfig) -> mcp_client.BuiltMcpServer:
|
||||
built["n"] += 1
|
||||
if built["n"] == 1:
|
||||
return first
|
||||
return _built_server(first)
|
||||
raise ConnectionError("cannot reconnect")
|
||||
|
||||
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
||||
|
|
@ -1282,9 +1296,12 @@ async def test_call_mcp_marks_connection_dead_when_reconnect_keeps_failing(
|
|||
assert isinstance(out, dict)
|
||||
assert out["success"] is False
|
||||
assert "unavailable" in out["content"]
|
||||
assert session.is_dead is True
|
||||
assert session.is_dead is False
|
||||
assert session.is_unavailable is True
|
||||
assert session.server is None
|
||||
assert session._task is not None and not session._task.done()
|
||||
|
||||
# A later call short-circuits to the same failed output without a new attempt.
|
||||
# A later call during cooldown short-circuits to the same failed output.
|
||||
again = await call_mcp.on_invoke_tool(
|
||||
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
||||
)
|
||||
|
|
@ -1309,18 +1326,21 @@ async def _pump_until(predicate: Callable[[], bool], *, limit: int = 100) -> Non
|
|||
raise AssertionError("condition not reached")
|
||||
|
||||
|
||||
def _quarantine_reached(session: SupervisedMcpSession, count: int) -> bool:
|
||||
return session.is_dead or session._quarantine_count >= count
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_idle_session_death_self_heals_on_reconnect(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# A session that dies while idle (its supervising task cancelled between calls,
|
||||
# modeling the transport scope dying with no call in flight) reconnects once on
|
||||
# its own and keeps serving, rather than staying dead until a later call would
|
||||
# have triggered a reconnect.
|
||||
# modeling the transport scope dying with no call in flight) is quarantined and
|
||||
# keeps serving, rather than ending its supervising task.
|
||||
first = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
built = iter([first, second])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built)))
|
||||
|
||||
session = await _started_session(_secret_config("fs"))
|
||||
registry = McpRegistry()
|
||||
|
|
@ -1328,15 +1348,16 @@ async def test_idle_session_death_self_heals_on_reconnect(
|
|||
|
||||
assert session._task is not None
|
||||
session._task.cancel() # idle transport death: no call in flight
|
||||
await _pump_until(lambda: session.server is second)
|
||||
await _pump_until(lambda: session.is_unavailable)
|
||||
assert session.is_dead is False
|
||||
assert session.server is None
|
||||
|
||||
# The reconnected session serves calls normally.
|
||||
# Once the cooldown expires, the next call reconnects onto a fresh session.
|
||||
session._unavailable_until = time.monotonic() - 1
|
||||
out = await call_mcp.on_invoke_tool(
|
||||
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
||||
)
|
||||
assert out == {"type": "text", "text": "routed:read_file"}
|
||||
|
||||
await session.aclose()
|
||||
|
||||
|
||||
|
|
@ -1344,26 +1365,24 @@ async def test_idle_session_death_self_heals_on_reconnect(
|
|||
async def test_flapping_idle_session_is_marked_dead_without_looping(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# If a session reconnects after an idle death but dies again before serving any
|
||||
# call, the supervisor stops reconnecting and marks the connection dead, so a
|
||||
# server that instantly drops on connect cannot spin in a reconnect loop.
|
||||
# Repeated idle deaths consume quarantine slots; the supervisor stays alive
|
||||
# until the configured permanent-death threshold is reached.
|
||||
first = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
built = iter([first, second])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(built)))
|
||||
|
||||
session = await _started_session(_secret_config("fs"))
|
||||
registry = McpRegistry()
|
||||
registry.add(name="fs", session=session, tool_count=1)
|
||||
|
||||
assert session._task is not None
|
||||
# First idle death heals onto the second server (only two builds ever happen).
|
||||
session._task.cancel()
|
||||
await _pump_until(lambda: session.server is second)
|
||||
assert session.is_dead is False
|
||||
|
||||
# Second idle death before any call is served: give up rather than reconnect.
|
||||
session._task.cancel()
|
||||
# Three idle deaths exhaust the quarantine budget.
|
||||
for count in range(1, 4):
|
||||
session._task.cancel()
|
||||
await _pump_until(partial(_quarantine_reached, session, count))
|
||||
if session.is_dead:
|
||||
break
|
||||
await _pump_until(lambda: session._task is not None and session._task.done())
|
||||
assert session.is_dead is True
|
||||
|
||||
|
|
@ -1420,7 +1439,7 @@ async def test_aclose_is_bounded_when_an_in_flight_call_hangs(
|
|||
# aclose falls back to cancelling the supervising task, and cleanup still runs.
|
||||
monkeypatch.setattr(mcp_session_mod, "_SHUTDOWN_TIMEOUT", 0.2)
|
||||
server = _HangingCallServer("fs", [_mcp_tool("read_file")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
|
||||
session = await _started_session(_secret_config("fs"))
|
||||
call = asyncio.create_task(session.dispatch("read_file", {}, label="fs_read_file"))
|
||||
|
|
@ -1444,7 +1463,7 @@ async def test_aclose_cleans_up_when_connect_is_cancelled_mid_await(
|
|||
# is cancelled; aclose must not raise on it and must still cancel + clean up the
|
||||
# partially connected supervisor.
|
||||
server = _HangingConnectServer("fs", [_mcp_tool("read_file")])
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
||||
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server))
|
||||
|
||||
session = SupervisedMcpSession(_secret_config("fs"))
|
||||
start = asyncio.create_task(session.start())
|
||||
|
|
@ -1469,14 +1488,14 @@ async def test_a_session_death_is_contained_and_other_connections_survive(
|
|||
healthy = FakeMCPServer("healthy", [_mcp_tool("read_file")])
|
||||
dying_builds = {"n": 0}
|
||||
|
||||
def _build(config: McpConnectionConfig) -> MCPServer:
|
||||
def _build(config: McpConnectionConfig) -> mcp_client.BuiltMcpServer:
|
||||
if config.name == "healthy":
|
||||
return healthy
|
||||
return _built_server(healthy)
|
||||
# The dying connection connects once, then its rebuild raises, so it ends
|
||||
# up marked dead rather than recovering.
|
||||
dying_builds["n"] += 1
|
||||
if dying_builds["n"] == 1:
|
||||
return dying
|
||||
return _built_server(dying)
|
||||
raise ConnectionError("cannot reconnect")
|
||||
|
||||
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
||||
|
|
@ -1513,11 +1532,13 @@ async def test_reconnect_reuses_the_stored_config_and_never_logs_the_token(
|
|||
# inventory list_mcps emits.
|
||||
seen_tokens: list[str | None] = []
|
||||
|
||||
def _build(config: McpConnectionConfig) -> MCPServer:
|
||||
def _build(config: McpConnectionConfig) -> mcp_client.BuiltMcpServer:
|
||||
seen_tokens.append(config.auth.token if config.auth else None)
|
||||
if len(seen_tokens) == 1:
|
||||
return _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
||||
return FakeMCPServer("fs", [_mcp_tool("read_file")])
|
||||
return _built_server(
|
||||
_DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
||||
)
|
||||
return _built_server(FakeMCPServer("fs", [_mcp_tool("read_file")]))
|
||||
|
||||
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
||||
|
||||
|
|
@ -1566,23 +1587,37 @@ class _RaisingMCPServer(FakeMCPServer):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_on_dead_fires_once_on_the_death_transition() -> None:
|
||||
# An adopted session with no config cannot reconnect, so the first failed
|
||||
# call marks it dead; the on-dead callback fires exactly once, on the edge.
|
||||
async def test_session_on_dead_fires_once_on_the_death_transition(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# An adopted session with no config cannot reconnect, so repeated transient
|
||||
# exhaustion eventually marks it dead; the callback fires once on that edge.
|
||||
server = _RaisingMCPServer("db", [_mcp_tool("read")])
|
||||
session = SupervisedMcpSession.adopt(server, name="db")
|
||||
fires: list[int] = []
|
||||
session.set_on_dead(lambda: fires.append(1))
|
||||
|
||||
clock = [100.0]
|
||||
monkeypatch.setattr("strix.tools.mcp.session.time.monotonic", lambda: clock[0])
|
||||
monkeypatch.setattr(mcp_session_mod, "_retry_delay", lambda _attempt, _retry_after: 0)
|
||||
|
||||
async def no_sleep(_delay: float) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(asyncio, "sleep", no_sleep)
|
||||
out = await session.dispatch("read", {}, label="db_read")
|
||||
|
||||
assert session.is_dead is True
|
||||
assert session.is_dead is False
|
||||
assert isinstance(out, dict) and out.get("success") is False
|
||||
assert fires == [1]
|
||||
assert fires == []
|
||||
|
||||
# A later call to the already-dead session must not fire the callback again.
|
||||
clock[0] += 31
|
||||
await session.dispatch("read", {}, label="db_read")
|
||||
assert session.is_dead is False
|
||||
clock[0] += 61
|
||||
await session.dispatch("read", {}, label="db_read")
|
||||
assert fires == [1]
|
||||
assert session.is_dead is True
|
||||
|
||||
|
||||
def test_registry_statuses_report_the_live_dead_flag_and_provider() -> None:
|
||||
|
|
|
|||
363
tests/test_mcp_resilience.py
Normal file
363
tests/test_mcp_resilience.py
Normal file
|
|
@ -0,0 +1,363 @@
|
|||
"""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, cast
|
||||
|
||||
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 _built_server(server: Any) -> Any:
|
||||
return mcp_client.BuiltMcpServer(server, None)
|
||||
|
||||
|
||||
def _http_error(status: int, *, retry_after: str | None = None) -> httpx.HTTPStatusError:
|
||||
request = httpx.Request(
|
||||
"POST",
|
||||
"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
|
||||
|
||||
|
||||
@pytest.mark.parametrize("control_flow", [SystemExit, KeyboardInterrupt])
|
||||
@pytest.mark.asyncio
|
||||
async def test_control_flow_exceptions_propagate(
|
||||
control_flow: type[BaseException],
|
||||
) -> None:
|
||||
server = _sequence_server("control-flow", control_flow("stop"))
|
||||
session = mcp_session.SupervisedMcpSession.adopt(server, name="control-flow")
|
||||
|
||||
with pytest.raises(control_flow):
|
||||
await session.dispatch("read", {}, label="control_flow")
|
||||
|
||||
await session.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_after_parses_seconds_and_http_date() -> None:
|
||||
seconds = HttpStatusRecorder()
|
||||
await 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()
|
||||
await 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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recorder_only_keeps_non_sensitive_request_metadata() -> None:
|
||||
recorder = HttpStatusRecorder()
|
||||
response = _http_error(500, retry_after="3").response
|
||||
await 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: _built_server(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: _built_server(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
|
||||
assert session.server is None
|
||||
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: _built_server(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())
|
||||
second = _sequence_server("cancelled")
|
||||
builds = iter(
|
||||
[
|
||||
mcp_client.BuiltMcpServer(first, recorder),
|
||||
mcp_client.BuiltMcpServer(second, None),
|
||||
]
|
||||
)
|
||||
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()
|
||||
await 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_http_factory_awaits_response_recorder() -> None:
|
||||
built = mcp_client._build_server(_config("hook"))
|
||||
assert built.recorder is not None
|
||||
factory = cast("Any", built.server).params["httpx_client_factory"]
|
||||
client = factory(headers={}, timeout=httpx.Timeout(1), auth=None)
|
||||
|
||||
def response(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(429, headers={"Retry-After": "7"}, request=request)
|
||||
|
||||
client._transport = httpx.MockTransport(response)
|
||||
result = await client.get("https://provider.example/mcp?token=secret-query")
|
||||
assert result.status_code == 429
|
||||
info = built.recorder.take()
|
||||
assert info is not None
|
||||
assert info.kind == "rate_limit"
|
||||
assert info.status == 429
|
||||
assert info.retry_after == 7
|
||||
assert info.request_method == "GET"
|
||||
assert info.request_path == "/mcp"
|
||||
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()
|
||||
|
||||
|
||||
def test_same_name_semaphore_works_across_event_loops() -> None:
|
||||
async def run_once() -> None:
|
||||
server = FakeMCPServer("loop-cap", [_mcp_tool("read")])
|
||||
session = mcp_session.SupervisedMcpSession.adopt(
|
||||
server,
|
||||
name="loop-cap",
|
||||
config=_config("loop-cap", max_concurrent_calls=1),
|
||||
)
|
||||
assert await session.dispatch("read", {}, label="loop_cap_read") == {
|
||||
"type": "text",
|
||||
"text": "routed:read",
|
||||
}
|
||||
await session.aclose()
|
||||
|
||||
asyncio.run(run_once())
|
||||
asyncio.run(run_once())
|
||||
|
||||
|
||||
@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: _built_server(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