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 ab9ac0acd9b..9e316289c87 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 @@ -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", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index a2cee7b7d3c..1e9c63e2eb2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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