diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7509f62f7f4..274487af9f7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1361,6 +1361,18 @@ async def _byok_listing_auth_header( return await _get_byok_credential(mcp_server, user_api_key_auth) +def _listing_upstream_auth_header( + mcp_server: MCPServer, + mcp_auth_header: str | dict[str, str] | None, + byok_auth_header: str | dict[str, str] | None, +) -> str | dict[str, str] | None: + """The credential a tools/list sends upstream: an oauth2 server's is minted or signed, never the + stored BYOK one, the same rule ``_get_tools_from_mcp_servers`` applies to its own resolution.""" + if mcp_server.auth_type == MCPAuth.oauth2: + return mcp_auth_header + return byok_auth_header + + def _client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, @@ -4500,6 +4512,9 @@ class MCPServerManager: client = None resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header) + upstream_mcp_auth_header: Final = _listing_upstream_auth_header( + server, mcp_auth_header, resolved_mcp_auth_header + ) listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, mcp_auth_header=resolved_mcp_auth_header, @@ -4549,7 +4564,7 @@ class MCPServerManager: if ( get_mcp_jwt_signer() is not None and not has_static_authorization - and not resolved_mcp_auth_header + and not upstream_mcp_auth_header and not has_extra_authorization ): extra_headers = await inject_mcp_jwt_headers_for_upstream( @@ -4572,7 +4587,7 @@ class MCPServerManager: client = await self._create_mcp_client( server=server, - mcp_auth_header=resolved_mcp_auth_header, + mcp_auth_header=upstream_mcp_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, 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 eab57b0aee4..28a1098b436 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 @@ -7380,6 +7380,58 @@ class TestMCPServerManager: assert listed is not None and listed.description == "t" assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + @pytest.mark.asyncio + async def test_oauth2_byok_listing_leaves_the_minted_token_and_signer_in_place(self): + """The stored BYOK secret keys the catalog slot tools/call reads, but never reaches the oauth2 + upstream on tools/list: the client mints its M2M token and MCPJWTSigner still signs.""" + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache_key, + cache_byok_credential, + ) + from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache + + manager = MCPServerManager() + server = MCPServer( + server_id="cc1", + name="cc1", + transport=MCPTransport.http, + url="http://cc1", + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="csec", + token_url="http://cc1/token", + is_byok=True, + ) + alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="echo", description="m2m catalog", inputSchema={})] + ) + signer_headers = AsyncMock(return_value={"Authorization": "Bearer signed-jwt"}) + cache_byok_credential("alice", "cc1", "BYOK-ALICE-SECRET") + try: + with ( + patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ), + patch( # test-quality-ok: same singleton's header injection, asserted on by call + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.inject_mcp_jwt_headers_for_upstream", + signer_headers, + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=alice) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) + + client_kwargs = manager._create_mcp_client.await_args.kwargs + assert client_kwargs["mcp_auth_header"] is None, client_kwargs + assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} + signer_headers.assert_awaited_once() + call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET") + listed = manager.get_listed_tool(server, "echo", call_side) + assert listed is not None and listed.description == "m2m catalog" + @pytest.mark.parametrize( ("signer", "static_headers"), [