From 847bbd4fdaef0015f7060add08e73a3d2b904c0d Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Thu, 18 Dec 2025 09:50:56 +0900 Subject: [PATCH] fix: prefer _get_key_object_permission for key lookups and remove redundant checks --- .../mcp_server/auth/user_api_key_auth_mcp.py | 31 +---- .../auth/test_user_api_key_auth_mcp.py | 130 ++++++++++++++++++ 2 files changed, 132 insertions(+), 29 deletions(-) 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 d6df3b76f1..b43f421717 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 @@ -525,30 +525,9 @@ class MCPRequestHandler: async def _get_allowed_mcp_servers_for_key( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[str]: - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - if user_api_key_auth is None: - return [] - - if user_api_key_auth.object_permission_id is None: - return [] - - if prisma_client is None: - verbose_logger.debug("prisma_client is None") - return [] - try: - key_object_permission = await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, + key_object_permission = await MCPRequestHandler._get_key_object_permission( + user_api_key_auth ) if key_object_permission is None: return [] @@ -583,12 +562,6 @@ class MCPRequestHandler: 1. First checks if object_permission is already loaded on the team 2. If not, fetches from DB using object_permission_id if it exists """ - if user_api_key_auth is None: - return [] - - if user_api_key_auth.team_id is None: - return [] - try: # Use the helper method that properly handles fetching from DB if needed object_permissions = await MCPRequestHandler._get_team_object_permission( 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 7927aa7f48..21782f4218 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 @@ -1258,3 +1258,133 @@ async def test_get_allowed_mcp_servers_for_team_with_no_object_permission(): # Verify the helper was called mock_get_team_perm.assert_called_once_with(mock_user_auth) + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_for_team_without_user_auth_returns_empty(): + """Ensure helper returns empty list when no user auth is provided.""" + + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(None) + + assert result == [] + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_for_team_without_team_id_returns_empty(): + """Ensure helper returns empty list when user lacks a team_id.""" + + mock_user_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + team_id=None, + ) + + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( + mock_user_auth + ) + + assert result == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "user_api_key_auth, prisma_client_value, scenario", + [ + (None, object(), "no_user"), + ( + UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + object(), + "no_object_permission_id", + ), + ( + UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + object_permission_id="perm-123", + ), + None, + "no_prisma_client", + ), + ], +) +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( + "litellm.proxy.auth.auth_checks.get_object_permission", + 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 + ) + + assert result == [] + mock_get_perm.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_none(): + """Ensure [] is returned when get_object_permission yields None.""" + + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + object_permission_id="perm-123", + ) + + mock_prisma = object() + + with patch( + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + ) as mock_get_perm: + mock_get_perm.return_value = None + + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( + user_api_key_auth + ) + + assert result == [] + mock_get_perm.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): + """Ensure in-memory object_permission is used without hitting the DB.""" + + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + perms = LiteLLM_ObjectPermissionTable( + object_permission_id="perm-in-memory", + mcp_servers=["direct-server"], + mcp_access_groups=["grp-alpha"], + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + object_permission=perms, + ) + + with patch( + "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: + mock_access_groups.return_value = ["group-server"] + + 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() + mock_access_groups.assert_called_once_with(["grp-alpha"])