fix: prefer _get_key_object_permission for key lookups and remove redundant checks

This commit is contained in:
Yuta Saito
2025-12-18 09:50:56 +09:00
parent cdcbccb30d
commit 847bbd4fda
2 changed files with 132 additions and 29 deletions
@@ -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(
@@ -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"])