From 91f6661b37e88e60ff19d848e0c0edbba8d8423c Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 25 Apr 2026 12:39:44 -0700 Subject: [PATCH] [Fix] Align MCP OAuth proxy endpoints with per-server access policy Bring `/server/oauth/{server_id}/authorize`, `/token`, and `/register` in line with `fetch_mcp_server`: the helper that resolves the server now also applies the per-caller access policy. Admin-view callers are unrestricted; non-admins must have the server in their allowed-servers set; servers resolved from the admin-only `/server/oauth/session` temporary cache reject non-admins. --- .../mcp_management_endpoints.py | 36 ++++- .../test_mcp_management_endpoints.py | 152 ++++++++++++++++-- 2 files changed, 164 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index a68c8ca9fa8..fca08e591fd 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1448,15 +1448,15 @@ if MCP_AVAILABLE: return _redact_mcp_credentials(temp_record) async def _get_cached_temporary_mcp_server_or_404( - server_id: str, request: Optional[Request] = None + server_id: str, + user_api_key_dict: UserAPIKeyAuth, + request: Optional[Request] = None, ) -> MCPServer: server = await get_cached_temporary_mcp_server(server_id) + resolved_from_temp_cache = server is not None if server is None: # Fall back to real DB/config server (e.g. for the user-side OAuth flow # which calls these endpoints with a real server_id, not a temp session id). - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils client_ip = IPAddressUtils.get_mcp_client_ip(request) if request else None @@ -1470,6 +1470,28 @@ if MCP_AVAILABLE: status_code=status.HTTP_404_NOT_FOUND, detail={"error": f"MCP server {server_id} not found"}, ) + + # Per-server access policy mirrors `fetch_mcp_server`: admin-view + # callers are unrestricted; non-admins must have the server in their + # allowed-servers set. Temporary cached servers come from the + # admin-only `/server/oauth/session` setup flow and are not exposed + # to non-admins. + if not _user_has_admin_view(user_api_key_dict): + if resolved_from_temp_cache: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": f"Access denied to MCP server {server_id}"}, + ) + allowed_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_dict + ) + ) + if server.server_id not in allowed_server_ids: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": f"Access denied to MCP server {server_id}"}, + ) return server @router.get( @@ -1490,7 +1512,7 @@ if MCP_AVAILABLE: scope: Optional[str] = None, ): mcp_server = await _get_cached_temporary_mcp_server_or_404( - server_id, request=request + server_id, user_api_key_dict, request=request ) # Use the server's stored client_id when the caller doesn't supply one resolved_client_id = mcp_server.client_id or client_id or "" @@ -1536,7 +1558,7 @@ if MCP_AVAILABLE: scope: Optional[str] = Form(None), ): mcp_server = await _get_cached_temporary_mcp_server_or_404( - server_id, request=request + server_id, user_api_key_dict, request=request ) resolved_client_id = mcp_server.client_id or client_id or "" if not resolved_client_id: @@ -1574,7 +1596,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): mcp_server = await _get_cached_temporary_mcp_server_or_404( - server_id, request=request + server_id, user_api_key_dict, request=request ) request_data = await _read_request_body(request=request) data: dict = {**request_data} diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 442265d3af0..821e2002906 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1341,12 +1341,15 @@ class TestTemporaryMCPSessionEndpoints: ) server = generate_mock_mcp_server_config_record(server_id="cached") + admin_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + ) with patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server", return_value=server, ) as get_cached: - result = await _get_cached_temporary_mcp_server_or_404("cached") + result = await _get_cached_temporary_mcp_server_or_404("cached", admin_auth) assert result is server get_cached.assert_awaited_once_with("cached") @@ -1356,10 +1359,95 @@ class TestTemporaryMCPSessionEndpoints: return_value=None, ): with pytest.raises(HTTPException) as exc_info: - await _get_cached_temporary_mcp_server_or_404("missing") + await _get_cached_temporary_mcp_server_or_404("missing", admin_auth) assert exc_info.value.status_code == 404 + @pytest.mark.asyncio + async def test_get_cached_temporary_mcp_server_non_admin_denied(self): + """Non-admin without access to the server gets 403, not the server.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _get_cached_temporary_mcp_server_or_404, + ) + + registry_server = generate_mock_mcp_server_config_record(server_id="server-x") + non_admin = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER, + ) + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id.return_value = registry_server + mock_manager.get_mcp_server_by_name.return_value = None + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[]) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server", + return_value=None, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await _get_cached_temporary_mcp_server_or_404("server-x", non_admin) + + assert exc_info.value.status_code == 403 + mock_manager.get_allowed_mcp_servers.assert_awaited_once_with(non_admin) + + @pytest.mark.asyncio + async def test_get_cached_temporary_mcp_server_non_admin_allowed(self): + """Non-admin with the server in their allowed set gets the server.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _get_cached_temporary_mcp_server_or_404, + ) + + registry_server = generate_mock_mcp_server_config_record(server_id="server-x") + non_admin = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER, + ) + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id.return_value = registry_server + mock_manager.get_mcp_server_by_name.return_value = None + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server-x"]) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server", + return_value=None, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + result = await _get_cached_temporary_mcp_server_or_404( + "server-x", non_admin + ) + + assert result is registry_server + + @pytest.mark.asyncio + async def test_get_cached_temporary_mcp_server_temp_cache_non_admin_denied(self): + """Servers resolved from the admin-only temp cache reject non-admins.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _get_cached_temporary_mcp_server_or_404, + ) + + temp_server = generate_mock_mcp_server_config_record(server_id="temp-cache") + non_admin = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server", + return_value=temp_server, + ): + with pytest.raises(HTTPException) as exc_info: + await _get_cached_temporary_mcp_server_or_404("temp-cache", non_admin) + + assert exc_info.value.status_code == 403 + @pytest.mark.asyncio async def test_add_session_mcp_server_caches_and_redacts_credentials(self): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( @@ -1472,6 +1560,9 @@ class TestTemporaryMCPSessionEndpoints: request = MagicMock() server = generate_mock_mcp_server_config_record(server_id="server-1") authorize_response = MagicMock() + admin_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + ) with ( patch( @@ -1486,6 +1577,7 @@ class TestTemporaryMCPSessionEndpoints: result = await mcp_authorize( request=request, server_id="server-1", + user_api_key_dict=admin_auth, client_id="client-id", redirect_uri="https://example.com/callback", state="state123", @@ -1496,7 +1588,7 @@ class TestTemporaryMCPSessionEndpoints: ) assert result is authorize_response - get_server.assert_awaited_once_with("server-1", request=request) + get_server.assert_awaited_once_with("server-1", admin_auth, request=request) authorize_mock.assert_awaited_once_with( request=request, mcp_server=server, @@ -1518,6 +1610,9 @@ class TestTemporaryMCPSessionEndpoints: request = MagicMock() server = generate_mock_mcp_server_config_record(server_id="server-1") exchange_response = {"access_token": "token"} + admin_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + ) with ( patch( @@ -1532,6 +1627,7 @@ class TestTemporaryMCPSessionEndpoints: result = await mcp_token( request=request, server_id="server-1", + user_api_key_dict=admin_auth, grant_type="authorization_code", code="code-123", redirect_uri="https://example.com/callback", @@ -1543,7 +1639,7 @@ class TestTemporaryMCPSessionEndpoints: ) assert result is exchange_response - get_server.assert_awaited_once_with("server-1", request=request) + get_server.assert_awaited_once_with("server-1", admin_auth, request=request) exchange_mock.assert_awaited_once_with( request=request, mcp_server=server, @@ -1566,6 +1662,9 @@ class TestTemporaryMCPSessionEndpoints: request = MagicMock() server = generate_mock_mcp_server_config_record(server_id="server-1") exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"} + admin_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + ) with ( patch( @@ -1580,6 +1679,7 @@ class TestTemporaryMCPSessionEndpoints: result = await mcp_token( request=request, server_id="server-1", + user_api_key_dict=admin_auth, grant_type="refresh_token", code=None, redirect_uri=None, @@ -1591,7 +1691,7 @@ class TestTemporaryMCPSessionEndpoints: ) assert result is exchange_response - get_server.assert_awaited_once_with("server-1", request=request) + get_server.assert_awaited_once_with("server-1", admin_auth, request=request) exchange_mock.assert_awaited_once_with( request=request, mcp_server=server, @@ -1620,6 +1720,9 @@ class TestTemporaryMCPSessionEndpoints: "response_types": ["code"], "token_endpoint_auth_method": "client_secret_basic", } + admin_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + ) with ( patch( @@ -1635,10 +1738,14 @@ class TestTemporaryMCPSessionEndpoints: AsyncMock(return_value=register_response), ) as register_mock, ): - result = await mcp_register(request=request, server_id="server-1") + result = await mcp_register( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + ) assert result is register_response - get_server.assert_awaited_once_with("server-1", request=request) + get_server.assert_awaited_once_with("server-1", admin_auth, request=request) read_body.assert_awaited_once_with(request=request) register_mock.assert_awaited_once_with( request=request, @@ -1664,12 +1771,15 @@ class TestTemporaryMCPSessionEndpoints: original_cache = mgmt_endpoints.litellm.cache mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend) try: - with patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers", - {}, - ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper", - return_value=serialized, + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers", + {}, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper", + return_value=serialized, + ), ): result = await get_cached_temporary_mcp_server("from-redis") finally: @@ -1734,7 +1844,9 @@ class TestTemporaryMCPSessionEndpoints: _get_temporary_mcp_server_from_redis, ) - server = generate_mock_mcp_server_config_record(server_id="from-redis-encrypted") + server = generate_mock_mcp_server_config_record( + server_id="from-redis-encrypted" + ) serialized = json.dumps(server.model_dump(mode="json")) mock_cache_backend = SimpleNamespace( async_get_cache=AsyncMock(return_value="encrypted-payload") @@ -1808,7 +1920,9 @@ class TestTemporaryMCPSessionEndpoints: _get_temporary_mcp_server_from_redis, ) - mock_cache_backend = SimpleNamespace(async_get_cache=AsyncMock(return_value="enc")) + mock_cache_backend = SimpleNamespace( + async_get_cache=AsyncMock(return_value="enc") + ) original_cache = mgmt_endpoints.litellm.cache mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend) try: @@ -1823,12 +1937,16 @@ class TestTemporaryMCPSessionEndpoints: assert result is None @pytest.mark.asyncio - async def test_get_temporary_mcp_server_from_redis_returns_none_on_decrypt_none(self): + async def test_get_temporary_mcp_server_from_redis_returns_none_on_decrypt_none( + self, + ): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( _get_temporary_mcp_server_from_redis, ) - mock_cache_backend = SimpleNamespace(async_get_cache=AsyncMock(return_value="enc")) + mock_cache_backend = SimpleNamespace( + async_get_cache=AsyncMock(return_value="enc") + ) original_cache = mgmt_endpoints.litellm.cache mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend) try: