diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 504219a64e1..508ab447326 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -3,8 +3,9 @@ import base64 import importlib import json import os +import selectors import sys -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Callable from pathlib import Path from types import ModuleType from typing import Final @@ -84,8 +85,8 @@ class _MockTransportClient(MCPClient): class _ManualClockLoop(asyncio.SelectorEventLoop): """An event loop whose clock moves only when the test advances it, so timeouts fire on test-controlled conditions""" - def __init__(self) -> None: - super().__init__() + def __init__(self, selector: selectors.BaseSelector | None = None) -> None: + super().__init__(selector) self._now = 0.0 def time(self) -> float: @@ -95,6 +96,28 @@ class _ManualClockLoop(asyncio.SelectorEventLoop): self._now += seconds +class _AutojumpSelector(selectors.DefaultSelector): + def __init__(self, advance: Callable[[float], None]) -> None: + super().__init__() + self._advance = advance + + def select(self, timeout: float | None = None) -> list[tuple[selectors.SelectorKey, int]]: + ready: Final = super().select(0) + if ready or timeout == 0: + return ready + if timeout is None: + return super().select(None) + self._advance(timeout) + return [] + + +class _AutojumpClockLoop(_ManualClockLoop): + """A manual-clock loop that jumps to the next timer only once no callback or I/O event is left to run""" + + def __init__(self) -> None: + super().__init__(_AutojumpSelector(self.advance)) + + class _FakeExceptionGroup(Exception): """Duck-typed stand-in for an anyio/builtin ExceptionGroup. @@ -1890,16 +1913,26 @@ async def test_transport_parsing_failure_is_preserved(transport: MCPTransport, f ) -@pytest.mark.asyncio -async def test_sse_read_failure_is_preserved() -> None: - client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2) - with pytest.raises(httpx2.ReadError, match="secret-read-error"): - await asyncio.wait_for( - client._execute_session_operation( - _diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), lambda session: session.list_tools() - ), - timeout=3, - ) +def test_sse_read_failure_is_preserved() -> None: + loop: Final = _AutojumpClockLoop() + + async def run() -> None: + client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2) + with pytest.raises(httpx2.ReadError, match="secret-read-error"): + await asyncio.wait_for( + client._execute_session_operation( + _diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), + lambda session: session.list_tools(), + ), + timeout=3, + ) + + try: + loop.run_until_complete(run()) + finally: + loop.run_until_complete(loop.shutdown_asyncgens()) + loop.run_until_complete(loop.shutdown_default_executor()) + loop.close() @pytest.mark.asyncio @@ -2638,9 +2671,12 @@ def test_public_mcp_import_preserves_incompatible_sdk_error() -> None: @pytest.mark.parametrize("grouped", (False, True)) @pytest.mark.parametrize("raise_on_error", (False, True)) @pytest.mark.parametrize("termination", ("ok", "failure", "hang")) -async def test_outer_deadline_delivers_session_termination(termination: str, grouped: bool, raise_on_error: bool) -> None: +async def test_outer_deadline_delivers_session_termination( + termination: str, grouped: bool, raise_on_error: bool +) -> None: deleted: Final = asyncio.Event() started: Final = asyncio.Event() + caller_deadline: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future() async def respond(request: httpx2.Request) -> httpx2.Response: await anyio.lowlevel.checkpoint() @@ -2672,26 +2708,30 @@ async def test_outer_deadline_delivers_session_termination(termination: str, gro if payload.method == "tools/list": return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": []}}) started.set() + caller_deadline.result().deadline = anyio.current_time() await anyio.sleep_forever() raise AssertionError("cancelled request resumed") client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", timeout=30) - async def invoke(): - with anyio.fail_after(0.2): - pending: Final = client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error) + async def invoke() -> None: + with anyio.fail_after(None) as deadline: + caller_deadline.set_result(deadline) + pending: Final = client.call_tool( + CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error + ) if grouped: await asyncio.gather(pending) else: await pending - before: Final = anyio.current_time() - with pytest.raises(TimeoutError): - await invoke() + with anyio.fail_after(20): + with pytest.raises(TimeoutError): + await invoke() assert started.is_set() assert deleted.is_set(), "Cancellation must deliver DELETE before returning to the caller" - assert anyio.current_time() - before < 6.5 + assert anyio.current_time() - caller_deadline.result().deadline < 6.5 assert await client.list_tools(raise_on_error=True) == []