diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index c6bb720eecc..46ae99a3d64 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -4,10 +4,14 @@ Management routes (/v1/mcp/server CRUD) go through the shared Gateway transport. The MCP protocol itself (initialize, tools/list, tools/call over streamable HTTP) goes through the official mcp SDK, the same client library production MCP hosts use, aimed at the gateway's per-server URL namespace -{PROXY}/{alias}/mcp. The gateway accepts the LiteLLM virtual key as either +{PROXY}/{alias}/mcp. + +Every protocol method takes the request headers as a plain dict, built inside +the test body, so the exact wire format is visible where it is asserted. The +gateway accepts the LiteLLM virtual key as either `x-litellm-api-key: Bearer sk-...` or `Authorization: Bearer sk-...` (both -Bearer-prefixed on the MCP routes, matching the docs); `McpAuth` models the -two header styles so each test states which one it drives. +Bearer-prefixed on the MCP routes, matching the docs); `McpHeaderName` names +those two documented header styles for the tests' parametrized matrices. """ from __future__ import annotations @@ -34,17 +38,6 @@ McpHeaderName = Literal["x-litellm-api-key", "Authorization"] ToolArguments = Mapping[str, str | float] -@dataclass(frozen=True, slots=True) -class McpAuth: - """One of the two documented ways to present a LiteLLM key to the MCP gateway.""" - - header_name: McpHeaderName - key: str - - def headers(self) -> dict[str, str]: - return {self.header_name: f"Bearer {self.key}"} - - @dataclass(frozen=True, slots=True) class McpToolNames: names: tuple[str, ...] @@ -155,20 +148,20 @@ class McpClient: response_type=NoBody, ) - def list_tools_once(self, alias: str, auth: McpAuth) -> ListToolsOutcome: + def list_tools_once(self, alias: str, headers: dict[str, str]) -> ListToolsOutcome: try: - return McpToolNames(names=asyncio.run(_list_tool_names(_mcp_url(alias), auth.headers()))) + return McpToolNames(names=asyncio.run(_list_tool_names(_mcp_url(alias), headers))) except Exception as exc: # noqa: BLE001 - the SDK raises ExceptionGroup-wrapped transport errors; modelled as a value return McpDenied(status_code=_http_status(exc), message=str(exc)) - def poll_tool_names(self, alias: str, auth: McpAuth) -> tuple[str, ...]: + def poll_tool_names(self, alias: str, headers: dict[str, str]) -> tuple[str, ...]: """tools/list to the shared deadline: a just-created server record propagates to the gateway asynchronously and a just-created key can lag the auth cache, so the first attempts may 401 or list nothing.""" deadline = time.monotonic() + self.gateway.poll_timeout outcome: ListToolsOutcome = McpDenied(status_code=None, message="never attempted") while time.monotonic() < deadline: - outcome = self.list_tools_once(alias, auth) + outcome = self.list_tools_once(alias, headers) match outcome: case McpToolNames(names=names) if names: return names @@ -178,13 +171,13 @@ class McpClient: f"MCP tools for {alias!r} never listed within {self.gateway.poll_timeout}s; last outcome: {outcome}" ) - def call_tool(self, alias: str, auth: McpAuth, tool: str, arguments: ToolArguments) -> McpToolText: + def call_tool(self, alias: str, headers: dict[str, str], tool: str, arguments: ToolArguments) -> McpToolText: """One tools/call over its own fresh MCP session, so concurrent callers behave like independent clients.""" - return asyncio.run(_call_tool(_mcp_url(alias), auth.headers(), tool, arguments)) + return asyncio.run(_call_tool(_mcp_url(alias), headers, tool, arguments)) - def stub_stats(self, alias: str, auth: McpAuth, stats_tool: str, marker: str) -> StubToolStats: - outcome = self.call_tool(alias, auth, stats_tool, {"marker": marker}) + def stub_stats(self, alias: str, headers: dict[str, str], stats_tool: str, marker: str) -> StubToolStats: + outcome = self.call_tool(alias, headers, stats_tool, {"marker": marker}) return StubToolStats.model_validate_json(outcome.text) diff --git a/tests/e2e/mcp/test_mcp_concurrency_e2e.py b/tests/e2e/mcp/test_mcp_concurrency_e2e.py index 9c1f87ded18..ccff22d5414 100644 --- a/tests/e2e/mcp/test_mcp_concurrency_e2e.py +++ b/tests/e2e/mcp/test_mcp_concurrency_e2e.py @@ -23,8 +23,8 @@ import pytest from e2e_config import MCP_STUB_URL, unique_marker from lifecycle import ResourceManager -from mcp_client import McpAuth, McpClient, McpToolText -from models import McpServerCreateBody +from mcp_client import McpClient, McpToolText +from models import KeyGenerateBody, McpServerCreateBody pytestmark = pytest.mark.e2e @@ -35,7 +35,7 @@ class TestMcpServerMaxConcurrency: @pytest.mark.covers("mcp.call_tool.api_key.caps_concurrency") def test_max_concurrent_requests_caps_in_flight_upstream_calls( - self, client: McpClient, resources: ResourceManager, scoped_key: str + self, client: McpClient, resources: ResourceManager ) -> None: """Two servers against the same stub: one capped at 2 concurrent requests, one uncapped control. The tools/list polls are the settle @@ -46,7 +46,10 @@ class TestMcpServerMaxConcurrency: max_concurrent = 2 burst_size = 6 slow_call_seconds = 2.0 - auth = McpAuth(header_name="x-litellm-api-key", key=scoped_key) + + key = client.gateway.generate_key(KeyGenerateBody()) + resources.defer(lambda: client.gateway.delete_key(key)) + headers = {"x-litellm-api-key": f"Bearer {key}"} capped_alias = f"e2emcpcap{unique_marker()}" capped = client.create_server( @@ -66,14 +69,14 @@ class TestMcpServerMaxConcurrency: assert client.server_info(capped.server_id).max_concurrent_requests == max_concurrent assert client.server_info(control.server_id).max_concurrent_requests is None - _ = client.poll_tool_names(capped_alias, auth) - _ = client.poll_tool_names(control_alias, auth) + _ = client.poll_tool_names(capped_alias, headers) + _ = client.poll_tool_names(control_alias, headers) def burst_of_slow_calls(alias: str, text: str, marker: str) -> list[McpToolText]: arguments = {"text": text, "marker": marker, "sleep_seconds": slow_call_seconds} with ThreadPoolExecutor(max_workers=burst_size) as pool: futures = [ - pool.submit(client.call_tool, alias, auth, f"{alias}-slow_echo", arguments) + pool.submit(client.call_tool, alias, headers, f"{alias}-slow_echo", arguments) for _ in range(burst_size) ] return [future.result() for future in futures] @@ -85,7 +88,7 @@ class TestMcpServerMaxConcurrency: f"queued calls must all still succeed under the cap: {capped_results}" ) - capped_stats = client.stub_stats(capped_alias, auth, f"{capped_alias}-stats", capped_marker) + capped_stats = client.stub_stats(capped_alias, headers, f"{capped_alias}-stats", capped_marker) assert capped_stats.completed == burst_size assert capped_stats.max_in_flight == max_concurrent, ( f"stub saw {capped_stats.max_in_flight} overlapping calls; the cap of {max_concurrent} " @@ -96,7 +99,7 @@ class TestMcpServerMaxConcurrency: control_results = burst_of_slow_calls(control_alias, "control", control_marker) assert all(result.is_error is False and result.text == "control" for result in control_results) - control_stats = client.stub_stats(control_alias, auth, f"{control_alias}-stats", control_marker) + control_stats = client.stub_stats(control_alias, headers, f"{control_alias}-stats", control_marker) assert control_stats.completed == burst_size assert control_stats.max_in_flight == burst_size, ( f"uncapped control saw {control_stats.max_in_flight} overlapping calls, expected all {burst_size}; " diff --git a/tests/e2e/mcp/test_mcp_tool_access_e2e.py b/tests/e2e/mcp/test_mcp_tool_access_e2e.py index e15c9500155..9a1090e8bdf 100644 --- a/tests/e2e/mcp/test_mcp_tool_access_e2e.py +++ b/tests/e2e/mcp/test_mcp_tool_access_e2e.py @@ -6,14 +6,16 @@ mcp.call_tool.bearer.succeeds (the `Authorization` header), and mcp.list_tools.{api_key,bearer}.rejects_unknown_key (the ingress gate), all against the deterministic mcp-stub compose service (tests/e2e/mcp/stub/). -Each test walks the full lifecycle: register the server over the management -API and defer its deletion, assert the recorded state round-trips -(GET /v1/mcp/server/{id} echoes what was configured), then drive the real MCP -protocol through the gateway the way a production MCP host does (initialize, -tools/list, tools/call over streamable HTTP at {PROXY}/{alias}/mcp) and -assert the enforced behavior. Both auth headers carry `Bearer `; that is -the documented contract on the MCP routes (a bare key in `x-litellm-api-key` -is accepted on LLM routes but 401s here). +Each test body spells out every step a human QA run would take, in order: +create the server over the management API, read the record back, mint a +virtual key over /key/generate, build the exact wire header +(`
: Bearer sk-...`, the documented contract on the MCP routes; a bare +key in `x-litellm-api-key` is accepted on LLM routes but 401s here), then +drive the real MCP protocol through the gateway the way a production MCP host +does (initialize, tools/list, tools/call over streamable HTTP at +{PROXY}/{alias}/mcp) and assert the enforced behavior. Teardown is the only +thing delegated (resources.defer), because it must run even when an +assertion fails. The two header styles are one behavior matrix, not two behaviors: the same test body runs once per header via parametrize, so the specs cannot drift @@ -29,8 +31,8 @@ import pytest from e2e_config import MCP_STUB_URL, unique_marker from lifecycle import ResourceManager -from mcp_client import McpAuth, McpClient, McpDenied, McpHeaderName -from models import McpServerCreateBody +from mcp_client import McpClient, McpDenied, McpHeaderName +from models import KeyGenerateBody, McpServerCreateBody pytestmark = pytest.mark.e2e @@ -76,7 +78,7 @@ class TestMcpToolAccess: @pytest.mark.parametrize("header_name", AUTH_HEADER_MATRIX) def test_list_and_call_tools( - self, header_name: McpHeaderName, client: McpClient, resources: ResourceManager, scoped_key: str + self, header_name: McpHeaderName, client: McpClient, resources: ResourceManager ) -> None: alias = f"e2emcp{unique_marker()}" created = client.create_server(McpServerCreateBody(alias=alias, url=MCP_STUB_URL, allow_all_keys=True)) @@ -87,19 +89,22 @@ class TestMcpToolAccess: assert stored.url == MCP_STUB_URL assert stored.allow_all_keys is True - auth = McpAuth(header_name=header_name, key=scoped_key) - names = client.poll_tool_names(alias, auth) + key = client.gateway.generate_key(KeyGenerateBody()) + resources.defer(lambda: client.gateway.delete_key(key)) + + headers = {header_name: f"Bearer {key}"} + names = client.poll_tool_names(alias, headers) expected = tuple(sorted(f"{alias}-{tool}" for tool in STUB_TOOLS)) assert names == expected, f"gateway listed {names}, expected exactly {expected}" payload = f"e2e-{unique_marker()}" - result = client.call_tool(alias, auth, f"{alias}-echo", {"text": payload}) + result = client.call_tool(alias, headers, f"{alias}-echo", {"text": payload}) assert result.is_error is False, f"echo call errored: {result.text[:300]}" assert result.text == payload @pytest.mark.parametrize("header_name", REJECTION_MATRIX) def test_unrecognized_key_is_turned_away( - self, header_name: McpHeaderName, client: McpClient, resources: ResourceManager, scoped_key: str + self, header_name: McpHeaderName, client: McpClient, resources: ResourceManager ) -> None: """Settle the server with a real key first, then present an unknown key over the same header: the refusal must be an explicit 401 from @@ -109,8 +114,10 @@ class TestMcpToolAccess: created = client.create_server(McpServerCreateBody(alias=alias, url=MCP_STUB_URL, allow_all_keys=True)) resources.defer(lambda: client.delete_server(created.server_id)) - _ = client.poll_tool_names(alias, McpAuth(header_name=header_name, key=scoped_key)) + key = client.gateway.generate_key(KeyGenerateBody()) + resources.defer(lambda: client.gateway.delete_key(key)) + _ = client.poll_tool_names(alias, {header_name: f"Bearer {key}"}) - denied = client.list_tools_once(alias, McpAuth(header_name=header_name, key="sk-not-a-real-key")) + denied = client.list_tools_once(alias, {header_name: "Bearer sk-not-a-real-key"}) assert isinstance(denied, McpDenied), f"unrecognized key was served tools: {denied}" assert denied.status_code == 401, f"expected 401 for an unrecognized key, got {denied}"