mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: simplify testing
This commit is contained in:
parent
4a09507c58
commit
0cd61a6a6a
2 changed files with 74 additions and 113 deletions
|
|
@ -24,98 +24,80 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
@pytest.mark.asyncio
|
||||
class TestMCPRequestHandler:
|
||||
@pytest.mark.parametrize(
|
||||
"user_api_key_auth,object_permission_id,prisma_client_available,db_result,expected_result",
|
||||
"key_servers,team_servers,expected_result,scenario",
|
||||
[
|
||||
# Test case 1: user_api_key_auth is None
|
||||
(None, None, True, None, []),
|
||||
# Test case 2: object_permission_id is None
|
||||
(UserAPIKeyAuth(), None, True, None, []),
|
||||
# Test case 3: prisma_client is None
|
||||
# Test case 1: No key servers, no team servers
|
||||
([], [], [], "no_permissions"),
|
||||
# Test case 2: Key has servers, no team servers
|
||||
(["server1", "server2"], [], ["server1", "server2"], "key_only"),
|
||||
# Test case 3: No key servers, team has servers (inherit from team)
|
||||
(
|
||||
UserAPIKeyAuth(object_permission_id="test-id"),
|
||||
"test-id",
|
||||
False,
|
||||
None,
|
||||
[],
|
||||
["team_server1", "team_server2"],
|
||||
["team_server1", "team_server2"],
|
||||
"inherit_from_team",
|
||||
),
|
||||
# Test case 4: Database query returns None
|
||||
(UserAPIKeyAuth(object_permission_id="test-id"), "test-id", True, None, []),
|
||||
# Test case 5: Database query returns object with mcp_servers
|
||||
# Test case 4: Key and team both have servers (intersection)
|
||||
(
|
||||
UserAPIKeyAuth(object_permission_id="test-id"),
|
||||
"test-id",
|
||||
True,
|
||||
MagicMock(mcp_servers=["server1", "server2"]),
|
||||
["server1", "server2"],
|
||||
["server1", "team_server"],
|
||||
["server1"],
|
||||
"intersection",
|
||||
),
|
||||
# Test case 6: Database query returns object with None mcp_servers
|
||||
# Test case 5: Key and team have no overlap (empty result)
|
||||
(
|
||||
UserAPIKeyAuth(object_permission_id="test-id"),
|
||||
"test-id",
|
||||
True,
|
||||
MagicMock(mcp_servers=None),
|
||||
["server1", "server2"],
|
||||
["team_server1", "team_server2"],
|
||||
[],
|
||||
"no_overlap",
|
||||
),
|
||||
# Test case 7: Database query returns object with empty mcp_servers
|
||||
# Test case 6: Key and team have complete overlap
|
||||
(
|
||||
UserAPIKeyAuth(object_permission_id="test-id"),
|
||||
"test-id",
|
||||
True,
|
||||
MagicMock(mcp_servers=[]),
|
||||
[],
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
["server1", "server2"],
|
||||
"complete_overlap",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_get_allowed_mcp_servers_for_key(
|
||||
async def test_get_allowed_mcp_servers(
|
||||
self,
|
||||
user_api_key_auth,
|
||||
object_permission_id,
|
||||
prisma_client_available,
|
||||
db_result,
|
||||
key_servers,
|
||||
team_servers,
|
||||
expected_result,
|
||||
scenario,
|
||||
):
|
||||
"""Test _get_allowed_mcp_servers_for_key with various scenarios"""
|
||||
"""Test get_allowed_mcp_servers with various key/team permission scenarios"""
|
||||
|
||||
# Setup user_api_key_auth object_permission_id if provided
|
||||
if user_api_key_auth and object_permission_id:
|
||||
user_api_key_auth.object_permission_id = object_permission_id
|
||||
# Create a mock user
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
)
|
||||
|
||||
# Mock prisma_client
|
||||
mock_prisma_client = MagicMock() if prisma_client_available else None
|
||||
mock_find_unique = None
|
||||
# Mock the helper methods instead of database calls
|
||||
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:
|
||||
# Set up return values
|
||||
mock_key_servers.return_value = key_servers
|
||||
mock_team_servers.return_value = team_servers
|
||||
|
||||
if mock_prisma_client:
|
||||
# Mock the database query
|
||||
mock_find_unique = AsyncMock(return_value=db_result)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = (
|
||||
mock_find_unique
|
||||
)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
||||
# Assert the result (order-independent comparison)
|
||||
assert sorted(result) == sorted(expected_result)
|
||||
|
||||
# Verify database call was made correctly when expected
|
||||
if (
|
||||
user_api_key_auth
|
||||
and user_api_key_auth.object_permission_id
|
||||
and prisma_client_available
|
||||
and mock_find_unique
|
||||
):
|
||||
mock_find_unique.assert_called_once_with(
|
||||
where={
|
||||
"object_permission_id": user_api_key_auth.object_permission_id
|
||||
}
|
||||
# Call the method
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
elif mock_find_unique:
|
||||
# If prisma_client exists but conditions aren't met, no call should be made
|
||||
if not user_api_key_auth or not user_api_key_auth.object_permission_id:
|
||||
mock_find_unique.assert_not_called()
|
||||
|
||||
# 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)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"team_servers,key_servers,expected_servers,scenario",
|
||||
|
|
|
|||
|
|
@ -689,9 +689,6 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
"""Test that a user cannot call a tool from a server they don't have access to"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server import call_mcp_tool
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -703,45 +700,27 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
object_permission_id="key-permission-123",
|
||||
)
|
||||
|
||||
# Mock the database calls that determine access permissions
|
||||
# Mock get_object_permission to return no MCP servers for the key
|
||||
# Mock global_mcp_server_manager.get_mcp_server_names_from_ids to return
|
||||
# a list that doesn't include "restricted_server" (the server the user is trying to access)
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission"
|
||||
) as mock_get_object_permission:
|
||||
# Mock get_team_object to return no MCP access for the team
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object"
|
||||
) as mock_get_team_object:
|
||||
# Mock object permission - key has no MCP server access
|
||||
mock_key_permission = MagicMock()
|
||||
mock_key_permission.mcp_servers = [] # No direct server access
|
||||
mock_key_permission.mcp_access_groups = [] # No access groups
|
||||
mock_get_object_permission.return_value = mock_key_permission
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_names_from_ids"
|
||||
) as mock_get_server_names:
|
||||
# User has access to "allowed_server" but not "restricted_server"
|
||||
mock_get_server_names.return_value = ["allowed_server", "another_server"]
|
||||
|
||||
# Mock team object - team also has no MCP access
|
||||
mock_team = MagicMock()
|
||||
mock_team.object_permission = None # Team has no MCP permissions
|
||||
mock_get_team_object.return_value = mock_team
|
||||
# Try to call a tool from "restricted_server" - should raise HTTPException with 403 status
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await call_mcp_tool(
|
||||
name="restricted_server-send_email",
|
||||
arguments={
|
||||
"to": "test@example.com",
|
||||
"subject": "Test",
|
||||
"body": "Test",
|
||||
},
|
||||
user_api_key_auth=mock_user_auth,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
)
|
||||
|
||||
# Mock _get_mcp_servers_from_access_groups to return empty list
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
|
||||
) as mock_get_servers_from_groups:
|
||||
mock_get_servers_from_groups.return_value = []
|
||||
|
||||
# Try to call a tool - should raise HTTPException with 403 status
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await call_mcp_tool(
|
||||
name="restricted_server-send_email",
|
||||
arguments={
|
||||
"to": "test@example.com",
|
||||
"subject": "Test",
|
||||
"body": "Test",
|
||||
},
|
||||
user_api_key_auth=mock_user_auth,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
)
|
||||
|
||||
# Verify the exception details
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "User not allowed to call this tool" in exc_info.value.detail
|
||||
# Verify the exception details
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "User not allowed to call this tool" in exc_info.value.detail
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue