diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 05b5cc67d83..37d69011390 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -190,6 +190,10 @@ def _should_strip_caller_authorization( Strip rules: - **M2M (client_credentials) servers**: never forward the caller's ``Authorization`` — the proxy fetches its own upstream token. + - **Migrated per-user OAuth (authorization_code) servers**: never forward + the caller's ``Authorization`` — the v2 resolver injects the stored + per-user token, so a caller-supplied bearer cannot override another + user's stored credential. Delegate / pass-through keep forwarding it. - **OAuth pass-through servers**: strip when the ``Authorization`` header is actually the LiteLLM API key — either because admission validated it (``user_api_key_auth.api_key`` is set) and the caller @@ -202,6 +206,15 @@ def _should_strip_caller_authorization( """ if mcp_server.has_client_credentials: return True + if ( + mcp_server.auth_type == MCPAuth.oauth2 + and to_server_spec(mcp_server) is not None + ): + # Migrated per-user OAuth (authorization_code): the v2 resolver injects the + # stored token, so a caller-forwarded Authorization must not be forwarded + # 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: return False @@ -221,6 +234,18 @@ def _should_strip_caller_authorization( ) +def _without_authorization( + headers: Optional[dict[str, str]], +) -> Optional[dict[str, str]]: + """A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or + None if nothing remains. Drops only the credential, keeping other forwarded headers. + """ + if not headers: + return None + filtered = {k: v for k, v in headers.items() if k.lower() != "authorization"} + return filtered or None + + def _extract_upstream_auth_failure( exc: BaseException, ) -> Optional[Tuple[int, Optional[str]]]: @@ -3509,6 +3534,16 @@ class MCPServerManager: extra_headers = None else: extra_headers = oauth2_headers + # Migrated authorization_code: the v2 resolver injects the stored per-user + # token, so drop the caller-forwarded Authorization (apply-if-absent would + # otherwise let it shadow the resolved token). Delegate keeps it. Centralized + # via _should_strip_caller_authorization to match _prepare_mcp_server_headers. + if extra_headers and _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ): + extra_headers = _without_authorization(extra_headers) if mcp_server.extra_headers and raw_headers: if extra_headers is None: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e5e5b709121..55a0641c887 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -283,6 +283,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _should_strip_caller_authorization, + _without_authorization, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( @@ -1435,6 +1436,16 @@ if MCP_AVAILABLE: else: # Copy to avoid mutating the original dict (important for parallel fetching) extra_headers = oauth2_headers.copy() if oauth2_headers else None + # Migrated authorization_code: the v2 resolver injects the stored per-user + # token, so drop the caller-forwarded Authorization (apply-if-absent would + # otherwise let it shadow the resolved token). Delegate keeps it. Centralized + # via _should_strip_caller_authorization to match _call_regular_mcp_tool. + if extra_headers and _should_strip_caller_authorization( + mcp_server=server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ): + extra_headers = _without_authorization(extra_headers) if server.extra_headers and raw_headers: if extra_headers is None: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ca326d197ae..34e932b6ae7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -284,8 +284,11 @@ def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorizatio assert extra_headers is None -def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers(): - """Interactive OAuth still forwards the user's OAuth token in extra_headers.""" +def test_prepare_mcp_server_headers_oauth2_interactive_drops_caller_authorization(): + """A v2-migrated interactive OAuth (authorization_code) server must NOT forward the + caller's Authorization: the resolver injects the stored per-user token, so a + caller-supplied bearer must not override another user's stored credential. Non-auth + headers are still carried; only the credential is dropped.""" try: from litellm.proxy._experimental.mcp_server.server import ( _prepare_mcp_server_headers, @@ -293,7 +296,7 @@ def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers(): except ImportError: pytest.skip("MCP server not available") - user_oauth = {"Authorization": "Bearer upstream-user-token"} + caller_oauth = {"Authorization": "Bearer caller-supplied-token"} server = MCPServer( server_id="3lo-server", @@ -307,12 +310,13 @@ def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers(): server=server, mcp_server_auth_headers=None, mcp_auth_header=None, - oauth2_headers=user_oauth, + oauth2_headers=caller_oauth, raw_headers=None, ) assert server_auth_header is None - assert extra_headers == user_oauth + # Caller's Authorization is dropped (only key present) -> extra_headers is None. + assert extra_headers is None def test_prepare_mcp_server_headers_m2m_skips_authorization_from_raw_extra_headers(): @@ -2813,8 +2817,10 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): @pytest.mark.asyncio @pytest.mark.no_parallel -async def test_oauth2_headers_passed_to_mcp_client(): - """Test that OAuth2 headers are properly passed through to the MCP client for OAuth2 servers like github_mcp""" +async def test_oauth2_caller_headers_not_forwarded_for_migrated_server(): + """A v2-migrated authorization_code server (like github_mcp) must NOT forward the + caller's oauth2 Authorization to the MCP client — the resolver injects the stored + per-user token, so a caller-supplied bearer cannot override another user's credential.""" try: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -2928,20 +2934,13 @@ async def test_oauth2_headers_passed_to_mcp_client(): assert captured_client_args["server"].server_id == oauth2_server.server_id assert captured_client_args["server"].auth_type == MCPAuth.oauth2 - # Most importantly: verify that OAuth2 headers were passed as extra_headers - assert ( - captured_client_args["extra_headers"] is not None - ), "Expected extra_headers to be passed for OAuth2 server" - assert ( - captured_client_args["extra_headers"] == oauth2_headers - ), f"Expected OAuth2 headers to be passed as extra_headers, got {captured_client_args['extra_headers']}" - - # Verify the Authorization header specifically - assert "Authorization" in captured_client_args["extra_headers"] - assert ( - captured_client_args["extra_headers"]["Authorization"] - == "Bearer github_oauth_token_12345" - ) + # Security: a v2-migrated authorization_code server must NOT forward the caller's + # oauth2 Authorization upstream. The v2 resolver injects the stored per-user token, + # so a caller-supplied bearer cannot override another user's stored credential. + extra_headers = captured_client_args["extra_headers"] + assert extra_headers is None or "Authorization" not in { + k.lower() for k in extra_headers + }, f"Caller Authorization must not be forwarded, got {extra_headers}" @pytest.mark.asyncio 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 a1ba103edb9..598be9d2f8b 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 @@ -31,6 +31,8 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _deserialize_json_dict, _deserialize_json_list, _normalize_mcp_server_cost_info, + _should_strip_caller_authorization, + _without_authorization, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -600,6 +602,84 @@ class TestMCPServerManager: assert captured_extra_headers == {"x-request-id": "req-123"} + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_v2_authz_code_drops_caller_authorization( + self, + ): + """A v2-migrated per-user OAuth (authorization_code) server must NOT seed a + caller-forwarded Authorization into extra_headers — the resolver injects the + stored per-user token, and apply-if-absent would otherwise let the caller's + header override another user's stored credential (matches v1's overwrite).""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + # oauth2, not M2M, not delegate => to_server_spec maps it to AuthorizationCodeConfig + server = MCPServer( + server_id="server-authz-code-call", + name="authz-code-server", + url="https://example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + # Migrated authorization_code => the centralized strip decision says drop the + # caller's Authorization (the v2 resolver injects the stored token). + assert ( + _should_strip_caller_authorization( + mcp_server=server, raw_headers=None, user_api_key_auth=None + ) + is True + ) + + 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 + nonlocal captured_extra_headers + 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={"Authorization": "Bearer caller-supplied-token"}, + raw_headers={"authorization": "Bearer caller-supplied-token"}, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + # The caller's Authorization must not reach extra_headers; the v2 resolver is + # the sole Authorization source for this server. + assert captured_extra_headers != "unset" + if captured_extra_headers: + assert "authorization" not in {k.lower() for k in captured_extra_headers} + + def test_without_authorization_drops_only_the_credential(self): + # None / empty -> None + assert _without_authorization(None) is None + assert _without_authorization({}) is None + # Only Authorization present -> nothing left -> None (case-insensitive) + assert _without_authorization({"authorization": "Bearer x"}) is None + # Authorization dropped, other headers kept + assert _without_authorization( + {"Authorization": "Bearer x", "X-Trace-Id": "t"} + ) == {"X-Trace-Id": "t"} + @pytest.mark.asyncio async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header( self,