From b273b979160d9b0db852940445e5a7f59f719996 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 02:47:31 +0200 Subject: [PATCH] style: apply ruff format to MCP client test tree file --- .../test_mcp_client.py | 120 +++++++++++++----- 1 file changed, 86 insertions(+), 34 deletions(-) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 3860f0f44bb..d16de968c25 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -597,7 +597,9 @@ class TestExecuteSessionOperationSurfacesTransportError: raise _FakeExceptionGroup("transport", [_FakeExceptionGroup("reader", failures)]) self._make_session(session_class, initialize) - expected: Final = close_error if failure_phase == "early" else connect_error if failure_phase == "mixed" else cancelled + expected: Final = ( + close_error if failure_phase == "early" else connect_error if failure_phase == "mixed" else cancelled + ) with pytest.raises(type(expected)) as caught: await client._execute_session_operation(self._make_transport(close_transport), AsyncMock(), http_client) assert caught.value is expected @@ -617,7 +619,6 @@ class TestExecuteSessionOperationSurfacesTransportError: result = await client._execute_session_operation(transport_ctx, _op) assert result == "done" - @pytest.mark.asyncio @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_session_entry_failure_still_closes_transport(self, session_class): @@ -1235,11 +1236,7 @@ def test_mcp_extra_matches_proxy_extra_and_supports_streamable_http(): sdk2_names: Final = frozenset(("mcp", "httpx2", "pydantic")) mcp_extra: Final = {Requirement(req).name: req for req in extras["mcp"]} - assert mcp_extra == { - name: req - for req in extras["proxy"] - if (name := Requirement(req).name) in sdk2_names - } + assert mcp_extra == {name: req for req in extras["proxy"] if (name := Requirement(req).name) in sdk2_names} specifier: Final = Requirement(mcp_extra["mcp"]).specifier assert not specifier.contains("1.28.1") @@ -1639,7 +1636,9 @@ async def test_http_response_handler_preserves_success_and_http_errors(status_co @pytest.mark.asyncio async def test_http_status_check_allows_auth_refresh_before_rejecting() -> None: - from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ClientCredentialsBearerAuth + from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( + ClientCredentialsBearerAuth, + ) seen = [] @@ -1892,14 +1891,20 @@ async def test_sse_read_failure_is_preserved() -> None: @pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"]) @pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio]) @pytest.mark.parametrize("mode", ["ok", "closed", "silent"]) -async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None: +async def test_transport_completion_and_normal_messages( + transport: MCPTransport, mode: str, protocol_version: str +) -> None: from mcp import ClientSession from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message logging_callback: Final = AsyncMock() read_timeout: Final = 0.2 if mode == "silent" else 30 client: Final = MCPClient( - server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version + server_url="https://example.com/sse", + transport_type=transport, + timeout=read_timeout, + logging_callback=logging_callback, + protocol_version=protocol_version, ) async def operation(session: ClientSession) -> CallToolResult: @@ -2461,8 +2466,15 @@ def test_client_import_before_proxy_credentials_succeeds_in_fresh_process(): import subprocess result = subprocess.run( - [sys.executable, "-c", "import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)"], - capture_output=True, text=True, timeout=60, check=False, + [ + sys.executable, + "-c", + "import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)", + ], + capture_output=True, + text=True, + timeout=60, + check=False, ) assert result.returncode == 0, result.stderr assert result.stdout.strip() == "MCPServerManager" @@ -2495,8 +2507,10 @@ async def test_request_auth_preview_uses_the_same_effective_headers_as_egress() from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth client: Final = MCPClient( - server_url="https://upstream.example/mcp", auth_type=MCPAuth.bearer_token, - resolved_auth=StaticHeaderAuth("Bearer resolved"), extra_headers={"X-Trace": "trace"}, + server_url="https://upstream.example/mcp", + auth_type=MCPAuth.bearer_token, + resolved_auth=StaticHeaderAuth("Bearer resolved"), + extra_headers={"X-Trace": "trace"}, ) request: Final = await client.prepare_request_auth() assert request.method == "POST" @@ -2520,18 +2534,29 @@ async def test_expired_session_preserves_sdk_error_and_next_operation_reinitiali return httpx2.Response(202) requests.append((payload["method"], request.headers.get("mcp-session-id"))) if payload["method"] == "initialize": - return httpx2.Response(200, headers={"mcp-session-id": f"session-{len(requests)}"}, json={ - "jsonrpc": "2.0", "id": payload["id"], "result": { - "protocolVersion": "2025-06-18", "capabilities": {}, - "serverInfo": {"name": "expiry-test", "version": "1"}, + return httpx2.Response( + 200, + headers={"mcp-session-id": f"session-{len(requests)}"}, + json={ + "jsonrpc": "2.0", + "id": payload["id"], + "result": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "serverInfo": {"name": "expiry-test", "version": "1"}, + }, }, - }) + ) if len(requests) == 2: if rpc_error: - return httpx2.Response(404, json={ - "jsonrpc": "2.0", "id": payload["id"], - "error": {"code": METHOD_NOT_FOUND, "message": "Tool catalog unavailable"}, - }) + return httpx2.Response( + 404, + json={ + "jsonrpc": "2.0", + "id": payload["id"], + "error": {"code": METHOD_NOT_FOUND, "message": "Tool catalog unavailable"}, + }, + ) return httpx2.Response(404) return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload["id"], "result": {"tools": []}}) @@ -2547,7 +2572,12 @@ async def test_expired_session_preserves_sdk_error_and_next_operation_reinitiali streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools() ) assert result.tools == [] - assert requests == [("initialize", None), ("tools/list", "session-1"), ("initialize", None), ("tools/list", "session-3")] + assert requests == [ + ("initialize", None), + ("tools/list", "session-1"), + ("initialize", None), + ("tools/list", "session-3"), + ] @pytest.mark.asyncio @@ -2604,7 +2634,9 @@ 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() @@ -2644,7 +2676,9 @@ 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(): - pending: Final = asyncio.ensure_future(client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)) + 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() @@ -2858,7 +2892,9 @@ async def test_cancellation_delivers_termination_over_tcp( listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0) port: Final = listener.sockets[0].getsockname()[1] client: Final = MCPClient( - server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30 + server_url=f"http://127.0.0.1:{port}/mcp", + protocol_version=protocol_version, + timeout=2 if cancel_mode == "read_timeout" else 30, ) async def calls(): @@ -2942,16 +2978,32 @@ async def test_configured_upstream_revision_is_offered_and_checked(revision, acc assert payload.params["protocolVersion"] == offered assert ("sampling" in payload.params["capabilities"]) == callbacks assert ("elicitation" in payload.params["capabilities"]) == callbacks - return httpx2.Response(200, json={ - "jsonrpc": "2.0", "id": payload.id, - "result": {"protocolVersion": offered if accepted else "unsupported", - "capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}}, - }) + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": { + "protocolVersion": offered if accepted else "unsupported", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "upstream", "version": "1"}, + }, + }, + ) assert accepted, "No operation may execute after failed version negotiation" - return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}}) + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}, + }, + ) client = _MockTransportClient( - respond, server_url="https://example.com/mcp", protocol_version=revision, + respond, + server_url="https://example.com/mcp", + protocol_version=revision, sampling_callback=AsyncMock() if callbacks else None, elicitation_callback=AsyncMock() if callbacks else None, )