style: apply ruff format to MCP client test tree file

This commit is contained in:
jibanez-staticduo 2026-10-01 02:47:31 +02:00
parent 1625735ba4
commit b273b97916
No known key found for this signature in database

View file

@ -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,
)