mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(e2e): build wire headers and mint keys inline in the mcp test bodies
This commit is contained in:
parent
9ff988d7ef
commit
d32974c282
3 changed files with 51 additions and 48 deletions
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}; "
|
||||
|
|
|
|||
|
|
@ -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 <key>`; 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
|
||||
(`<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}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue