mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
style: apply ruff format to MCP client test tree file
This commit is contained in:
parent
1625735ba4
commit
b273b97916
1 changed files with 86 additions and 34 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue