diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index d2986a3cd82..9f28da19292 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -220,6 +220,12 @@ class MCPRequestHandler: # when EVERY target is auth_type=oauth2 with delegate_auth_to_upstream # set; fails closed otherwise. validated_user_api_key_auth = UserAPIKeyAuth() + elif MCPRequestHandler._target_servers_are_true_passthrough( + path=request_route, + mcp_servers=mcp_servers, + client_ip=IPAddressUtils.get_mcp_client_ip(request), + ): + validated_user_api_key_auth = UserAPIKeyAuth() elif oauth2_headers: # Authorization on a non-delegated server: the bearer must be a real # LiteLLM credential, so a failed validation is a genuine 401/403 and @@ -399,6 +405,33 @@ class MCPRequestHandler: return False return True + @staticmethod + def _target_servers_are_true_passthrough( + path: str, mcp_servers: Optional[list[str]], client_ip: Optional[str] + ) -> bool: + """ + True only when EVERY MCP server the request targets is ``auth_type == true_passthrough``. + Fails closed when any target does not opt in or cannot be resolved. + + Used by :meth:`process_mcp_request` to skip LiteLLM admission auth entirely: the gateway is a + transparent proxy and the caller's ``Authorization`` is an upstream token, never a LiteLLM key. + Mirrors :meth:`_target_servers_delegate_auth_to_upstream`; a mixed-target request keeps normal auth. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + + target_names = MCPRequestHandler._resolve_target_server_names(path=path, mcp_servers_header=mcp_servers) + if not target_names: + return False + + for name in target_names: + server = global_mcp_server_manager.get_mcp_server_by_name(name, client_ip=client_ip) + if server is None or server.auth_type != MCPAuth.true_passthrough: + return False + return True + @staticmethod def _resolve_target_server_names(path: str, mcp_servers_header: Optional[List[str]]) -> List[str]: """ diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7c88c903324..fd978e5bfc9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -78,6 +78,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_ ) from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( AuthorizationCodeConfig, + PassthroughConfig, ServerSpec, TokenExchangeConfig, ) @@ -213,6 +214,13 @@ def _should_strip_caller_authorization( pass-through cold-start case (RFC 9728) the bearer in ``Authorization`` is the upstream OAuth token and must be forwarded, so we keep it. + - **oauth_delegate servers**: admission always runs and there is no + anonymous path, so the caller's separate ``Authorization`` is + forwarded only when a distinct ``x-litellm-api-key`` carried + admission. Without that header the ``Authorization`` *was* the + admission credential — a virtual key, an IdP JWT, or an SSO / OIDC / + session token whose ``api_key`` is ``None`` — and must never reach + the upstream, so it is stripped regardless of the ``api_key`` value. """ if mcp_server.auth_type == MCPAuth.oauth2_token_exchange: # OBO: the inbound Authorization is the subject token. It is exchanged at the IdP and only the @@ -226,11 +234,13 @@ def _should_strip_caller_authorization( # upstream — it would override another user's stored credential. Delegate and # pass-through return None from to_server_spec and keep forwarding the bearer. return True - if not mcp_server.is_oauth_passthrough: + if not (mcp_server.is_oauth_passthrough or mcp_server.is_oauth_delegate): return False normalized_raw_headers = {str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str)} has_explicit_litellm_admission_header = normalized_raw_headers.get("x-litellm-api-key") is not None + if mcp_server.is_oauth_delegate: + return not has_explicit_litellm_admission_header admission_consumed_authorization_as_litellm_key = ( user_api_key_auth is not None and bool(getattr(user_api_key_auth, "api_key", None)) @@ -323,6 +333,41 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header +def _client_forwarded_authorization_headers( + mcp_server: MCPServer, + oauth2_headers: Optional[dict[str, str]], + raw_headers: Optional[dict[str, str]], + user_api_key_auth: Optional[UserAPIKeyAuth], +) -> Optional[dict[str, str]]: + """Egress headers for the client-forwarded-token modes (``true_passthrough`` / ``oauth_delegate``). + + Forwards the caller's ``Authorization`` to the upstream, stripped when + ``_should_strip_caller_authorization`` says it was consumed as the LiteLLM admission key. Shared by + ``_call_regular_mcp_tool`` and ``server.py``'s ``_prepare_mcp_server_headers`` so the two egress + paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split. + """ + extra_headers = oauth2_headers.copy() if oauth2_headers else None + if extra_headers and _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ): + return _without_authorization(extra_headers) + return extra_headers + + +def _take_forwarded_authorization( + headers: Optional[dict[str, str]], +) -> tuple[Optional[str], Optional[dict[str, str]]]: + """Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the + remaining headers, so the passthrough resolver arm is the single Authorization source rather than + the header also riding in ``extra_headers`` (which the resolved auth would then defer to).""" + if not headers: + return None, headers + value = next((v for k, v in headers.items() if k.lower() == "authorization"), None) + return value, _without_authorization(headers) + + def _extract_upstream_auth_failure( exc: BaseException, ) -> Optional[tuple[int, Optional[str]]]: @@ -1623,15 +1668,18 @@ class MCPServerManager: delegate_server_ids = [ server.server_id for server in self.get_registry().values() - if getattr(server, "auth_type", None) == MCPAuth.oauth2 - and getattr(server, "delegate_auth_to_upstream", False) is True - # M2M servers must not be exposed anonymously: an - # unauthenticated caller would get LiteLLM to proxy tool - # calls using its stored client_credentials. Resolve the flow - # rather than reading has_client_credentials so an unstamped - # M2M-shape row (null column, verbatim-read as non-M2M) still - # fails closed here, matching the anonymous-delegate auth gate. - and MCPServerManager.effective_oauth2_flow(server) != "client_credentials" + if ( + getattr(server, "auth_type", None) == MCPAuth.oauth2 + and getattr(server, "delegate_auth_to_upstream", False) is True + # M2M servers must not be exposed anonymously: an + # unauthenticated caller would get LiteLLM to proxy tool + # calls using its stored client_credentials. Resolve the flow + # rather than reading has_client_credentials so an unstamped + # M2M-shape row (null column, verbatim-read as non-M2M) still + # fails closed here, matching the anonymous-delegate auth gate. + and MCPServerManager.effective_oauth2_flow(server) != "client_credentials" + ) + or getattr(server, "auth_type", None) == MCPAuth.true_passthrough ] combined_servers.update(delegate_server_ids) @@ -2229,16 +2277,17 @@ class MCPServerManager: spec = None if transport == MCPTransport.stdio else to_server_spec(server) provider = cred_provider or self._cred_provider # A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path - # so it wins - except for the per-user modes the v2 resolver owns (authorization_code's - # stored token and token_exchange's RFC 8693 minted token). A caller must not be able to - # substitute another user's stored credential, nor silently disable the OBO exchange and - # forward an arbitrary bearer upstream, so we keep the v2 spec and ignore the override for - # both; the REST tools preview supplies its not-yet-persisted token through the resolver - # (cred_provider), never this path. + # so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's + # stored token, token_exchange's RFC 8693 minted token, and the passthrough modes' + # forwarded caller token). A caller must not be able to substitute another user's stored + # credential, nor silently disable the OBO exchange and forward an arbitrary bearer + # upstream, so we keep the v2 spec and ignore the override for these; the REST tools + # preview supplies its not-yet-persisted token through the resolver (cred_provider), + # never this path. if ( spec is not None and mcp_auth_header - and not isinstance(spec.config, (AuthorizationCodeConfig, TokenExchangeConfig)) + and not isinstance(spec.config, (AuthorizationCodeConfig, PassthroughConfig, TokenExchangeConfig)) ): spec = None auth_value = ( @@ -2305,11 +2354,14 @@ class MCPServerManager: server_url = server.url or "" if spec is not None: + inbound_token = subject_token + if isinstance(spec.config, PassthroughConfig): + inbound_token, extra_headers = _take_forwarded_authorization(extra_headers) resolved_auth, extra_headers = await self._resolve_v2_auth( server=server, spec=spec, provider=provider, - subject_token=subject_token, + subject_token=inbound_token, user_api_key_auth=user_api_key_auth, extra_headers=extra_headers, ) @@ -3730,6 +3782,13 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, ): extra_headers = _without_authorization(extra_headers) + elif mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate: + extra_headers = _client_forwarded_authorization_headers( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) if mcp_server.extra_headers and raw_headers: if extra_headers is None: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index 05896bfff74..e87e8081ced 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( AuthorizationCodeConfig, CredError, NoneConfig, + PassthroughConfig, ServerSpec, SharedKey, Subject, @@ -62,9 +63,10 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live modes: ``none``, the static-header family (``api_key`` plus the Authorization schemes, - all shared-key), ``oauth2`` per-user tokens (``authorization_code``), and - ``oauth2_token_exchange`` (OBO); client_credentials (M2M), delegated/passthrough - oauth2, and SigV4 return None and stay on v1. + all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2_token_exchange`` + (OBO), and the client-forwarded token modes ``true_passthrough`` / ``oauth_delegate`` + (``PassthroughConfig``); client_credentials (M2M), delegated/passthrough oauth2, and SigV4 + return None and stay on v1. """ if server.is_byok: return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type) @@ -94,6 +96,8 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: ) # client_credentials (M2M) and delegate/passthrough oauth2 stay on v1 return None + case MCPAuth.true_passthrough | MCPAuth.oauth_delegate: + return ServerSpec(server_id=server.server_id, resource=resource, config=PassthroughConfig()) case MCPAuth.oauth2_token_exchange: return _token_exchange_spec(server, resource) case MCPAuth.aws_sigv4: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index c82ce1037d6..ecfd471190c 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -7,10 +7,11 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly at runtime instead of returning `None`. -`none` and `api_key` (shared-key source) are live, as is `authorization_code`, which reads the -user's token from the injected `OAuthTokenStore`, and `token_exchange`, which swaps the caller's -inbound token through the injected `TokenExchanger`. The remaining arms are `not_implemented` stubs -that each land in a follow-up PR with their seam. Pure v2: no imports from v1. +`none`, `api_key` (shared-key source), and `passthrough` (forwards the caller's own inbound token) +are live, as is `authorization_code`, which reads the user's token from the injected +`OAuthTokenStore`, and `token_exchange`, which swaps the caller's inbound token through the +injected `TokenExchanger`. The remaining arms are `not_implemented` stubs that each land in a +follow-up PR with their seam. Pure v2: no imports from v1. """ from __future__ import annotations @@ -97,7 +98,7 @@ class UpstreamCredentialProvider: case ApiKeyConfig() as config: return self._api_key(config) case PassthroughConfig(): - return _not_implemented(AuthSpecKind.passthrough) + return self._passthrough(subject) case ClientCredentialsConfig(): return _not_implemented(AuthSpecKind.client_credentials) case TokenExchangeConfig() as config: @@ -118,6 +119,18 @@ class UpstreamCredentialProvider: """ return await self._authz_token(subject, server) is not None + def _passthrough(self, subject: Subject) -> Result[httpx.Auth, CredError]: + """Forward the caller's own upstream credential verbatim; the gateway mints nothing. + + The inbound token is the caller's already-disambiguated ``Authorization`` (never the LiteLLM + admission credential; the edge adapter drops that before building the ``Subject``). When it is + absent the request is sent unauthenticated so the upstream's own 401 surfaces, rather than the + gateway challenging on the upstream's behalf. + """ + if subject.inbound_token is None: + return Ok(NoOpAuth()) + return Ok(StaticHeaderAuth(subject.inbound_token.get_secret_value(), header_name="Authorization")) + def _api_key(self, config: ApiKeyConfig) -> Result[httpx.Auth, CredError]: match config.key_source: case SharedKey() as source: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e3812522ded..96fd014fbe1 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -331,6 +331,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, + _client_forwarded_authorization_headers, _should_strip_caller_authorization, _without_authorization, global_mcp_server_manager, @@ -1561,6 +1562,13 @@ if MCP_AVAILABLE: user_api_key_auth=user_api_key_auth, ): extra_headers = _without_authorization(extra_headers) + elif server.is_true_passthrough or server.is_oauth_delegate: + extra_headers = _client_forwarded_authorization_headers( + mcp_server=server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) if server.extra_headers and raw_headers: if extra_headers is None: diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 9c564a3c7a6..d273e8ec4db 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -38,6 +38,8 @@ class MCPAuth(str, enum.Enum): aws_sigv4 = "aws_sigv4" token = "token" oauth2_token_exchange = "oauth2_token_exchange" + true_passthrough = "true_passthrough" + oauth_delegate = "oauth_delegate" # RFC 8693 default subject_token_type. A NULL column / omitted config key means @@ -60,6 +62,8 @@ MCPAuthType = Optional[ MCPAuth.aws_sigv4, MCPAuth.token, MCPAuth.oauth2_token_exchange, + MCPAuth.true_passthrough, + MCPAuth.oauth_delegate, ] ] diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 522a6f09165..f102ab5b7b9 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -152,6 +152,18 @@ class MCPServer(BaseModel): """True if this is an OAuth2 server that relies on per-user tokens (no client_credentials).""" return self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials + @property + def is_true_passthrough(self) -> bool: + """True for the transparent-proxy mode: LiteLLM performs no admission auth and forwards the + client's ``Authorization`` to the upstream unchanged.""" + return self.auth_type == MCPAuth.true_passthrough + + @property + def is_oauth_delegate(self) -> bool: + """True for the delegated-upstream-OAuth mode: LiteLLM still admits the caller (API key / SSO / + JWT) but forwards the caller's separate upstream ``Authorization`` unchanged, minting nothing.""" + return self.auth_type == MCPAuth.oauth_delegate + @property def requires_per_user_auth(self) -> bool: """ @@ -167,6 +179,9 @@ class MCPServer(BaseModel): if self.needs_user_oauth_token: return True + if self.is_true_passthrough or self.is_oauth_delegate: + return True + # PAT passthrough: auth_type is none but extra_headers includes auth headers if self.auth_type == MCPAuth.none and self.extra_headers: auth_header_names = {"authorization", "x-api-key", "api-key", "apikey"} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 3db0f8540f9..c984ccb783e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1,14 +1,12 @@ import json import os import sys -from unittest.mock import AsyncMock, MagicMock, call as mock_call, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.testclient import TestClient -sys.path.insert( - 0, os.path.abspath("../../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path from starlette.datastructures import Headers @@ -78,20 +76,14 @@ class TestMCPRequestHandler: ) # Mock the helper methods instead of database calls - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key_servers: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team_servers: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key_servers: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team_servers: # Set up return values mock_key_servers.return_value = key_servers mock_team_servers.return_value = team_servers # Call the method - result = await MCPRequestHandler.get_allowed_mcp_servers( - user_api_key_auth=mock_user_auth - ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=mock_user_auth) # Assert the result (order-independent comparison) assert sorted(result) == sorted(expected_result) @@ -148,20 +140,14 @@ class TestMCPRequestHandler: ) # Mock the helper functions - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key_servers: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team_servers: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key_servers: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team_servers: # Configure mocks to return the test data mock_key_servers.return_value = key_servers mock_team_servers.return_value = team_servers # Call the method - result = await MCPRequestHandler.get_allowed_mcp_servers( - user_api_key_auth - ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) # Assert the result (order-independent comparison) assert sorted(result) == sorted(expected_servers) @@ -186,9 +172,7 @@ class TestMCPRequestHandler: ): """The require_key_mcp_access_defined general setting flips an empty key from inheriting its team's MCP servers (default) to inheriting none.""" - auth = UserAPIKeyAuth( - api_key="test-key", user_id="test-user", team_id="test-team" - ) + auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", team_id="test-team") with ( patch.object( MCPRequestHandler, @@ -272,23 +256,22 @@ class TestMCPRequestHandler: async def test_no_mcp_servers_sentinel_returns_empty(self, team_servers): """A key scoped to the no-mcp-servers sentinel resolves to zero servers, overriding team inheritance and never leaking the sentinel marker.""" - user_api_key_auth = UserAPIKeyAuth( - api_key="test-key", user_id="test-user", team_id="test-team" - ) + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", team_id="test-team") key_object_permission = MagicMock() - key_object_permission.mcp_servers = [ - SpecialMCPServerNames.no_mcp_servers.value - ] + key_object_permission.mcp_servers = [SpecialMCPServerNames.no_mcp_servers.value] - with patch.object( - MCPRequestHandler, - "_get_key_object_permission", - return_value=key_object_permission, - ), patch.object( - MCPRequestHandler, - "_get_allowed_mcp_servers_for_team", - new_callable=AsyncMock, - return_value=team_servers, + with ( + patch.object( + MCPRequestHandler, + "_get_key_object_permission", + return_value=key_object_permission, + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_team", + new_callable=AsyncMock, + return_value=team_servers, + ), ): result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) @@ -309,9 +292,7 @@ class TestMCPRequestHandler: "_get_key_object_permission", return_value=key_object_permission, ): - result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( - user_api_key_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) assert result == [SpecialMCPServerNames.no_mcp_servers.value] @@ -320,9 +301,7 @@ class TestMCPRequestHandler: # Test case: None values in database mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.return_value = ( - None - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.return_value = None mock_prisma_client.db.litellm_teamtable.find_unique.return_value = None user_api_key_auth = UserAPIKeyAuth( @@ -337,9 +316,7 @@ class TestMCPRequestHandler: assert result == [] # Test case: Exception handling - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.side_effect = ( - Exception("DB Error") - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.side_effect = Exception("DB Error") with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) @@ -384,15 +361,9 @@ class TestMCPRequestHandler: access_group_ids=["grp-mcp"], ) with ( - patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key, - patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team, - patch.object( - MCPRequestHandler, "_get_key_access_group_mcp_server_extras" - ) as mock_grants, + patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key, + patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team, + patch.object(MCPRequestHandler, "_get_key_access_group_mcp_server_extras") as mock_grants, ): mock_key.return_value = key_servers mock_team.return_value = team_servers @@ -414,13 +385,9 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=[]), ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - auth - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth) assert result == [] # expand_permission_list must not be reached when there are no raw ids. mock_mgr.expand_permission_list.assert_not_called() @@ -433,14 +400,10 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=["alias-a", "srv-b"]), ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.expand_permission_list.return_value = ["srv-a", "srv-b"] - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - auth - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth) assert sorted(result) == ["srv-a", "srv-b"] mock_mgr.expand_permission_list.assert_called_once_with(["alias-a", "srv-b"]) @@ -451,9 +414,7 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", new=AsyncMock(side_effect=Exception("db down")), ): - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - auth - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth) assert result == [] @pytest.mark.parametrize( @@ -743,9 +704,7 @@ class TestMCPRequestHandler: # Verify MCP servers mcp_servers_header = extracted_headers.get(SpecialHeaders.mcp_servers.value) mcp_servers = None - if ( - mcp_servers_header is not None - ): # Changed from 'if mcp_servers_header:' to handle empty strings + if mcp_servers_header is not None: # Changed from 'if mcp_servers_header:' to handle empty strings try: # First try to parse as JSON array for backward compatibility try: @@ -754,16 +713,12 @@ class TestMCPRequestHandler: mcp_servers = None except (json.JSONDecodeError, TypeError, ValueError): # If JSON parsing fails, treat as comma-separated list - mcp_servers = [ - s.strip() for s in mcp_servers_header.split(",") if s.strip() - ] + mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()] except Exception: mcp_servers = None # If we got an empty string or parsing resulted in no servers, return empty list - if mcp_servers_header == "" or ( - mcp_servers is not None and len(mcp_servers) == 0 - ): + if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0): mcp_servers = [] assert mcp_servers == expected_result["mcp_servers"] @@ -833,9 +788,7 @@ class TestMCPOAuth2AuthFlow: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", new_callable=AsyncMock, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = oauth2_server ( @@ -851,10 +804,7 @@ class TestMCPOAuth2AuthFlow: # The upstream token is never validated as a LiteLLM key ... mock_auth.assert_not_called() # ... and is preserved for upstream forwarding. - assert ( - oauth2_headers.get("Authorization") - == "Bearer atlassian-oauth2-access-token-xyz" - ) + assert oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-access-token-xyz" async def test_explicit_litellm_key_with_oauth2_authorization(self): """ @@ -893,9 +843,7 @@ class TestMCPOAuth2AuthFlow: assert call_args.kwargs["api_key"] == "sk-litellm-valid-key" # OAuth2 headers should still contain the Authorization token - assert ( - oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-token" - ) + assert oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-token" async def test_litellm_key_in_authorization_backward_compat(self): """ @@ -997,9 +945,7 @@ class TestMCPOAuth2AuthFlow: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_proxy_exception, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = oauth2_server with pytest.raises(ProxyException) as exc_info: @@ -1068,9 +1014,7 @@ class TestMCPPublicRouteGuard: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): # Explicit unresolvable target — proves auth still fails even # when the registry has no info to fall back to. @@ -1101,9 +1045,7 @@ class TestMCPPublicRouteGuard: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = None with pytest.raises(HTTPException) as exc_info: @@ -1157,9 +1099,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp._is_mcp_passthrough_cold_start" ) as mock_cold_start, @@ -1199,9 +1139,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1229,9 +1167,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1263,18 +1199,14 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", return_value="203.0.113.10", ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = None with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(scope) assert exc_info.value.status_code == 401 - mock_mgr.get_mcp_server_by_name.assert_any_call( - "passthrough_server", client_ip="203.0.113.10" - ) + mock_mgr.get_mcp_server_by_name.assert_any_call("passthrough_server", client_ip="203.0.113.10") async def test_cold_start_propagates_non_401_http_error(self): from fastapi import HTTPException @@ -1294,9 +1226,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_forbidden, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1329,9 +1259,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_server_error, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1357,9 +1285,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1367,9 +1293,7 @@ class TestMCPPassthroughColdStartAdmission: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) - mock_mgr.get_mcp_server_by_name.assert_any_call( - "passthrough_server", client_ip="" - ) + mock_mgr.get_mcp_server_by_name.assert_any_call("passthrough_server", client_ip="") async def test_cold_start_allows_proxy_exception_401_for_path_target(self): from litellm.proxy._types import ProxyException @@ -1394,9 +1318,7 @@ class TestMCPPassthroughColdStartAdmission: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPPassthroughColdStartAdmission._make_passthrough_server() @@ -1404,9 +1326,7 @@ class TestMCPPassthroughColdStartAdmission: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) - mock_mgr.get_mcp_server_by_name.assert_any_call( - "passthrough_server", client_ip="" - ) + mock_mgr.get_mcp_server_by_name.assert_any_call("passthrough_server", client_ip="") @pytest.mark.asyncio @@ -1450,12 +1370,10 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.api_key) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPOAuth2FallbackTargetGating._make_server( + MCPAuth.api_key ) with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(scope) @@ -1483,9 +1401,7 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = None with pytest.raises(HTTPException) as exc_info: @@ -1523,12 +1439,10 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.oauth2) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPOAuth2FallbackTargetGating._make_server( + MCPAuth.oauth2 ) with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(scope) @@ -1562,15 +1476,11 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPOAuth2FallbackTargetGating._make_server( - auth_type=MCPAuth.none, - is_oauth_passthrough=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPOAuth2FallbackTargetGating._make_server( + auth_type=MCPAuth.none, + is_oauth_passthrough=True, ) auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) @@ -1598,9 +1508,7 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", return_value="203.0.113.10", ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = None with pytest.raises(HTTPException) as exc_info: @@ -1612,9 +1520,7 @@ class TestMCPOAuth2FallbackTargetGating: # resolve to ``None`` (hidden by client IP) so neither bypass # opens. Use ``assert_any_call`` to assert the IP-scoped lookup # happened without locking the count. - mock_mgr.get_mcp_server_by_name.assert_any_call( - "hidden_oauth2_server", client_ip="203.0.113.10" - ) + mock_mgr.get_mcp_server_by_name.assert_any_call("hidden_oauth2_server", client_ip="203.0.113.10") async def test_fallback_blocked_when_any_target_in_header_is_not_oauth2(self): """ @@ -1649,9 +1555,7 @@ class TestMCPOAuth2FallbackTargetGating: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup with pytest.raises(HTTPException) as exc_info: @@ -1732,17 +1636,10 @@ class TestMCPDelegateAuthToUpstream: delegate_auth_to_upstream=True, available_on_public_internet=True, ) - assert ( - manager._build_mcp_server_table(delegated).delegate_auth_to_upstream is True - ) + assert manager._build_mcp_server_table(delegated).delegate_auth_to_upstream is True - not_delegated = delegated.model_copy( - update={"delegate_auth_to_upstream": False} - ) - assert ( - manager._build_mcp_server_table(not_delegated).delegate_auth_to_upstream - is False - ) + not_delegated = delegated.model_copy(update={"delegate_auth_to_upstream": False}) + assert manager._build_mcp_server_table(not_delegated).delegate_auth_to_upstream is False def test_build_mcp_server_table_preserves_oauth_passthrough(self): """Registry → API list rows must expose oauth_passthrough for the UI. @@ -1773,9 +1670,7 @@ class TestMCPDelegateAuthToUpstream: assert row.delegate_auth_to_upstream is False not_passthrough = passthrough.model_copy(update={"oauth_passthrough": False}) - assert ( - manager._build_mcp_server_table(not_passthrough).oauth_passthrough is False - ) + assert manager._build_mcp_server_table(not_passthrough).oauth_passthrough is False async def test_delegate_skips_litellm_auth_with_no_authorization(self): """ @@ -1796,15 +1691,11 @@ class TestMCPDelegateAuthToUpstream: patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.oauth2, - delegate_auth_to_upstream=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, ) auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) @@ -1834,15 +1725,11 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", new_callable=AsyncMock, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.oauth2, - delegate_auth_to_upstream=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, ) ( auth_result, @@ -1880,15 +1767,11 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.oauth2, - delegate_auth_to_upstream=False, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=False, ) with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(scope) @@ -1919,15 +1802,11 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.api_key, - delegate_auth_to_upstream=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.api_key, + delegate_auth_to_upstream=True, ) with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(scope) @@ -1971,9 +1850,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup with pytest.raises(HTTPException) as exc_info: @@ -2003,9 +1880,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth_fails, ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = None with pytest.raises(HTTPException) as exc_info: @@ -2034,15 +1909,11 @@ class TestMCPDelegateAuthToUpstream: new_callable=AsyncMock, return_value=UserAPIKeyAuth(user_id="real-user"), ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.oauth2, - delegate_auth_to_upstream=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, ) auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) @@ -2074,15 +1945,11 @@ class TestMCPDelegateAuthToUpstream: new_callable=AsyncMock, return_value=UserAPIKeyAuth(user_id="real-user"), ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): - mock_mgr.get_mcp_server_by_name.return_value = ( - TestMCPDelegateAuthToUpstream._make_server( - auth_type=MCPAuth.oauth2, - delegate_auth_to_upstream=True, - ) + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, ) ( auth_result, @@ -2135,9 +2002,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = m2m_server # No delegate bypass → normal auth is attempted → 401 raised @@ -2188,9 +2053,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = legacy_m2m_server with pytest.raises(HTTPException) as exc_info: @@ -2235,9 +2098,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = pkce_server auth, *_rest = await MCPRequestHandler.process_mcp_request(scope) @@ -2278,9 +2139,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = internal_server auth, *_rest = await MCPRequestHandler.process_mcp_request(scope) @@ -2426,6 +2285,104 @@ class TestMCPDelegateAuthToUpstream: assert "public-server" in result assert "internal-server" in result + async def test_true_passthrough_skips_litellm_auth_anonymously(self): + """auth_type=true_passthrough performs no admission auth: the caller's Authorization is an + upstream token forwarded unchanged and user_api_key_auth is never called.""" + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/true_passthrough_server", + "headers": [(b"authorization", b"Bearer upstream-token")], + } + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.true_passthrough, + ) + ( + auth_result, + _, + _, + _, + oauth2_headers, + _, + ) = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert auth_result.api_key is None + assert oauth2_headers.get("Authorization") == "Bearer upstream-token" + mock_auth.assert_not_called() + + async def test_true_passthrough_mixed_targets_fail_closed(self): + """One true_passthrough target mixed with a non-passthrough target must NOT skip admission.""" + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [(b"x-mcp-servers", b"tp_server,plain_server")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + def mock_lookup(name, client_ip=None): + if name == "tp_server": + return TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.true_passthrough, + ) + return TestMCPDelegateAuthToUpstream._make_server(auth_type=MCPAuth.api_key) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + + async def test_get_allowed_servers_includes_true_passthrough(self): + """Anonymous callers can reach true_passthrough servers; admission is delegated upstream.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + tp_server = MCPServer( + server_id="tp-server", + name="tp_server", + transport="http", + auth_type=MCPAuth.true_passthrough, + available_on_public_internet=True, + ) + manager.registry = {tp_server.server_id: tp_server} + + with patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + result = await manager.get_allowed_mcp_servers(None) + + assert "tp-server" in result + def test_extract_target_server_names_matches_routing_parser(self): """ Regression: _extract_target_server_names_from_path must match the @@ -2464,13 +2421,12 @@ class TestMCPDelegateAuthToUpstream: ("/", []), ] for path_input, expected in cases: - assert ( - MCPRequestHandler._extract_target_server_names_from_path(path_input) - == expected - ), f"path={path_input!r} → expected {expected!r}" - assert ( - _get_mcp_servers_in_path(path_input) or [] - ) == expected, f"path={path_input!r} → routing expected {expected!r}" + assert MCPRequestHandler._extract_target_server_names_from_path(path_input) == expected, ( + f"path={path_input!r} → expected {expected!r}" + ) + assert (_get_mcp_servers_in_path(path_input) or []) == expected, ( + f"path={path_input!r} → routing expected {expected!r}" + ) async def test_delegate_does_not_bypass_on_extra_path_segment(self): """ @@ -2514,9 +2470,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name with pytest.raises(HTTPException) as exc_info: @@ -2579,9 +2533,7 @@ class TestMCPDelegateAuthToUpstream: "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_auth_raises, ) as mock_auth, - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_mgr, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, ): mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name # Bypass MUST NOT fire — path-derived target is the non-delegate @@ -2601,15 +2553,12 @@ class TestMCPDelegateAuthToUpstream: empty-list case, which fails closed). """ # Path matches /mcp/... — header is ignored. - assert MCPRequestHandler._resolve_target_server_names( - path="/mcp/foo", mcp_servers_header=["evil"] - ) == ["foo"] - assert MCPRequestHandler._resolve_target_server_names( - path="/mcp/foo,bar", mcp_servers_header=["evil"] - ) == ["foo", "bar"] - assert MCPRequestHandler._resolve_target_server_names( - path="/foo/mcp", mcp_servers_header=["evil"] - ) == ["foo"] + assert MCPRequestHandler._resolve_target_server_names(path="/mcp/foo", mcp_servers_header=["evil"]) == ["foo"] + assert MCPRequestHandler._resolve_target_server_names(path="/mcp/foo,bar", mcp_servers_header=["evil"]) == [ + "foo", + "bar", + ] + assert MCPRequestHandler._resolve_target_server_names(path="/foo/mcp", mcp_servers_header=["evil"]) == ["foo"] # Path does not match — header is trusted. assert MCPRequestHandler._resolve_target_server_names( path="/.well-known/oauth-authorization-server", @@ -2653,16 +2602,12 @@ class TestMCPCustomHeaderName: (None, "", "x-mcp-auth"), ], ) - def test_get_mcp_client_side_auth_header_name( - self, env_var, general_setting, expected_header_name - ): + def test_get_mcp_client_side_auth_header_name(self, env_var, general_setting, expected_header_name): """Test that custom header name configuration works correctly""" # Mock the secret manager and general settings with patch("litellm.secret_managers.main.get_secret_str") as mock_get_secret: - with patch( - "litellm.proxy.proxy_server.general_settings" - ) as mock_general_settings: + with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings: # Configure mocks mock_get_secret.return_value = env_var mock_general_settings.get.return_value = general_setting @@ -2685,9 +2630,7 @@ class TestMCPCustomHeaderName: if env_var is None: # When env var is None, general settings should be checked (twice if not None) expected_general_calls = 2 if general_setting is not None else 1 - assert ( - mock_general_settings.get.call_count == expected_general_calls - ) + assert mock_general_settings.get.call_count == expected_general_calls for call in mock_general_settings.get.call_args_list: assert call.args == ("mcp_client_side_auth_header_name",) else: @@ -2728,9 +2671,7 @@ class TestMCPCustomHeaderName: ), ], ) - def test_get_mcp_auth_header_from_headers_with_custom_name( - self, custom_header_name, headers, expected_auth_header - ): + def test_get_mcp_auth_header_from_headers_with_custom_name(self, custom_header_name, headers, expected_auth_header): """Test that MCP auth header extraction uses custom header name""" # Mock the header name method @@ -2749,9 +2690,7 @@ class TestMCPCustomHeaderName: extracted_headers = MCPRequestHandler._safe_get_headers_from_scope(scope) # Call the method - result = MCPRequestHandler._get_mcp_auth_header_from_headers( - extracted_headers - ) + result = MCPRequestHandler._get_mcp_auth_header_from_headers(extracted_headers) # Assert the result assert result == expected_auth_header @@ -2818,9 +2757,7 @@ class TestMCPCustomHeaderName: from starlette.datastructures import Headers # Test case 1: No server-specific headers - headers = Headers( - {"x-litellm-api-key": "test-key", "content-type": "application/json"} - ) + headers = Headers({"x-litellm-api-key": "test-key", "content-type": "application/json"}) result = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) assert result == {} @@ -2904,17 +2841,13 @@ class TestMCPCustomHeaderName: assert result == {"github_mcp": {"Authorization": "Bearer github-mcp-token"}} # Test case 8: Edge case - empty header value - headers = Headers( - {"x-litellm-api-key": "test-key", "x-mcp-github-authorization": ""} - ) + headers = Headers({"x-litellm-api-key": "test-key", "x-mcp-github-authorization": ""}) result = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) assert result == {"github": {"Authorization": ""}} # Test case 9: Edge case - very long header value long_token = "Bearer " + "x" * 1000 - headers = Headers( - {"x-litellm-api-key": "test-key", "x-mcp-github-authorization": long_token} - ) + headers = Headers({"x-litellm-api-key": "test-key", "x-mcp-github-authorization": long_token}) result = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) assert result == {"github": {"Authorization": long_token}} @@ -2980,9 +2913,7 @@ class TestMCPAccessGroupsE2E: # Assert the results assert auth_result.api_key == "test-api-key" assert mcp_auth_header is None - assert ( - mcp_servers is None - ) # x-mcp-access-groups is not parsed as mcp_servers + assert mcp_servers is None # x-mcp-access-groups is not parsed as mcp_servers assert mcp_server_auth_headers == {} # Verify the mock was called @@ -3097,9 +3028,7 @@ def test_mcp_path_based_server_segregation(monkeypatch): # Use TestClient to make a request to /mcp/zapier,group1/tools client = TestClient(app) - response = client.get( - "/mcp/zapier,group1/tools", headers={"x-litellm-api-key": "test"} - ) + response = client.get("/mcp/zapier,group1/tools", headers={"x-litellm-api-key": "test"}) assert response.status_code == 200 assert response.json() == {"status": "ok"} @@ -3177,15 +3106,11 @@ async def test_get_team_object_permission_with_already_loaded_permission(): mock_prisma, ): with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: - with patch( - "litellm.proxy.auth.auth_checks.get_object_permission" - ) as mock_get_perm: + with patch("litellm.proxy.auth.auth_checks.get_object_permission") as mock_get_perm: mock_get_team.return_value = mock_team_obj # Call the method - result = await MCPRequestHandler._get_team_object_permission( - mock_user_auth - ) + result = await MCPRequestHandler._get_team_object_permission(mock_user_auth) # Assert we got the object permission assert result == mock_object_permission @@ -3272,9 +3197,7 @@ async def test_get_team_object_permission_ui_session_team_skips_db_lookup(): mock_prisma = MagicMock() with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: - result = await MCPRequestHandler._get_team_object_permission( - mock_user_auth - ) + result = await MCPRequestHandler._get_team_object_permission(mock_user_auth) assert result is None mock_get_team.assert_not_called() @@ -3342,9 +3265,7 @@ async def test_get_allowed_tools_for_server_ui_session_team_keeps_key_restrictio detail={"error": "Team doesn't exist in db. Team=litellm-dashboard."}, ), ): - with patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=key_perm - ): + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm): result = await MCPRequestHandler.get_allowed_tools_for_server( server_id="server_1", user_api_key_auth=user_api_key_auth, @@ -3410,9 +3331,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): return_value=["group-server1", "group-server2"], ) as mock_get_access_group_servers, ): - result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( - mock_user_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(mock_user_auth) assert set(result) == { "direct-server1", @@ -3455,9 +3374,7 @@ async def test_get_allowed_mcp_servers_for_team_with_no_object_permission(): return_value=mock_team, ), ): - result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( - mock_user_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(mock_user_auth) assert result == [] @@ -3507,9 +3424,7 @@ async def test_get_allowed_mcp_servers_for_team_without_team_id_returns_empty(): ), ], ) -async def test_get_allowed_mcp_servers_for_key_guard_conditions( - user_api_key_auth, prisma_client_value, scenario -): +async def test_get_allowed_mcp_servers_for_key_guard_conditions(user_api_key_auth, prisma_client_value, scenario): """Ensure guard clauses return [] before hitting get_object_permission.""" with patch( @@ -3517,9 +3432,7 @@ async def test_get_allowed_mcp_servers_for_key_guard_conditions( new_callable=AsyncMock, ) as mock_get_perm: with patch("litellm.proxy.proxy_server.prisma_client", prisma_client_value): - result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( - user_api_key_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) assert result == [] mock_get_perm.assert_not_called() @@ -3546,9 +3459,7 @@ async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_non ): mock_get_perm.return_value = None - result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( - user_api_key_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) assert result == [] mock_get_perm.assert_awaited_once() @@ -3591,14 +3502,10 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: - with patch.object( - MCPRequestHandler, "_get_mcp_servers_from_access_groups" - ) as mock_access_groups: + with patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups") as mock_access_groups: mock_access_groups.return_value = ["group-server"] - result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( - user_api_key_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) assert set(result) == {"direct-server", "group-server"} mock_get_perm.assert_not_called() @@ -3619,21 +3526,13 @@ class TestAgentMCPPermissions: team_id="test-team", agent_id="agent-123", ) - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" - ) as mock_agent: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = ["server_1"] - result = await MCPRequestHandler.get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth - ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=user_api_key_auth) assert sorted(result) == ["server_1"] mock_agent.assert_called_once_with(user_api_key_auth) @@ -3644,21 +3543,13 @@ class TestAgentMCPPermissions: user_id="test-user", agent_id="agent-456", ) - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" - ) as mock_agent: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = [] # no agent-level restriction - result = await MCPRequestHandler.get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth - ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=user_api_key_auth) assert sorted(result) == ["server_1", "server_2"] mock_agent.assert_called_once_with(user_api_key_auth) @@ -3669,21 +3560,13 @@ class TestAgentMCPPermissions: user_id="test-user", agent_id="agent-789", ) - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_key" - ) as mock_key: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_team" - ) as mock_team: - with patch.object( - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" - ) as mock_agent: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: + with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = ["server_2", "server_3"] - result = await MCPRequestHandler.get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth - ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=user_api_key_auth) assert sorted(result) == ["server_2"] async def test_get_allowed_tools_for_server_agent_intersection(self): @@ -3696,9 +3579,7 @@ class TestAgentMCPPermissions: key_perm = MagicMock() key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} team_perm = None - with patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=key_perm - ): + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm): with patch.object( MCPRequestHandler, "_get_team_object_permission", @@ -3730,9 +3611,7 @@ class TestAgentMCPPermissions: ) key_perm = MagicMock() key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} - with patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=key_perm - ): + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm): with patch.object( MCPRequestHandler, "_get_team_object_permission", @@ -3762,9 +3641,7 @@ class TestAgentMCPPermissions: agent_row = MagicMock() agent_row.object_permission_id = "perm-xyz" prisma_client = MagicMock() - prisma_client.db.litellm_agentstable.find_unique = AsyncMock( - return_value=agent_row - ) + prisma_client.db.litellm_agentstable.find_unique = AsyncMock(return_value=agent_row) user_api_key_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", @@ -3782,9 +3659,7 @@ class TestAgentMCPPermissions: return_value=expected_perm, ) as mock_get_perm, ): - result = await MCPRequestHandler._get_agent_object_permission( - user_api_key_auth - ) + result = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) assert result is expected_perm mock_get_perm.assert_awaited_once() assert mock_get_perm.await_args.kwargs["object_permission_id"] == "perm-xyz" @@ -3804,9 +3679,7 @@ class TestAgentMCPPermissions: agent_row = MagicMock() agent_row.object_permission_id = None prisma_client = MagicMock() - prisma_client.db.litellm_agentstable.find_unique = AsyncMock( - return_value=agent_row - ) + prisma_client.db.litellm_agentstable.find_unique = AsyncMock(return_value=agent_row) user_api_key_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", @@ -3822,14 +3695,8 @@ class TestAgentMCPPermissions: new_callable=AsyncMock, ) as mock_get_perm, ): - assert ( - await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) - is None - ) - assert ( - await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) - is None - ) + assert await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) is None + assert await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) is None mock_get_perm.assert_not_awaited() prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once() @@ -3870,9 +3737,7 @@ async def test_tool_permission_servers_included_in_allowed_servers(): ) with ( - patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=perm - ), + patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=perm), patch.object( MCPRequestHandler, "_get_mcp_servers_from_access_groups", @@ -4105,9 +3970,7 @@ class TestOrgMCPPermissions: org_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} with ( - patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=key_perm - ), + patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm), patch.object( MCPRequestHandler, "_get_team_object_permission", @@ -4137,9 +4000,7 @@ class TestOrgMCPPermissions: org_perm.mcp_tool_permissions = {} with ( - patch.object( - MCPRequestHandler, "_get_key_object_permission", return_value=key_perm - ), + patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm), patch.object( MCPRequestHandler, "_get_team_object_permission", @@ -4231,9 +4092,7 @@ async def test_mcp_key_access_group_extras_when_team_authorized(): ] _start_patches(patches) try: - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == ["srv-stripe"] finally: _stop_patches(patches) @@ -4270,9 +4129,7 @@ async def test_mcp_key_access_group_extras_when_key_directly_authorized(): ] _start_patches(patches) try: - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == ["srv-stripe"] finally: _stop_patches(patches) @@ -4286,9 +4143,7 @@ async def test_mcp_key_access_group_extras_when_key_has_no_groups(): access_group_ids=[], team_id="team-a", ) - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == [] @@ -4315,9 +4170,7 @@ async def test_mcp_key_access_group_extras_when_group_has_no_servers(): ] _start_patches(patches) try: - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == [] finally: _stop_patches(patches) @@ -4351,9 +4204,7 @@ async def test_mcp_key_access_group_extras_granted_even_when_group_authorizes_ne ] _start_patches(patches) try: - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == ["srv-finance-only"] finally: _stop_patches(patches) @@ -4376,9 +4227,7 @@ async def test_mcp_key_access_group_extras_when_get_access_object_raises(): ] _start_patches(patches) try: - result = await MCPRequestHandler._get_key_access_group_mcp_server_extras( - valid_token - ) + result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(valid_token) assert result == [] finally: _stop_patches(patches) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py index fbe768e07c6..17960e917a4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( AuthorizationCodeConfig, CredError, NoneConfig, + PassthroughConfig, SharedKey, TokenExchangeConfig, ) @@ -234,6 +235,12 @@ def test_token_exchange_empty_subject_token_type_normalizes_to_default(): assert spec.config.subject_token_type == "urn:ietf:params:oauth:token-type:access_token" +@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate]) +def test_client_forwarded_modes_map_to_passthrough_config(auth_type): + spec = to_server_spec(_server(auth_type=auth_type)) + assert spec is not None and isinstance(spec.config, PassthroughConfig) + + @pytest.mark.parametrize( "server", [ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index f8e45b38b49..64226eea821 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -1,9 +1,9 @@ """Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed. -`none`, `api_key` (shared-key source), `authorization_code`, and `token_exchange` are implemented; -every other arm, plus the `api_key` BYOK source, returns a typed `not_implemented` error until its -mode lands. Parametrizing the stubs over one config each also guards reachability: a dropped `case` -would hit `assert_never` and raise instead of returning the stub. +`none`, `api_key` (shared-key source), `passthrough`, `authorization_code`, and `token_exchange` are +implemented; every other arm, plus the `api_key` BYOK source, returns a typed `not_implemented` error +until its mode lands. Parametrizing the stubs over one config each also guards reachability: a dropped +`case` would hit `assert_never` and raise instead of returning the stub. """ import httpx @@ -253,9 +253,24 @@ async def test_token_exchange_without_an_exchanger_fails_closed(): assert result.error.tag == "misconfigured" +@pytest.mark.asyncio +async def test_passthrough_forwards_the_inbound_token_verbatim(): + subject = Subject(tenant_id="", subject_id="", inbound_token=SecretStr("Bearer upstream-xyz")) + result = await UpstreamCredentialProvider().resolve_credentials(subject, _spec(PassthroughConfig())) + assert isinstance(result, Ok) + assert isinstance(result.ok, StaticHeaderAuth) + assert _emitted(result.ok)["Authorization"] == "Bearer upstream-xyz" + + +@pytest.mark.asyncio +async def test_passthrough_without_inbound_token_is_a_no_op(): + result = await UpstreamCredentialProvider().resolve_credentials(_SUBJECT, _spec(PassthroughConfig())) + assert isinstance(result, Ok) + assert isinstance(result.ok, NoOpAuth) + + _STUBBED = [ ("api_key_byok", ApiKeyConfig(key_source=Byok())), - ("passthrough", PassthroughConfig()), ("client_credentials", ClientCredentialsConfig()), ("aws_sigv4", AwsSigV4Config(region="us-east-1")), ] 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 6fa69b2f96c..a7a0379c7ad 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 @@ -1267,6 +1267,229 @@ class TestMCPServerManager: assert captured_extra_headers == {"Authorization": "Bearer upstream-oauth-bearer"} + async def _capture_call_extra_headers(self, server, oauth2_headers, raw_headers, user_api_key_auth): + manager = MCPServerManager() + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) + captured = {"extra_headers": "unset"} + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs + ): # pragma: no cover - helper + captured["extra_headers"] = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + proxy_logging_obj=None, + user_api_key_auth=user_api_key_auth, + ) + return captured["extra_headers"] + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_true_passthrough_forwards_authorization(self): + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-true-passthrough", + name="tp-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + ) + extra_headers = await self._capture_call_extra_headers( + server, + oauth2_headers={"Authorization": "Bearer upstream-token"}, + raw_headers={"authorization": "Bearer upstream-token"}, + user_api_key_auth=UserAPIKeyAuth(api_key=None), + ) + assert extra_headers == {"Authorization": "Bearer upstream-token"} + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_oauth_delegate_forwards_separate_authorization(self): + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-oauth-delegate", + name="od-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + ) + extra_headers = await self._capture_call_extra_headers( + server, + oauth2_headers={"Authorization": "Bearer upstream-token"}, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer upstream-token", + }, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + assert extra_headers == {"Authorization": "Bearer upstream-token"} + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_oauth_delegate_never_forwards_admission_key(self): + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-oauth-delegate-leak", + name="od-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + ) + extra_headers = await self._capture_call_extra_headers( + server, + oauth2_headers={"Authorization": "Bearer sk-litellm-key"}, + raw_headers={"authorization": "Bearer sk-litellm-key"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + assert not extra_headers or "authorization" not in {k.lower() for k in extra_headers} + + def test_should_strip_caller_authorization_new_modes(self): + from litellm.proxy._types import UserAPIKeyAuth + + true_passthrough = MCPServer( + server_id="tp", + name="tp", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + ) + assert ( + _should_strip_caller_authorization( + mcp_server=true_passthrough, + raw_headers={"authorization": "Bearer upstream"}, + user_api_key_auth=UserAPIKeyAuth(api_key=None), + ) + is False + ) + + oauth_delegate = MCPServer( + server_id="od", + name="od", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + ) + assert ( + _should_strip_caller_authorization( + mcp_server=oauth_delegate, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer upstream", + }, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + is False + ) + assert ( + _should_strip_caller_authorization( + mcp_server=oauth_delegate, + raw_headers={"authorization": "Bearer sk-litellm-key"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + is True + ) + + def test_should_strip_authorization_for_oauth_delegate_admitted_via_jwt_without_api_key(self): + """JWT / SSO / OIDC / session admission yields a UserAPIKeyAuth with a user_id but + api_key=None; the caller's Authorization was that credential and must be stripped for + oauth_delegate when no separate x-litellm-api-key carried admission (LIT-3794-class leak).""" + from litellm.proxy._types import UserAPIKeyAuth + + oauth_delegate = MCPServer( + server_id="od-jwt", + name="od", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + ) + assert ( + _should_strip_caller_authorization( + mcp_server=oauth_delegate, + raw_headers={"authorization": "Bearer eyJ-idp-jwt"}, + user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None), + ) + is True + ) + assert ( + _should_strip_caller_authorization( + mcp_server=oauth_delegate, + raw_headers={ + "x-litellm-api-key": "Bearer sk-1234", + "authorization": "Bearer upstream", + }, + user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None), + ) + is False + ) + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_oauth_delegate_never_forwards_jwt_admission(self): + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="od-jwt-e2e", + name="od", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + ) + extra_headers = await self._capture_call_extra_headers( + server, + oauth2_headers={"Authorization": "Bearer eyJ-idp-jwt"}, + raw_headers={"authorization": "Bearer eyJ-idp-jwt"}, + user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None), + ) + assert not extra_headers or "authorization" not in {k.lower() for k in extra_headers} + + def test_new_passthrough_modes_require_per_user_auth(self): + for auth_type in (MCPAuth.true_passthrough, MCPAuth.oauth_delegate): + server = MCPServer( + server_id="s", + name="s", + url="https://example.com", + transport=MCPTransport.http, + auth_type=auth_type, + ) + assert server.requires_per_user_auth is True + + @pytest.mark.asyncio + async def test_create_mcp_client_forwarded_modes_use_the_passthrough_arm(self): + manager = MCPServerManager() + server = MCPServer( + server_id="tp-egress", + name="tp", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + ) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth", + new_callable=AsyncMock, + ) as mock_resolve, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls, + ): + await manager._create_mcp_client(server=server, extra_headers={"Authorization": "Bearer upstream-token"}) + mock_resolve.assert_not_awaited() + kwargs = mock_client_cls.call_args.kwargs + emitted = httpx.Request("GET", "https://example.com/mcp") + flow = kwargs["resolved_auth"].auth_flow(emitted) + next(flow) + flow.close() + assert emitted.headers["Authorization"] == "Bearer upstream-token" + assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]} + @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): """Ensure prompts are fetched and prefixed when requested.""" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d9bba85bf4b..9765775af8f 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27008,7 +27008,7 @@ export interface components { /** Alias */ alias?: string | null; /** Auth Type */ - auth_type?: ("none" | "api_key" | "bearer_token" | "basic" | "authorization" | "oauth2" | "aws_sigv4" | "token" | "oauth2_token_exchange") | null; + auth_type?: ("none" | "api_key" | "bearer_token" | "basic" | "authorization" | "oauth2" | "aws_sigv4" | "token" | "oauth2_token_exchange" | "true_passthrough" | "oauth_delegate") | null; /** Mcp Info */ mcp_info?: { [key: string]: unknown;