fix: simplify testing

This commit is contained in:
Krrish Dholakia 2025-09-30 12:37:25 -07:00
parent 4a09507c58
commit 0cd61a6a6a
2 changed files with 74 additions and 113 deletions

View file

@ -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",

View file

@ -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