mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
feat: address greptile comments
This commit is contained in:
parent
beb426860d
commit
cf6b0e45a7
36 changed files with 698 additions and 37 deletions
|
|
@ -334,11 +334,16 @@ class MCPRequestHandler:
|
|||
"""
|
||||
Get list of allowed MCP servers for the given user/key based on permissions.
|
||||
|
||||
Permission hierarchy (all rules are intersections):
|
||||
1. Get allowed servers from key permissions
|
||||
2. Get allowed servers from team permissions
|
||||
3. Get allowed servers from end_user permissions
|
||||
4. Final result = intersection of key/team AND end_user (if end_user has permissions set)
|
||||
Permission hierarchy:
|
||||
1. Calculate key/team allowed servers:
|
||||
- If team has permissions: key inherits from team (or intersects if key has own permissions)
|
||||
- If team has no permissions: use key permissions
|
||||
2. Apply end_user permissions if end_user_id is set:
|
||||
- If require_end_user_mcp_access_defined=True and end_user has no permissions: block all access
|
||||
- If end_user has permissions:
|
||||
* If key/team has no restrictions (empty): use end_user permissions
|
||||
* If key/team has restrictions: intersect key/team AND end_user
|
||||
- If end_user has no permissions and flag is False: use key/team permissions
|
||||
|
||||
Returns:
|
||||
List[str]: List of allowed MCP servers by server id
|
||||
|
|
@ -386,40 +391,41 @@ class MCPRequestHandler:
|
|||
)
|
||||
)
|
||||
|
||||
# Check if require_end_user_mcp_access flag is enabled
|
||||
require_end_user_mcp_access = general_settings.get(
|
||||
"require_end_user_mcp_access_defined", False
|
||||
)
|
||||
|
||||
# If end_user has explicit MCP server permissions, apply intersection
|
||||
# If the flag is enabled and end_user has no permissions, block all access
|
||||
if require_end_user_mcp_access and len(allowed_mcp_servers_for_end_user) == 0:
|
||||
verbose_logger.debug(
|
||||
f"require_end_user_mcp_access_defined=True and end_user {user_api_key_auth.end_user_id} has no MCP permissions - blocking MCP access"
|
||||
)
|
||||
return []
|
||||
|
||||
# If end_user has explicit MCP server permissions, apply intersection logic
|
||||
if len(allowed_mcp_servers_for_end_user) > 0:
|
||||
verbose_logger.debug(
|
||||
f"End user {user_api_key_auth.end_user_id} has explicit MCP permissions: {allowed_mcp_servers_for_end_user}"
|
||||
)
|
||||
|
||||
# Check if require_end_user_mcp_access flag is enabled
|
||||
require_end_user_mcp_access = general_settings.get(
|
||||
"require_end_user_mcp_access_defined", False
|
||||
)
|
||||
|
||||
# If the flag is enabled and end_user has permissions, use ONLY end_user permissions
|
||||
if require_end_user_mcp_access:
|
||||
# If key/team has no restrictions (empty list), use end_user permissions
|
||||
if len(allowed_mcp_servers) == 0:
|
||||
verbose_logger.debug(
|
||||
"require_end_user_mcp_access_defined=True - using only end_user permissions"
|
||||
"No key/team MCP restrictions - using end_user permissions"
|
||||
)
|
||||
allowed_mcp_servers = allowed_mcp_servers_for_end_user
|
||||
else:
|
||||
# Otherwise, apply intersection: key/team AND end_user
|
||||
# Apply intersection: key/team AND end_user
|
||||
# This ensures key/team restrictions are ALWAYS respected
|
||||
filtered_servers = []
|
||||
for _mcp_server in allowed_mcp_servers:
|
||||
if _mcp_server in allowed_mcp_servers_for_end_user:
|
||||
filtered_servers.append(_mcp_server)
|
||||
allowed_mcp_servers = filtered_servers
|
||||
verbose_logger.debug(
|
||||
f"Applied end_user intersection filter. Final allowed servers: {allowed_mcp_servers}"
|
||||
f"Applied end_user intersection with key/team. Final allowed servers: {allowed_mcp_servers}"
|
||||
)
|
||||
# If flag is enabled but end_user has no permissions, block all access
|
||||
elif general_settings.get("require_end_user_mcp_access_defined", False):
|
||||
verbose_logger.debug(
|
||||
f"require_end_user_mcp_access_defined=True and end_user {user_api_key_auth.end_user_id} has no MCP permissions - blocking MCP access"
|
||||
)
|
||||
return []
|
||||
|
||||
return list(set(allowed_mcp_servers))
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -174,7 +174,8 @@ async def _handle_customer_object_permission_update(
|
|||
"""
|
||||
Handle object permission updates for customer endpoints.
|
||||
|
||||
Updates the update_end_user_table_data dict in place with the new object_permission_id.
|
||||
Updates the update_end_user_table_data dict in place with the new object_permission_id
|
||||
and removes the object_permission field to prevent it from being sent to the database.
|
||||
|
||||
Args:
|
||||
non_default_values: Dictionary containing the update values including object_permission
|
||||
|
|
@ -196,6 +197,10 @@ async def _handle_customer_object_permission_update(
|
|||
if object_permission_id is not None:
|
||||
update_end_user_table_data["object_permission_id"] = object_permission_id
|
||||
|
||||
# Remove object_permission from update_end_user_table_data to prevent DB write attempts
|
||||
# object_permission is a read-only relationship field, not a writable column
|
||||
update_end_user_table_data.pop("object_permission", None)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/end_user/new",
|
||||
|
|
@ -612,14 +617,6 @@ async def update_end_user(
|
|||
## Update user table, with update params + new budget id (if set) ##
|
||||
verbose_proxy_logger.debug("/customer/update: Received data = %s", data)
|
||||
|
||||
# Ensure object_permission is not in the update data
|
||||
# It should have been converted to object_permission_id by handle_update_object_permission_common
|
||||
if "object_permission" in update_end_user_table_data:
|
||||
verbose_proxy_logger.warning(
|
||||
f"object_permission still in update_end_user_table_data: {update_end_user_table_data.get('object_permission')}"
|
||||
)
|
||||
update_end_user_table_data.pop("object_permission", None)
|
||||
|
||||
if data.user_id is not None and len(data.user_id) > 0:
|
||||
update_end_user_table_data["user_id"] = data.user_id # type: ignore
|
||||
verbose_proxy_logger.debug("In update customer, user_id condition block.")
|
||||
|
|
|
|||
|
|
@ -15,9 +15,8 @@ sys.path.insert(
|
|||
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import \
|
||||
MCPRequestHandler
|
||||
from litellm.proxy._types import SpecialHeaders, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
|
|
@ -1176,7 +1175,8 @@ class TestMCPAccessGroupsE2E:
|
|||
@pytest.mark.asyncio
|
||||
def test_mcp_path_based_server_segregation(monkeypatch):
|
||||
# Import the MCP server FastAPI app and context getter
|
||||
from litellm.proxy._experimental.mcp_server.server import app, get_auth_context
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
app, get_auth_context)
|
||||
|
||||
captured_mcp_servers = {}
|
||||
|
||||
|
|
@ -1277,7 +1277,8 @@ async def test_get_team_object_permission_with_already_loaded_permission():
|
|||
Test that _get_team_object_permission returns the already loaded object_permission
|
||||
from the team object without making an additional DB call.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
|
||||
from litellm.proxy._types import (LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamTable)
|
||||
|
||||
# Create mock object permission
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
|
|
@ -1340,7 +1341,8 @@ async def test_get_team_object_permission_with_core_auth_auto_loading():
|
|||
the team object returned by get_team_object() should already have object_permission loaded
|
||||
when an object_permission_id exists.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
|
||||
from litellm.proxy._types import (LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamTable)
|
||||
|
||||
# Create mock object permission
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
|
|
@ -1595,3 +1597,659 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
|
|||
assert set(result) == {"direct-server", "group-server"}
|
||||
mock_get_perm.assert_not_called()
|
||||
mock_access_groups.assert_called_once_with(["grp-alpha"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestEndUserMCPPermissions:
|
||||
"""Test suite for end_user MCP permission functionality"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_servers,team_servers,end_user_servers,require_flag,expected_result,scenario",
|
||||
[
|
||||
# Test case 1: End user with permissions, key/team with permissions - intersection
|
||||
(
|
||||
["server1", "server2", "server3"],
|
||||
["server1", "server2", "server3", "server4"],
|
||||
["server2", "server3", "server5"],
|
||||
False,
|
||||
["server2", "server3"],
|
||||
"end_user_intersects_with_key_team",
|
||||
),
|
||||
# Test case 2: End user with permissions, key/team with no permissions (empty) - use end_user permissions
|
||||
(
|
||||
[],
|
||||
[],
|
||||
["server1", "server2"],
|
||||
False,
|
||||
["server1", "server2"],
|
||||
"end_user_with_empty_key_team",
|
||||
),
|
||||
# Test case 3: End user with no permissions, require_flag=True - fall back to key/team restrictions
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
[],
|
||||
True,
|
||||
["server1", "server2"],
|
||||
"end_user_no_perms_with_require_flag_true_fallback_to_key_team",
|
||||
),
|
||||
# Test case 4: End user with no permissions, require_flag=False - use key/team permissions
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
[],
|
||||
False,
|
||||
["server1", "server2"],
|
||||
"end_user_no_perms_with_require_flag_false",
|
||||
),
|
||||
# Test case 5: End user with permissions, require_flag=True - still intersect with key/team
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2", "server3"],
|
||||
["server2", "server3", "server4"],
|
||||
True,
|
||||
["server2"],
|
||||
"end_user_with_require_flag_true_still_intersects",
|
||||
),
|
||||
# Test case 6: End user has permissions but no overlap with key/team
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
["server3", "server4"],
|
||||
False,
|
||||
[],
|
||||
"end_user_no_overlap_with_key_team",
|
||||
),
|
||||
# Test case 7: Key has servers, team empty, end_user has overlap with key
|
||||
(
|
||||
["server1", "server2", "server3"],
|
||||
[],
|
||||
["server2", "server3", "server4"],
|
||||
False,
|
||||
["server2", "server3"],
|
||||
"key_only_intersects_with_end_user",
|
||||
),
|
||||
# Test case 8: All have different servers
|
||||
(
|
||||
["server1"],
|
||||
["server1", "server2"],
|
||||
["server3"],
|
||||
False,
|
||||
[],
|
||||
"all_different_servers",
|
||||
),
|
||||
# Test case 9: require_flag=True, end_user has permissions but key/team empty
|
||||
(
|
||||
[],
|
||||
[],
|
||||
["server1", "server2"],
|
||||
True,
|
||||
["server1", "server2"],
|
||||
"require_flag_true_with_empty_key_team_uses_end_user",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_get_allowed_mcp_servers_with_end_user(
|
||||
self,
|
||||
key_servers,
|
||||
team_servers,
|
||||
end_user_servers,
|
||||
require_flag,
|
||||
expected_result,
|
||||
scenario,
|
||||
):
|
||||
"""Test get_allowed_mcp_servers with end_user permissions"""
|
||||
|
||||
# Create a mock user with end_user_id
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the helper methods
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_end_user"
|
||||
) as mock_end_user_servers:
|
||||
with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings:
|
||||
# Set up return values
|
||||
mock_key_servers.return_value = key_servers
|
||||
mock_team_servers.return_value = team_servers
|
||||
mock_end_user_servers.return_value = end_user_servers
|
||||
mock_general_settings.get.return_value = require_flag
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# Assert the result (order-independent comparison)
|
||||
assert sorted(result) == sorted(expected_result)
|
||||
|
||||
# Verify helper methods were called
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_end_user_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_servers,team_servers,expected_result,scenario",
|
||||
[
|
||||
# Test case 1: No end_user_id - should use key/team logic only
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2", "server3"],
|
||||
["server1", "server2"],
|
||||
"no_end_user_id_uses_key_team",
|
||||
),
|
||||
# Test case 2: No end_user_id with empty key
|
||||
(
|
||||
[],
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
"no_end_user_id_inherits_from_team",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_get_allowed_mcp_servers_without_end_user_id(
|
||||
self,
|
||||
key_servers,
|
||||
team_servers,
|
||||
expected_result,
|
||||
scenario,
|
||||
):
|
||||
"""Test get_allowed_mcp_servers when end_user_id is not set"""
|
||||
|
||||
# Create a mock user WITHOUT end_user_id
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id=None,
|
||||
)
|
||||
|
||||
# Mock the helper methods
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_end_user"
|
||||
) as mock_end_user_servers:
|
||||
# Set up return values
|
||||
mock_key_servers.return_value = key_servers
|
||||
mock_team_servers.return_value = team_servers
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# Assert the result (order-independent comparison)
|
||||
assert sorted(result) == sorted(expected_result)
|
||||
|
||||
# Verify helper methods were called correctly
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
# end_user method should NOT be called since no end_user_id
|
||||
mock_end_user_servers.assert_not_called()
|
||||
|
||||
async def test_get_allowed_mcp_servers_for_end_user_with_valid_user(self):
|
||||
"""Test _get_allowed_mcp_servers_for_end_user with valid end_user"""
|
||||
from litellm.proxy._types import (LiteLLM_EndUserTable,
|
||||
LiteLLM_ObjectPermissionTable)
|
||||
|
||||
# Create mock object permission
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-123",
|
||||
mcp_servers=["server1", "server2"],
|
||||
mcp_access_groups=["group1"],
|
||||
)
|
||||
|
||||
# Create mock end_user object
|
||||
mock_end_user = LiteLLM_EndUserTable(
|
||||
user_id="test-end-user",
|
||||
object_permission=mock_object_permission,
|
||||
object_permission_id="perm-123",
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the dependencies
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object") as mock_get_end_user:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
|
||||
) as mock_get_access_groups:
|
||||
# Configure mocks
|
||||
mock_get_end_user.return_value = mock_end_user
|
||||
mock_get_access_groups.return_value = ["group-server1", "group-server2"]
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_end_user(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
# Assert the result contains both direct and access group servers
|
||||
assert set(result) == {"server1", "server2", "group-server1", "group-server2"}
|
||||
|
||||
# Verify methods were called
|
||||
mock_get_end_user.assert_called_once()
|
||||
mock_get_access_groups.assert_called_once_with(["group1"])
|
||||
|
||||
async def test_get_allowed_mcp_servers_for_end_user_with_no_permission(self):
|
||||
"""Test _get_allowed_mcp_servers_for_end_user when end_user has no object_permission"""
|
||||
from litellm.proxy._types import LiteLLM_EndUserTable
|
||||
|
||||
# Create mock end_user object without object_permission
|
||||
mock_end_user = LiteLLM_EndUserTable(
|
||||
user_id="test-end-user",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the dependencies
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object") as mock_get_end_user:
|
||||
# Configure mock
|
||||
mock_get_end_user.return_value = mock_end_user
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_end_user(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
# Assert empty list is returned
|
||||
assert result == []
|
||||
|
||||
# Verify method was called
|
||||
mock_get_end_user.assert_called_once()
|
||||
|
||||
async def test_get_allowed_mcp_servers_for_end_user_without_end_user_id(self):
|
||||
"""Test _get_allowed_mcp_servers_for_end_user when no end_user_id is set"""
|
||||
|
||||
# Create mock user auth without end_user_id
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
end_user_id=None,
|
||||
)
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_end_user(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
# Assert empty list is returned
|
||||
assert result == []
|
||||
|
||||
async def test_get_allowed_mcp_servers_for_end_user_without_prisma_client(self):
|
||||
"""Test _get_allowed_mcp_servers_for_end_user when prisma_client is None"""
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock prisma_client as None
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_end_user(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
# Assert empty list is returned
|
||||
assert result == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestEndUserMCPPermissionsFallbackBehavior:
|
||||
"""Test suite for end_user MCP permission fallback behavior.
|
||||
|
||||
Tests the updated behavior where when end_user has NO explicit MCP permissions defined,
|
||||
the system falls back to key/team restrictions instead of blocking all access.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_servers,team_servers,end_user_servers,require_flag,expected_result,scenario",
|
||||
[
|
||||
# Test case 1: End user with NO permissions, require_flag=True - should fall back to key/team
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2", "server3"],
|
||||
[],
|
||||
True,
|
||||
["server1", "server2"],
|
||||
"no_end_user_perms_require_flag_true_fallback_to_key_team_intersection",
|
||||
),
|
||||
# Test case 2: End user with NO permissions, require_flag=False - should fall back to key/team
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2", "server3"],
|
||||
[],
|
||||
False,
|
||||
["server1", "server2"],
|
||||
"no_end_user_perms_require_flag_false_fallback_to_key_team",
|
||||
),
|
||||
# Test case 3: End user with NO permissions, only key has permissions - should use key permissions
|
||||
(
|
||||
["server1", "server2", "server3"],
|
||||
[],
|
||||
[],
|
||||
True,
|
||||
["server1", "server2", "server3"],
|
||||
"no_end_user_perms_only_key_has_permissions",
|
||||
),
|
||||
# Test case 4: End user with NO permissions, only team has permissions - should inherit from team
|
||||
(
|
||||
[],
|
||||
["team_server1", "team_server2"],
|
||||
[],
|
||||
True,
|
||||
["team_server1", "team_server2"],
|
||||
"no_end_user_perms_only_team_has_permissions",
|
||||
),
|
||||
# Test case 5: End user with NO permissions, both key and team empty - should return empty
|
||||
(
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
True,
|
||||
[],
|
||||
"no_permissions_anywhere_returns_empty",
|
||||
),
|
||||
# Test case 6: End user with NO permissions, key and team have different servers
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server3", "server4"],
|
||||
[],
|
||||
True,
|
||||
[],
|
||||
"no_end_user_perms_key_team_no_overlap_returns_empty",
|
||||
),
|
||||
# Test case 7: End user HAS permissions, require_flag=True - should still intersect
|
||||
(
|
||||
["server1", "server2", "server3"],
|
||||
["server1", "server2", "server3"],
|
||||
["server2", "server3", "server4"],
|
||||
True,
|
||||
["server2", "server3"],
|
||||
"end_user_has_perms_require_flag_true_still_intersects",
|
||||
),
|
||||
# Test case 8: End user HAS permissions but empty key/team - should use end_user permissions
|
||||
(
|
||||
[],
|
||||
[],
|
||||
["server1", "server2"],
|
||||
True,
|
||||
["server1", "server2"],
|
||||
"end_user_has_perms_empty_key_team_uses_end_user",
|
||||
),
|
||||
# Test case 9: End user NO permissions, complex key/team intersection
|
||||
(
|
||||
["server1", "server2", "server3", "server4"],
|
||||
["server2", "server3", "server5"],
|
||||
[],
|
||||
True,
|
||||
["server2", "server3"],
|
||||
"no_end_user_perms_complex_key_team_intersection",
|
||||
),
|
||||
# Test case 10: End user HAS permissions with no overlap - should return empty
|
||||
(
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
["server5", "server6"],
|
||||
True,
|
||||
[],
|
||||
"end_user_has_perms_no_overlap_returns_empty",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_end_user_permission_fallback_to_key_team(
|
||||
self,
|
||||
key_servers,
|
||||
team_servers,
|
||||
end_user_servers,
|
||||
require_flag,
|
||||
expected_result,
|
||||
scenario,
|
||||
):
|
||||
"""Test that end_user permissions fall back to key/team when not defined"""
|
||||
|
||||
# Create a mock user with end_user_id
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the helper methods
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_end_user"
|
||||
) as mock_end_user_servers:
|
||||
with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings:
|
||||
# Set up return values
|
||||
mock_key_servers.return_value = key_servers
|
||||
mock_team_servers.return_value = team_servers
|
||||
mock_end_user_servers.return_value = end_user_servers
|
||||
mock_general_settings.get.return_value = require_flag
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# Assert the result (order-independent comparison)
|
||||
assert sorted(result) == sorted(expected_result), (
|
||||
f"Test scenario '{scenario}' failed: "
|
||||
f"Expected {sorted(expected_result)}, got {sorted(result)}"
|
||||
)
|
||||
|
||||
# Verify helper methods were called
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_end_user_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
async def test_end_user_no_permissions_with_require_flag_true_detailed(self):
|
||||
"""
|
||||
Detailed test: When require_end_user_mcp_access_defined=True and end_user has NO permissions,
|
||||
should fall back to key/team restrictions (not block all access).
|
||||
|
||||
This is the main behavior change being tested.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_EndUserTable
|
||||
|
||||
# Create mock end_user object without object_permission
|
||||
mock_end_user = LiteLLM_EndUserTable(
|
||||
user_id="test-end-user",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the dependencies
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object") as mock_get_end_user:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings:
|
||||
# Configure mocks
|
||||
mock_get_end_user.return_value = mock_end_user
|
||||
mock_key_servers.return_value = ["server1", "server2"]
|
||||
mock_team_servers.return_value = ["server1", "server2", "server3"]
|
||||
mock_general_settings.get.return_value = True # require_flag=True
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# CRITICAL: Should fall back to key/team intersection, not block all
|
||||
assert sorted(result) == ["server1", "server2"], (
|
||||
"When end_user has NO permissions and require_flag=True, "
|
||||
"should fall back to key/team restrictions, not block all access"
|
||||
)
|
||||
|
||||
# Verify methods were called
|
||||
mock_get_end_user.assert_called_once()
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
async def test_end_user_no_permissions_without_require_flag(self):
|
||||
"""
|
||||
Test: When require_end_user_mcp_access_defined=False and end_user has NO permissions,
|
||||
should use key/team restrictions (default behavior).
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_EndUserTable
|
||||
|
||||
# Create mock end_user object without object_permission
|
||||
mock_end_user = LiteLLM_EndUserTable(
|
||||
user_id="test-end-user",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the dependencies
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object") as mock_get_end_user:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings:
|
||||
# Configure mocks
|
||||
mock_get_end_user.return_value = mock_end_user
|
||||
mock_key_servers.return_value = ["server1", "server2", "server3"]
|
||||
mock_team_servers.return_value = []
|
||||
mock_general_settings.get.return_value = False # require_flag=False
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# Should use key permissions when team is empty
|
||||
assert sorted(result) == ["server1", "server2", "server3"]
|
||||
|
||||
# Verify methods were called
|
||||
mock_get_end_user.assert_called_once()
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
async def test_end_user_with_permissions_still_enforces_intersection(self):
|
||||
"""
|
||||
Test: When end_user HAS explicit permissions, should still enforce intersection
|
||||
with key/team restrictions (behavior unchanged).
|
||||
"""
|
||||
from litellm.proxy._types import (LiteLLM_EndUserTable,
|
||||
LiteLLM_ObjectPermissionTable)
|
||||
|
||||
# Create mock object permission for end_user
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-123",
|
||||
mcp_servers=["server2", "server3", "server4"],
|
||||
mcp_access_groups=[],
|
||||
)
|
||||
|
||||
# Create mock end_user object WITH object_permission
|
||||
mock_end_user = LiteLLM_EndUserTable(
|
||||
user_id="test-end-user",
|
||||
object_permission=mock_object_permission,
|
||||
object_permission_id="perm-123",
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
end_user_id="test-end-user",
|
||||
)
|
||||
|
||||
# Mock the dependencies
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object") as mock_get_end_user:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
|
||||
) as mock_key_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
|
||||
) as mock_team_servers:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
|
||||
) as mock_access_groups:
|
||||
with patch("litellm.proxy.proxy_server.general_settings") as mock_general_settings:
|
||||
# Configure mocks
|
||||
mock_get_end_user.return_value = mock_end_user
|
||||
mock_key_servers.return_value = ["server1", "server2", "server3"]
|
||||
mock_team_servers.return_value = ["server1", "server2", "server3"]
|
||||
mock_access_groups.return_value = [] # No access group servers
|
||||
mock_general_settings.get.return_value = True
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
|
||||
# Should return intersection of key/team AND end_user
|
||||
# key/team intersection = ["server1", "server2", "server3"]
|
||||
# end_user = ["server2", "server3", "server4"]
|
||||
# final intersection = ["server2", "server3"]
|
||||
assert sorted(result) == ["server2", "server3"]
|
||||
|
||||
# Verify methods were called
|
||||
mock_get_end_user.assert_called_once()
|
||||
mock_key_servers.assert_called_once_with(mock_user_auth)
|
||||
mock_team_servers.assert_called_once_with(mock_user_auth)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue