strix/tests/test_mcp_resilience.py

521 lines
18 KiB
Python

"""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(403), "permission"),
(_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"), "permission"),
],
)
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_classifies_permission_before_rate_limit() -> None:
error = ExceptionGroup("outer", [_http_error(429), _http_error(403)])
info = classify(error)
assert info.kind == "permission"
assert info.status == 403
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
def _list_tools_error_server(name: str, error: BaseException) -> Any:
server = FakeMCPServer(name, [_mcp_tool("read")])
async def list_tools(*_args: Any, **_kwargs: Any) -> Any:
raise error
server.list_tools = list_tools
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_success_resets_quarantine_strikes(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# A quarantine strike must be cleared by a successful revival, so transient
# failure bursts separated by successes do not accumulate toward permanent
# retirement. Without the reset, three such bursts would mark the connection
# dead even though it recovered between each one.
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("strikes", _http_error(500)),
_sequence_server("strikes", _http_error(500)),
_sequence_server("strikes", _http_error(500)),
_sequence_server("strikes"),
]
)
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(builds)))
session = mcp_session.SupervisedMcpSession(_config("strikes"))
assert await session.start()
# First burst exhausts three attempts and quarantines: one strike.
result = await session.dispatch("read", {}, label="strikes_read")
assert result["success"] is False
assert session._quarantine_count == 1
# The revive succeeds, which must clear the strike back to zero.
clock[0] += 31
result = await session.dispatch("read", {}, label="strikes_read")
assert result == {"type": "text", "text": "routed:read"}
assert session._quarantine_count == 0
assert session.is_dead is False
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.parametrize(
("status", "name"),
[(403, "permission-call"), (400, "protocol-call")],
)
@pytest.mark.asyncio
async def test_call_http_rejection_preserves_session(
monkeypatch: pytest.MonkeyPatch,
status: int,
name: str,
) -> None:
first = _sequence_server(name, _http_error(status))
second = _sequence_server(name)
builds = iter([first, second])
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(builds)))
session = mcp_session.SupervisedMcpSession(_config(name))
assert await session.start()
result = await session.dispatch("read", {}, label=f"{name}_read")
assert result["success"] is False
assert "not the connection" in result["content"]
assert session.is_dead is False
assert session.is_unavailable is False
assert session._quarantine_count == 0
result = await session.dispatch("read", {}, label=f"{name}_read")
assert result == {"type": "text", "text": "routed:read"}
assert session.is_dead is False
await session.aclose()
@pytest.mark.asyncio
async def test_call_jsonrpc_error_preserves_session(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# A JSON-RPC error is a well-formed reply to this request, so the session stays
# up: no reconnect, no retry, no quarantine. The streamable-HTTP client also
# synthesizes one (status-less "Session terminated") for an HTTP 404, which some
# providers return for a missing resource.
error = McpError(ErrorData(code=32600, message="Session terminated"))
builds = 0
def build(_config: Any) -> Any:
nonlocal builds
builds += 1
return _built_server(_sequence_server("rpc-error", error))
monkeypatch.setattr(mcp_client, "_build_server", build)
monkeypatch.setattr(mcp_session, "_retry_delay", _zero_delay)
session = mcp_session.SupervisedMcpSession(_config("rpc-error"))
assert await session.start()
result = await session.dispatch("read", {}, label="rpc_error_read")
assert result["success"] is False
assert "not the connection" in result["content"]
assert "still available" in result["content"]
assert session.is_dead is False
assert session.is_unavailable is False
assert session._quarantine_count == 0
assert builds == 1
await session.aclose()
@pytest.mark.asyncio
async def test_list_tools_during_quarantine_reports_temporary_state(
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("cooldown", _http_error(500)) for _ in range(3)])
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(next(builds)))
session = mcp_session.SupervisedMcpSession(_config("cooldown"))
assert await session.start()
await session.dispatch("read", {}, label="cooldown_read")
assert session.is_unavailable is True
with pytest.raises(mcp_session.McpConnectionUnavailableError) as excinfo:
await session.list_tools()
message = str(excinfo.value)
assert "temporarily unavailable" in message
assert "retrying in about 30 seconds" in message
assert "rest of this run" not in message
await session.aclose()
@pytest.mark.asyncio
async def test_call_http_403_during_list_tools_dies() -> None:
server = _list_tools_error_server("connect-403", _http_error(403))
session = mcp_session.SupervisedMcpSession.adopt(
server,
name="connect-403",
config=_config("connect-403"),
)
with pytest.raises(mcp_session.McpConnectionUnavailableError):
await session.list_tools()
assert session.is_dead is True
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()