mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
test(mcp): fire MCP client test deadlines on conditions instead of wall-clock time (#44166)
* test(mcp): fire MCP client test deadlines on conditions instead of wall-clock time Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): bound MCP client test failure paths independently of the triggered deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): measure outer-deadline cleanup budget from cancellation, not setup Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
fe9b6fd603
commit
615ed7900f
1 changed files with 61 additions and 21 deletions
|
|
@ -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) == []
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue