diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 2d597abf3b8..8e3d25552e0 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1038,6 +1038,20 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock: request.headers = {"Content-Type": "application/json"} request.client = MagicMock() request.client.host = "127.0.0.1" + request.scope = { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.3"}, + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/v1/batches", + "raw_path": b"/v1/batches", + "query_string": b"", + "root_path": "", + "headers": [(b"content-type", b"application/json"), (b"host", b"localhost")], + "client": ("127.0.0.1", 54321), + "server": ("localhost", 8000), + } request.body = AsyncMock(return_value=json.dumps(body).encode()) return request diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 1a56227b008..3860f0f44bb 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -2644,12 +2644,21 @@ async def test_outer_deadline_delivers_session_termination(termination: str, gro 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) - if grouped: - await asyncio.gather(pending) - else: + pending: Final = asyncio.ensure_future(client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)) + try: + with anyio.fail_after(2.0): + await started.wait() + with anyio.fail_after(0.2): + if grouped: + await asyncio.gather(pending) + else: + await pending + finally: + pending.cancel() + try: await pending + except BaseException: + pass before: Final = anyio.current_time() with pytest.raises(TimeoutError):