diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 91da7204435..d2851cf4fbf 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -107,6 +107,9 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Per-server_id generation, bumped on invalidation so a fetch that started before the server +# definition changed cannot repopulate the cache with the stale reply. +_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( tags=["mcp"], @@ -143,6 +146,7 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: @@ -2377,6 +2381,14 @@ async def fetch_upstream_oauth_protected_resource( cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: return cached[1] + generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) + + def store(payload: dict | None, ttl_seconds: int) -> None: + if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation: + return + stored_at: Final = time.time() + _OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload) + _prune_oauth_metadata_cache(stored_at) host_base: Final = f"{upstream.scheme}://{upstream.netloc}" candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"] @@ -2418,12 +2430,7 @@ async def fetch_upstream_oauth_protected_resource( ) continue if isinstance(payload, dict): - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_CACHE_TTL_SECONDS, - payload, - ) - _prune_oauth_metadata_cache(now) + store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS) return payload if len(network_errors) == len(candidates): @@ -2432,12 +2439,7 @@ async def fetch_upstream_oauth_protected_resource( # Negative-result caching: when no candidate yielded a usable payload, # remember that for a shorter TTL so we don't re-fetch on every # subsequent discovery request (and so the per-key lock can be pruned). - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, - None, - ) - _prune_oauth_metadata_cache(now) + store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS) return None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 11159a1c72b..e5ec69bf818 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4607,39 +4607,23 @@ class MCPServerManager: self._listed_tools_by_server_id.pop(server_id, None) invalidate_oauth_metadata_cache(server_id) - def _discovers_per_caller(self, server: MCPServer) -> bool: - return ( - server.requires_per_user_auth - or self._references_per_user_env_var(server) - or server.delegate_auth_to_upstream - or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) - or self._signs_caller_identity_upstream(server) - ) - - @staticmethod - def _signs_caller_identity_upstream(server: MCPServer) -> bool: - """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may - tailor its catalog to the caller even though the server itself is configured as shared.""" - from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server - get_mcp_jwt_signer, - ) - - if get_mcp_jwt_signer() is None: - return False - return not any(k.lower() == "authorization" for k in (server.static_headers or {})) - def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: """Key the listed-tool cache by every request input that can change the upstream catalog. Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or - exchanged as the OBO subject), and the server-specific auth header all reach - upstream, so two callers differing in any of them may be shown different tools. Shared - servers with none of those stay on the shared (``None``) slot. OpenAPI servers list from - the process-wide registry. + exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT + MCPJWTSigner mints for tools/list all reach upstream, so two callers differing in any of + them may be shown different tools. Shared servers with none of those stay on the shared + (``None``) slot. OpenAPI servers list from the process-wide registry. """ if server.spec_path or caller is None: return None auth: Final = caller.user_api_key_auth + signed_caller: Final = ( + f"{auth.user_id}:{auth.api_key}" + if auth is not None and self._signs_caller_identity_upstream(server) + else None + ) forwarded: Final = dict(self._forwarded_header_values(server, caller.raw_headers)) or None header_env: Final = self._build_stdio_env(server, caller.raw_headers) stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env @@ -4649,7 +4633,19 @@ class MCPServerManager: else None ) _, digest = self._discovery_key(server, auth, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer) - return digest + if signed_caller is None: + return digest + return hashlib.sha256(f"{digest}:{signed_caller}".encode()).hexdigest() + + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return not any(k.lower() == "authorization" for k in (server.static_headers or {})) @staticmethod def _forwarded_header_values( @@ -4688,11 +4684,18 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = self._discovers_per_caller(server) + per_user: Final = ( + server.requires_per_user_auth + or self._references_per_user_env_var(server) + or server.delegate_auth_to_upstream + or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + ) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None + (user_api_key_auth.user_id, user_api_key_auth.api_key) + if per_user and user_api_key_auth is not None + else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 66c30a62bde..67cdbbff5a4 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -1,3 +1,4 @@ +import itertools import uuid from pathlib import Path from typing import Final @@ -7,10 +8,13 @@ import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( McpCaller, + McpPeer, call_tool, delete_mcp, forget_mcp, + listed_tools, mcp_peer, + openapi_peer, register_mcp, tool_calls, tool_names, @@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga scenario.cleanups.callback(forget_mcp, gateway, winner) +def _openapi_server_lists_and_calls_only_its_own_tools( + gateway: Gateway, key: str, peer: McpPeer, identity: str +) -> None: + listed: Final = set(listed_tools(gateway, key, identity)) + assert listed == {"getpet", "createpet"}, (identity, listed) + peer.drain() + called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"}) + assert called.status_code == 200, called.text + assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity + + +def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None: + with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario: + stem: Final = "pet" + uuid.uuid4().hex[:8] + servers: Final = tuple( + (peer, alias, register_mcp(scenario, peer, alias)) + for peer, alias in ((short, stem), (long, stem + "store")) + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]}) + for peer, _, identity in servers: + _openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity) + aggregate: Final = McpCaller(gateway, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + assert sorted(aggregate.tools) == sorted( + f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet")) + ), aggregate.tools + assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream" + + +def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with openapi_peer() as peer: + config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + name: Final = "pet store " + uuid.uuid4().hex[:8] + config["mcp_servers"] = {name: peer.registration()} + path: Final = tmp_path / "openapi-space.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + _openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity) + aggregate: Final = McpCaller(candidate, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + prefix: Final = name.replace(" ", "_") + assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools + + def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 4c46c706054..3a023f6d38e 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -1,17 +1,23 @@ import json import queue +import threading import uuid -from urllib.parse import parse_qs, urlsplit -from typing import Final, Literal +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field from pathlib import Path +from typing import Final, Literal +from urllib.parse import parse_qs, urlsplit import pytest - -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy -from integration._support.wire import Reply, Request, wire_server +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +_Upstream = Callable[[Request], Reply] @pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint") @@ -104,6 +110,105 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat assert updated.status_code == 202, updated.text +@dataclass(frozen=True, slots=True) +class _Hold: + armed: threading.Event = field(default_factory=threading.Event) + released: threading.Event = field(default_factory=threading.Event) + + +def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream: + def issuer() -> str: + return origin() + ("/idp-after" if moved.is_set() else "/idp-before") + + def respond(request: Request) -> Reply: + if "oauth-authorization-server" in request.target or "openid-configuration" in request.target: + current: Final = issuer() + return Reply( + body=json.dumps( + { + "issuer": current, + "authorization_endpoint": current + "/authorize", + "token_endpoint": current + "/token", + } + ).encode() + ) + if request.target.startswith("/.well-known/oauth-protected-resource"): + body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode() + if hold is not None and hold.armed.is_set(): + assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released" + return Reply(body=body) + return Reply(status=404, body=b'{"error":"unexpected"}') + + return respond + + +def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str: + return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough") + + +def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]: + def observed() -> tuple[Request, ...]: + seen.extend(wire.drain()) + return tuple(seen) + + return observed + + +def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool: + return any( + "oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests + ) + + +def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp") + assert response.status_code == 200, response.text + return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) + + +def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + moved.set() + wire.drain() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), ( + "the save must send protected-resource discovery back to the upstream" + ) + + +def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: + moved: Final = threading.Event() + hold: Final = _Hold() + with ( + wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire, + gateway.scenario() as scenario, + ThreadPoolExecutor(max_workers=1) as pool, + ): + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + seen: Final[list[Request]] = [] + observed: Final = _wire_requests(wire, seen) + eventually(observed, _registration_discovery_settled, seconds=10) + settled: Final = len(seen) + hold.armed.set() + stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias) + eventually(observed, lambda requests: len(requests) > settled, seconds=10) + assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:] + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + hold.released.set() + assert stale.result(timeout=30) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + + @pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server") @pytest.mark.parametrize("transition", ("revoke", "expire")) def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0a92b548d62..7a2af8dfcca 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12646,3 +12646,46 @@ async def test_update_server_drops_cached_upstream_oauth_metadata(): finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) + + +@pytest.mark.asyncio +async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="stale-write-server", name="stale_write", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["old-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await in_flight == {"authorization_servers": ["old-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f8709572835..a502ae2a3e1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -114,7 +114,6 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} - @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -222,8 +221,6 @@ def _reload_mcp_manager_module(): return reloaded - - @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1397,7 +1394,9 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( + self, config_only_mcp_manager_factory, monkeypatch, caplog + ): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4683,7 +4682,9 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): + async def test_openapi_health_loads_spec_without_mcp_handshake( + self, respx_mock, monkeypatch, auth_type, is_byok, scheme + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4733,14 +4734,28 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), + ( + httpx.Response(401, text="secret response content"), + "unhealthy", + "OpenAPI specification request failed (HTTP 401)", + ), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), - (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), + ( + httpx.ConnectError("secret network details"), + "unhealthy", + "OpenAPI specification could not be loaded (ConnectError)", + ), + ( + httpx.Response(200, text="secret invalid JSON body"), + "unhealthy", + "OpenAPI specification could not be loaded (JSONDecodeError)", + ), ], ) - async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): + async def test_openapi_health_reports_safe_failures( + self, respx_mock, monkeypatch, failure, expected_status, expected_error + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4898,7 +4913,10 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"]) async def test_health_check_server_oauth2_reports_reachability( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + oauth2_flow: Literal["authorization_code", "client_credentials"] | None, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() @@ -4927,14 +4945,28 @@ class TestMCPServerManager: assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type", [ - MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token, - MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, - ]) + @pytest.mark.parametrize( + "auth_type", + [ + MCPAuth.bearer_token, + MCPAuth.api_key, + MCPAuth.basic, + MCPAuth.authorization, + MCPAuth.token, + MCPAuth.oauth2_token_exchange, + MCPAuth.oauth2_id_jag, + MCPAuth.true_passthrough, + MCPAuth.oauth_delegate, + ], + ) @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse]) @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503]) async def test_health_check_without_credentials_accepts_any_http_response( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse], + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + auth_type: MCPAuthType, + transport: Literal[MCPTransport.http, MCPTransport.sse], response_code: int, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") @@ -4965,6 +4997,7 @@ class TestMCPServerManager: self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + class UnreadBody(httpx.AsyncByteStream): def __init__(self) -> None: self.read = False @@ -4979,17 +5012,28 @@ class TestMCPServerManager: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events", + server_id="streaming-health", + name="streaming-health", + transport=MCPTransport.sse, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/events", ) manager.registry[server.server_id] = server bodies: Final = (UnreadBody(), UnreadBody()) - route: Final = respx_mock.get(server.url).mock(side_effect=[ - httpx.Response(response_code, stream=body, headers={ - "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/", - "Location": "http://127.0.0.1/private", - }) for body in bodies - ]) + route: Final = respx_mock.get(server.url).mock( + side_effect=[ + httpx.Response( + response_code, + stream=body, + headers={ + "Content-Type": "text/event-stream", + "Set-Cookie": "health=secret; Path=/", + "Location": "http://127.0.0.1/private", + }, + ) + for body in bodies + ] + ) first: Final = await manager.health_check_server(server.server_id) second: Final = await manager.health_check_server(server.server_id) @@ -5000,19 +5044,28 @@ class TestMCPServerManager: assert all("cookie" not in call.request.headers for call in route.calls) @pytest.mark.asyncio - @pytest.mark.parametrize(("transport", "url"), [ - (MCPTransport.stdio, "https://mcp.example.test"), - (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"), - (MCPTransport.http, "ftp://mcp.example.test"), - (MCPTransport.http, "https://user:secret@mcp.example.test"), - (MCPTransport.http, "https://mcp.example.test:bad/mcp"), - ]) + @pytest.mark.parametrize( + ("transport", "url"), + [ + (MCPTransport.stdio, "https://mcp.example.test"), + (MCPTransport.http, None), + (MCPTransport.http, ""), + (MCPTransport.http, "not-a-url"), + (MCPTransport.http, "ftp://mcp.example.test"), + (MCPTransport.http, "https://user:secret@mcp.example.test"), + (MCPTransport.http, "https://mcp.example.test:bad/mcp"), + ], + ) async def test_health_reachability_rejects_unprobeable_urls_without_requests( self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None ) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url, + server_id="unprobeable", + name="unprobeable", + transport=transport, + auth_type=MCPAuth.oauth2, + url=url, ) manager.registry[server.server_id] = server @@ -5023,19 +5076,26 @@ class TestMCPServerManager: assert not respx_mock.calls @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [ - httpx.ConnectError("TLS/connection failure with secret details"), - httpx.ReadTimeout("secret timeout details"), - httpx.RemoteProtocolError("secret malformed response"), - ]) + @pytest.mark.parametrize( + "failure", + [ + httpx.ConnectError("TLS/connection failure with secret details"), + httpx.ReadTimeout("secret timeout details"), + httpx.RemoteProtocolError("secret malformed response"), + ], + ) async def test_health_reachability_reports_no_response_without_secret_details( self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="failed-health", name="failed-health", transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret", + server_id="failed-health", + name="failed-health", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + is_byok=True, + url="https://mcp.example.test/secret?token=secret", ) manager.registry[server.server_id] = server route: Final = respx_mock.get(server.url).mock(side_effect=failure) @@ -5051,8 +5111,11 @@ class TestMCPServerManager: monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="bad-tls", name="bad-tls", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test", + server_id="bad-tls", + name="bad-tls", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test", ) manager.registry[server.server_id] = server @@ -5070,8 +5133,11 @@ class TestMCPServerManager: monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1) manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="slow-health", name="slow-health", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow", + server_id="slow-health", + name="slow-health", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/slow", ) manager.registry[server.server_id] = server started: Final = asyncio.Event() @@ -5124,8 +5190,11 @@ class TestMCPServerManager: server_ids: Final = [f"health-{index}" for index in range(server_count)] manager.registry = { server_id: MCPServer( - server_id=server_id, name=server_id, transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}", + server_id=server_id, + name=server_id, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=f"https://health.example.test/{server_id}", ) for server_id in server_ids } @@ -5452,8 +5521,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5538,8 +5614,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers @@ -7257,7 +7340,9 @@ class TestMCPServerManager: is_byok=True, ) manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) - manager._fetch_tools_with_timeout = AsyncMock(return_value=[MCPTool(name="turn", description="t", inputSchema={})]) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="t", inputSchema={})] + ) with patch("litellm.proxy.proxy_server.prisma_client", None): await manager._get_tools_from_server( @@ -7301,20 +7386,13 @@ class TestMCPServerManager: else: assert for_bob is None - def test_per_caller_slot_identity_is_the_token_not_the_api_key(self): - """Two callers sharing user_id split on the hashed token the admission validator stamps; - the raw api_key never enters the identity, so a caller carrying only that token lands on - the same slot.""" - from litellm.proxy._types import hash_token - + def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self): + """Two keys sharing a user_id get different signed JWTs, so they split; the same key + presented again lands on its own slot.""" manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - alice = ListedToolsCaller( - user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha") - ) - bob = ListedToolsCaller( - user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-beta") - ) + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-beta")) with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", @@ -7325,10 +7403,8 @@ class TestMCPServerManager: ) assert manager.get_listed_tool(server, "srv-turn", bob) is None - same_token = ListedToolsCaller( - user_api_key_auth=UserAPIKeyAuth(user_id="same-user", token=hash_token("sk-alpha")) - ) - listed = manager.get_listed_tool(server, "srv-turn", same_token) + same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + listed = manager.get_listed_tool(server, "srv-turn", same_key) assert listed is not None and listed.description == "slot a" @@ -10148,15 +10224,11 @@ class TestGetPublicMCPServers: if registered_in == "both" else server ) - manager.config_mcp_servers = ( - {server.server_id: config_server} if registered_in in ("config", "both") else {} - ) + manager.config_mcp_servers = {server.server_id: config_server} if registered_in in ("config", "both") else {} manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} original_server: Final = server.model_dump() original_config_server: Final = config_server.model_dump() - expected_public: Final = registered_in != "neither" and ( - public_ids == [server.server_id] or implicitly_public - ) + expected_public: Final = registered_in != "neither" and (public_ids == [server.server_id] or implicitly_public) with ( patch("litellm.public_mcp_servers", public_ids), @@ -10164,9 +10236,7 @@ class TestGetPublicMCPServers: ): public_servers: Final = manager.get_public_mcp_servers() assert manager.is_mcp_server_public(server.server_id) is expected_public - assert [item.server_id for item in public_servers] == ( - [server.server_id] if expected_public else [] - ) + assert [item.server_id for item in public_servers] == ([server.server_id] if expected_public else []) assert manager.is_mcp_server_public("server-alias") is False assert manager.is_mcp_server_public("missing-server") is False @@ -11062,7 +11132,9 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): + async def call_tool( + self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False + ): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -13335,7 +13407,9 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( + self, config_only_mcp_manager_factory, aliasing_entry_first: bool + ): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -13351,7 +13425,9 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( + self, config_only_mcp_manager_factory + ): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -13439,7 +13515,9 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( + self, config_only_mcp_manager_factory, caplog + ): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -13610,7 +13688,9 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another( + self, config_only_mcp_manager_factory, caplog + ): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13982,7 +14062,8 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner( + _mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -14053,7 +14134,9 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques @pytest.mark.asyncio @pytest.mark.parametrize("transport", ["http", "stdio"]) -async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ctx, transport: Literal["http", "stdio"]) -> None: +async def test_debug_reports_legacy_signing_and_non_http_transport( + _mcp_request_ctx, transport: Literal["http", "stdio"] +) -> None: from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request @@ -14097,12 +14180,16 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -14122,13 +14209,18 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + server_id="repeated-stale", + name="stale", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -14148,13 +14240,20 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + server_id="resolved-replacement", + name="replacement", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy( + update={ + "url": "https://new.example.com/mcp", + "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + } ) - replacement: Final = original.model_copy(update={ - "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -14162,8 +14261,11 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", name="publication", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="stale-publication", + name="publication", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -14179,9 +14281,13 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + server_id="expiring-session", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -14282,7 +14388,9 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + assert ( + result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + ) assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -14294,8 +14402,11 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, + server_id="cancelled-cache", + name="cancelled-cache", + transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", + auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -14424,7 +14535,9 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + return MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http + ) @pytest.mark.asyncio @@ -14592,7 +14705,9 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) +@pytest.mark.parametrize( + "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) +) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -14611,44 +14726,6 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) - from litellm.proxy._types import hash_token - - same_user_other_token: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None - ) - same_token_no_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None - ) - with_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None - ) - assert same_user_other_token != with_key - assert same_token_no_key == with_key - - -@pytest.mark.parametrize( - ("signer", "static_headers", "shared"), - [ - pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), - pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), - pytest.param(None, None, True, id="no-signer-stays-shared"), - ], -) -def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None: - manager: Final = MCPServerManager() - server: Final = _discovery_server().model_copy(update={"static_headers": static_headers}) - alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice") - bob: Final = UserAPIKeyAuth(user_id="bob", token="hashed-bob") - - with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam - "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", - return_value=signer, - ): - for_alice: Final = manager._discovery_key(server, alice, None, None, None, None) - for_bob: Final = manager._discovery_key(server, bob, None, None, None, None) - - assert (for_alice == for_bob) is shared - def _register_local_tool(name: str, description: str) -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry @@ -14947,26 +15024,45 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,credential", [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ]) + @pytest.mark.parametrize( + "auth_type,credential", + [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ], + ) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, credential: str | None, dispatch: str, + self, + tmp_path: Path, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, + credential: str | None, + dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}})) + spec_path.write_text( + json.dumps( + { + "openapi": "3.0.0", + "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}, + } + ) + ) server: Final = MCPServer( - server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, + server_id="dispatch-auth", + name="dispatch-auth", + url="https://upstream.example", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14989,14 +15085,21 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", - transport=transport, auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", + server_id="incomplete-obo", + name="incomplete-obo", + url="https://upstream.example/mcp", + transport=transport, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", + client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", + authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header="Bearer override", subject_token=subject, + server, + mcp_auth_header="Bearer override", + subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -15009,8 +15112,11 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-static", + name="empty-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -15018,16 +15124,22 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,headers", [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ]) + @pytest.mark.parametrize( + "auth_type,headers", + [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ], + ) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", name="header-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="header-static", + name="header-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -15036,29 +15148,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="openapi-empty", + name="openapi-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, - user_api_key_auth=None, forwarded_headers=None, + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=None, + forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,slot,value", [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ]) + @pytest.mark.parametrize( + "auth_type,slot,value", + [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ], + ) async def test_raw_static_credentials_are_forwarded_unchanged( - self, auth_type: MCPAuthType, slot: str, value: str, + self, + auth_type: MCPAuthType, + slot: str, + value: str, ) -> None: - server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) + server = MCPServer( + server_id="raw-key", + name="raw-key", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, + ) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -15072,17 +15203,24 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, respx_mock: MockRouter, value: str, source: str, + self, + respx_mock: MockRouter, + value: str, + source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.authorization, + server_id="raw-empty", + name="raw-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, + server, + mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -15090,9 +15228,15 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, - token_exchange_endpoint="https://idp.example/token") + server = MCPServer( + server_id="obo-byok", + name="obo-byok", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + is_byok=True, + token_exchange_endpoint="https://idp.example/token", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -15100,41 +15244,66 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) + server = MCPServer( + server_id="override", + name="override", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=configured, + ) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) + server = MCPServer( + server_id="empty-header", + name="empty-header", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=token, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", authentication_token="key") + server = MCPServer( + server_id="custom", + name="custom", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", + authentication_token="key", + ) client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize("static_headers,accepted", [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ]) + @pytest.mark.parametrize( + "static_headers,accepted", + [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ], + ) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, + server_id="static-slot", + name="static-slot", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -15146,21 +15315,36 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize("static,forwarded,caller", [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ]) + @pytest.mark.parametrize( + "static,forwarded,caller", + [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ], + ) async def test_openapi_static_credentials_remain_supported( - self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None + self, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], + forwarded: dict[str, str] | None, + caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, _request_extra_headers, create_tool_function, + _request_auth_header, + _request_extra_headers, + create_tool_function, ) + tool: Final = create_tool_function( - "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, + "/echo", + "get", + {}, + "https://upstream.example", + headers=static, + auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -15194,8 +15378,13 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key) + server = MCPServer( + server_id="cancel", + name="cancel", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + ) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -15204,8 +15393,14 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") + server = MCPServer( + server_id="blank-static", + name="blank-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=" ", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -15213,8 +15408,13 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic) + server = MCPServer( + server_id="bad-basic", + name="bad-basic", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -15223,34 +15423,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None) + server = MCPServer( + server_id="basic-scheme", + name="basic-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,default_slot", [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ]) + @pytest.mark.parametrize( + "auth_type,value,default_slot", + [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", name="alternate", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", + server_id="alternate", + name="alternate", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, + server, + mcp_auth_header=value if source == "caller" else None, + extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -15259,8 +15473,12 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", + server_id="both-empty", + name="both-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -15273,12 +15491,17 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, + server_id="caller-auth", + name="caller-auth", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=headers if source == "caller" else None, + server, + mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -15287,14 +15510,29 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize("value", [ - "", " ", "Bearer", "Basic", "token", "ApiKey", - "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", - ]) + @pytest.mark.parametrize( + "value", + [ + "", + " ", + "Bearer", + "Basic", + "token", + "ApiKey", + "Bearer Bearer", + "ApiKey ApiKey", + "token token", + "bEaReR BEARER", + "aPiKeY\tAPIKEY", + ], + ) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, + server_id="caller-empty", + name="caller-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -15305,8 +15543,11 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, + server_id="basic-pair", + name="basic-pair", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -15319,8 +15560,12 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, + server_id="basic-valid", + name="basic-valid", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -15329,17 +15574,27 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value", [ - (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), - ]) + @pytest.mark.parametrize( + "auth_type,value", + [ + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.bearer_token, "Bearer "), + (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), + (MCPAuth.token, "token "), + (MCPAuth.token, "TOKEN"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-scheme", + name="empty-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -15347,17 +15602,24 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,expected", [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ]) + @pytest.mark.parametrize( + "auth_type,value,expected", + [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ], + ) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", name="real-token", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, + server_id="real-token", + name="real-token", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -15396,16 +15658,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = {"observer": MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", - )} + manager.registry = { + "observer": MCPServer( + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", + ) + } manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for(manager.call_tool( - server_name="observer", name="execute", arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), - ), timeout=5) + result = await asyncio.wait_for( + manager.call_tool( + server_name="observer", + name="execute", + arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context( + {"metadata": {"guardrails": ["observe"] if selected else []}} + ), + ), + timeout=5, + ) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -15434,11 +15711,21 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + upstream = MCPServer( + server_id="explicit-empty", + name="explicit_empty", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + legacy_server.set_auth_context( + UserAPIKeyAuth(user_id="unrelated"), + raw_headers={"authorization": "unrelated-credential"}, + client_ip="192.0.2.99", + ) with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -15446,7 +15733,9 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await MCPServerManager()._create_mcp_client( + upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None + ) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -15469,16 +15758,28 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-a", + server_name="alpha", + alias="shared", + url="https://a.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-b", + server_name="beta", + alias="Shared", + url="https://b.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-c", + server_name="gamma", + alias="lonely", + url="https://c.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] @@ -15523,7 +15824,9 @@ class TestSharedIdentifierPrefixWarning: @pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): manager = config_only_mcp_manager_factory() - await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}}) + await manager.load_servers_from_config( + {"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}} + ) server = next(iter(manager.config_mcp_servers.values())) client = await manager._create_mcp_client(server) assert server.protocol_version == revision @@ -15535,9 +15838,13 @@ async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_m def test_runtime_protocol_metadata_preserves_explicit_precedence( revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None ) -> None: - server: Final = MCPServer.model_validate({ - "server_id": "preview", "name": "preview", "transport": "http", - "mcp_info": {"protocol_version": revision}, - **({"protocol_version": explicit} if explicit is not None else {}), - }) + server: Final = MCPServer.model_validate( + { + "server_id": "preview", + "name": "preview", + "transport": "http", + "mcp_info": {"protocol_version": revision}, + **({"protocol_version": explicit} if explicit is not None else {}), + } + ) assert server.protocol_version == (explicit if explicit is not None else revision)