mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
MCP fixes
* fix(oldteams.tsx): show policies when creating * fix(proxy/_types.py): ensure mcp rest endpoints can be called by virtual key ensures UI works with virtual key testing mcp endpoints * refactor: migrate get object permissions table logic to happen in user api key auth - allows functions to trust user api key object they receive has what they need * fix(rest_endpoints.py): filter for allowed tools based on what key has access to * fix(mcp_server_manager.py): ensure only allowed MCP's are returned to the user, via rest endpoints
This commit is contained in:
parent
b019638716
commit
5736fd32d9
11 changed files with 818 additions and 241 deletions
|
|
@ -6,7 +6,12 @@ from starlette.requests import Request
|
|||
from starlette.types import Scope
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, ProxyException, SpecialHeaders, UserAPIKeyAuth
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTable,
|
||||
ProxyException,
|
||||
SpecialHeaders,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
|
||||
|
|
@ -372,45 +377,31 @@ class MCPRequestHandler:
|
|||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_key_object_permission(
|
||||
def _get_key_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get key object_permission from cache or DB."""
|
||||
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,
|
||||
)
|
||||
"""
|
||||
Get key object_permission - already loaded by get_key_object() in main auth flow.
|
||||
|
||||
Note: object_permission is automatically populated when the key is fetched via
|
||||
get_key_object() in litellm/proxy/auth/auth_checks.py
|
||||
"""
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if user_api_key_auth.object_permission:
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
# Need to fetch from DB
|
||||
if user_api_key_auth.object_permission_id and prisma_client:
|
||||
return 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,
|
||||
)
|
||||
|
||||
return None
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_team_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get team object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_object_permission,
|
||||
get_team_object,
|
||||
)
|
||||
"""
|
||||
Get team object_permission - automatically loaded by get_team_object() in main auth flow.
|
||||
|
||||
Note: object_permission is automatically populated when the team is fetched via
|
||||
get_team_object() in litellm/proxy/auth/auth_checks.py
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
|
|
@ -423,7 +414,7 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
# First get the team object (which may have object_permission already loaded)
|
||||
# Get the team object (which has object_permission already loaded)
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -435,21 +426,7 @@ class MCPRequestHandler:
|
|||
if not team_obj:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if team_obj.object_permission:
|
||||
return team_obj.object_permission
|
||||
|
||||
# Need to fetch from DB using object_permission_id
|
||||
if team_obj.object_permission_id:
|
||||
return await get_object_permission(
|
||||
object_permission_id=team_obj.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,
|
||||
)
|
||||
|
||||
return None
|
||||
return team_obj.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def get_allowed_tools_for_server(
|
||||
|
|
@ -471,8 +448,8 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
try:
|
||||
# Get key and team object permissions
|
||||
key_obj_perm = await MCPRequestHandler._get_key_object_permission(
|
||||
# Get key and team object permissions (already loaded in main auth flow)
|
||||
key_obj_perm = MCPRequestHandler._get_key_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
team_obj_perm = await MCPRequestHandler._get_team_object_permission(
|
||||
|
|
@ -559,7 +536,8 @@ class MCPRequestHandler:
|
|||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
try:
|
||||
key_object_permission = await MCPRequestHandler._get_key_object_permission(
|
||||
# Get key object permission (already loaded in main auth flow)
|
||||
key_object_permission = MCPRequestHandler._get_key_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
if key_object_permission is None:
|
||||
|
|
@ -591,12 +569,10 @@ class MCPRequestHandler:
|
|||
"""
|
||||
Get allowed MCP servers for a team.
|
||||
|
||||
Uses the helper _get_team_object_permission which:
|
||||
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
|
||||
Note: object_permission is automatically loaded by get_team_object() in main auth flow.
|
||||
"""
|
||||
try:
|
||||
# Use the helper method that properly handles fetching from DB if needed
|
||||
# Get team object permission (already loaded in main auth flow)
|
||||
object_permissions = await MCPRequestHandler._get_team_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
|
|
|||
|
|
@ -70,9 +70,7 @@ try:
|
|||
from mcp.shared.tool_name_validation import (
|
||||
validate_tool_name, # pyright: ignore[reportAssignmentType]
|
||||
)
|
||||
from mcp.shared.tool_name_validation import (
|
||||
SEP_986_URL,
|
||||
)
|
||||
from mcp.shared.tool_name_validation import SEP_986_URL
|
||||
except ImportError:
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -673,24 +671,47 @@ class MCPServerManager:
|
|||
return [
|
||||
server.server_id
|
||||
for server in self.get_registry().values()
|
||||
if server.allow_all_keys
|
||||
if server.allow_all_keys is True
|
||||
]
|
||||
|
||||
async def get_allowed_mcp_servers(
|
||||
self, user_api_key_auth: Optional[UserAPIKeyAuth] = None
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get the allowed MCP Servers for the user
|
||||
Get the allowed MCP Servers for the user.
|
||||
|
||||
Priority:
|
||||
1. If object_permission.mcp_servers is explicitly set, use it (even for admins)
|
||||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
|
||||
# If admin, get all servers
|
||||
if user_api_key_auth and _user_has_admin_view(user_api_key_auth):
|
||||
return list(self.get_registry().keys())
|
||||
|
||||
allow_all_server_ids = self.get_allow_all_keys_server_ids()
|
||||
|
||||
try:
|
||||
# Check if object_permission.mcp_servers is explicitly set
|
||||
has_explicit_object_permission = False
|
||||
if user_api_key_auth and user_api_key_auth.object_permission:
|
||||
# Check if mcp_servers is explicitly set (not None, empty list is valid)
|
||||
if user_api_key_auth.object_permission.mcp_servers is not None:
|
||||
has_explicit_object_permission = True
|
||||
verbose_logger.debug(
|
||||
f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}"
|
||||
)
|
||||
|
||||
# If admin but NO explicit object permission, get all servers
|
||||
if (
|
||||
user_api_key_auth
|
||||
and _user_has_admin_view(user_api_key_auth)
|
||||
and not has_explicit_object_permission
|
||||
):
|
||||
verbose_logger.debug(
|
||||
"Admin user without explicit object_permission - returning all servers"
|
||||
)
|
||||
return list(self.get_registry().keys())
|
||||
|
||||
# Get allowed servers from object permissions (respects object_permission even for admins)
|
||||
allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
|
@ -2246,6 +2267,7 @@ class MCPServerManager:
|
|||
from litellm.proxy.proxy_server import (
|
||||
general_settings as proxy_general_settings,
|
||||
)
|
||||
|
||||
return proxy_general_settings
|
||||
except ImportError:
|
||||
# Fallback if proxy_server not available
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
ListMCPToolsRestAPIResponseObject,
|
||||
MCPServer,
|
||||
_tool_name_matches,
|
||||
execute_mcp_tool,
|
||||
filter_tools_by_allowed_tools,
|
||||
)
|
||||
|
|
@ -159,6 +160,7 @@ if MCP_AVAILABLE:
|
|||
server,
|
||||
server_auth_header,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper function to get tools for a single server."""
|
||||
tools = await global_mcp_server_manager._get_tools_from_server(
|
||||
|
|
@ -173,6 +175,29 @@ if MCP_AVAILABLE:
|
|||
if server.allowed_tools is not None and len(server.allowed_tools) > 0:
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions
|
||||
# This provides per-key/team/org control over which tools can be accessed
|
||||
if (
|
||||
user_api_key_auth
|
||||
and user_api_key_auth.object_permission
|
||||
and user_api_key_auth.object_permission.mcp_tool_permissions
|
||||
):
|
||||
allowed_tools_for_server = (
|
||||
user_api_key_auth.object_permission.mcp_tool_permissions.get(
|
||||
server.server_id
|
||||
)
|
||||
)
|
||||
if (
|
||||
allowed_tools_for_server is not None
|
||||
and len(allowed_tools_for_server) > 0
|
||||
):
|
||||
# Filter tools to only include those in the allowed list
|
||||
tools = [
|
||||
tool
|
||||
for tool in tools
|
||||
if _tool_name_matches(tool.name, allowed_tools_for_server)
|
||||
]
|
||||
|
||||
return _create_tool_response_objects(tools, server.mcp_info)
|
||||
|
||||
async def _resolve_allowed_mcp_servers_for_tool_call(
|
||||
|
|
@ -197,9 +222,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
allowed_mcp_servers: List[MCPServer] = []
|
||||
for allowed_server_id in allowed_server_ids_set:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
allowed_server_id
|
||||
)
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
|
||||
if server is not None:
|
||||
allowed_mcp_servers.append(server)
|
||||
return allowed_mcp_servers
|
||||
|
|
@ -276,9 +299,7 @@ if MCP_AVAILABLE:
|
|||
"message": f"The key is not allowed to access server {server_id}",
|
||||
},
|
||||
)
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
server_id
|
||||
)
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
return {
|
||||
"tools": [],
|
||||
|
|
@ -292,7 +313,10 @@ if MCP_AVAILABLE:
|
|||
|
||||
try:
|
||||
list_tools_result = await _get_tools_for_single_server(
|
||||
server, server_auth_header, raw_headers_from_request
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers_from_request,
|
||||
user_api_key_dict,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
|
|
@ -328,7 +352,10 @@ if MCP_AVAILABLE:
|
|||
|
||||
try:
|
||||
tools_result = await _get_tools_for_single_server(
|
||||
server, server_auth_header, raw_headers_from_request
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers_from_request,
|
||||
user_api_key_dict,
|
||||
)
|
||||
list_tools_result.extend(tools_result)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -420,7 +420,8 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/mcp/tools",
|
||||
"/mcp/tools/list",
|
||||
"/mcp/tools/call",
|
||||
# Read-only MCP discovery endpoint (virtual keys may be allowed here)
|
||||
"/mcp-rest/tools/list",
|
||||
"/mcp-rest/tools/call",
|
||||
"/v1/mcp/server",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class AgentRequestHandler:
|
|||
1. Key-level agent permissions
|
||||
2. Team-level agent permissions
|
||||
3. Agent access group resolution
|
||||
|
||||
|
||||
Follows the same inheritance logic as MCP:
|
||||
- If team has restrictions and key has restrictions: use intersection
|
||||
- If team has restrictions and key has none: inherit from team
|
||||
|
|
@ -35,7 +35,7 @@ class AgentRequestHandler:
|
|||
) -> List[str]:
|
||||
"""
|
||||
Get list of allowed agent IDs for the given user/key based on permissions.
|
||||
|
||||
|
||||
Returns:
|
||||
List[str]: List of allowed agent IDs. Empty list means no restrictions (allow all).
|
||||
"""
|
||||
|
|
@ -45,7 +45,9 @@ class AgentRequestHandler:
|
|||
await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
|
||||
)
|
||||
allowed_agents_for_team = (
|
||||
await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
|
||||
await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
user_api_key_auth
|
||||
)
|
||||
)
|
||||
|
||||
# If team has agent restrictions, handle inheritance and intersection logic
|
||||
|
|
@ -73,62 +75,48 @@ class AgentRequestHandler:
|
|||
) -> bool:
|
||||
"""
|
||||
Check if a specific agent is allowed for the given user/key.
|
||||
|
||||
|
||||
Args:
|
||||
agent_id: The agent ID to check
|
||||
user_api_key_auth: User authentication info
|
||||
|
||||
|
||||
Returns:
|
||||
bool: True if agent is allowed, False otherwise
|
||||
"""
|
||||
allowed_agents = await AgentRequestHandler.get_allowed_agents(user_api_key_auth)
|
||||
|
||||
|
||||
# Empty list means no restrictions - allow all
|
||||
if len(allowed_agents) == 0:
|
||||
return True
|
||||
|
||||
|
||||
return agent_id in allowed_agents
|
||||
|
||||
@staticmethod
|
||||
async def _get_key_object_permission(
|
||||
def _get_key_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> Optional[LiteLLM_ObjectPermissionTable]:
|
||||
"""Helper to get key object_permission from cache or DB."""
|
||||
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,
|
||||
)
|
||||
"""
|
||||
Get key object_permission - already loaded by get_key_object() in main auth flow.
|
||||
|
||||
Note: object_permission is automatically populated when the key is fetched via
|
||||
get_key_object() in litellm/proxy/auth/auth_checks.py
|
||||
"""
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if user_api_key_auth.object_permission:
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
# Need to fetch from DB
|
||||
if user_api_key_auth.object_permission_id and prisma_client:
|
||||
return 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,
|
||||
)
|
||||
|
||||
return None
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_team_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> Optional[LiteLLM_ObjectPermissionTable]:
|
||||
"""Helper to get team object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_object_permission,
|
||||
get_team_object,
|
||||
)
|
||||
"""
|
||||
Get team object_permission - automatically loaded by get_team_object() in main auth flow.
|
||||
|
||||
Note: object_permission is automatically populated when the team is fetched via
|
||||
get_team_object() in litellm/proxy/auth/auth_checks.py
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
|
|
@ -138,7 +126,7 @@ class AgentRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
# First get the team object (which may have object_permission already loaded)
|
||||
# Get the team object (which has object_permission already loaded)
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -150,21 +138,7 @@ class AgentRequestHandler:
|
|||
if not team_obj:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if team_obj.object_permission:
|
||||
return team_obj.object_permission
|
||||
|
||||
# Need to fetch from DB using object_permission_id
|
||||
if team_obj.object_permission_id:
|
||||
return await get_object_permission(
|
||||
object_permission_id=team_obj.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,
|
||||
)
|
||||
|
||||
return None
|
||||
return team_obj.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_agents_for_key(
|
||||
|
|
@ -172,31 +146,16 @@ class AgentRequestHandler:
|
|||
) -> List[str]:
|
||||
"""
|
||||
Get allowed agents for a key from its object_permission.
|
||||
"""
|
||||
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,
|
||||
)
|
||||
|
||||
Note: object_permission is already loaded by get_key_object() in main auth flow.
|
||||
"""
|
||||
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,
|
||||
# Get key object permission (already loaded in main auth flow)
|
||||
key_object_permission = AgentRequestHandler._get_key_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
|
@ -205,8 +164,10 @@ class AgentRequestHandler:
|
|||
direct_agents = key_object_permission.agents or []
|
||||
|
||||
# Get agents from access groups
|
||||
access_group_agents = await AgentRequestHandler._get_agents_from_access_groups(
|
||||
key_object_permission.agent_access_groups or []
|
||||
access_group_agents = (
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
key_object_permission.agent_access_groups or []
|
||||
)
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
|
|
@ -222,6 +183,8 @@ class AgentRequestHandler:
|
|||
) -> List[str]:
|
||||
"""
|
||||
Get allowed agents for a team from its object_permission.
|
||||
|
||||
Note: object_permission is already loaded by get_team_object() in main auth flow.
|
||||
"""
|
||||
if user_api_key_auth is None:
|
||||
return []
|
||||
|
|
@ -230,7 +193,7 @@ class AgentRequestHandler:
|
|||
return []
|
||||
|
||||
try:
|
||||
# Use the helper method that properly handles fetching from DB if needed
|
||||
# Get team object permission (already loaded in main auth flow)
|
||||
object_permissions = await AgentRequestHandler._get_team_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
|
@ -242,8 +205,10 @@ class AgentRequestHandler:
|
|||
direct_agents = object_permissions.agents or []
|
||||
|
||||
# Get agents from access groups
|
||||
access_group_agents = await AgentRequestHandler._get_agents_from_access_groups(
|
||||
object_permissions.agent_access_groups or []
|
||||
access_group_agents = (
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
object_permissions.agent_access_groups or []
|
||||
)
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
|
|
@ -284,9 +249,7 @@ class AgentRequestHandler:
|
|||
for agent in agents:
|
||||
agent_ids.add(agent.agent_id)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Error getting agents from access groups: {e}"
|
||||
)
|
||||
verbose_logger.debug(f"Error getting agents from access groups: {e}")
|
||||
return agent_ids
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -306,16 +269,16 @@ class AgentRequestHandler:
|
|||
)
|
||||
|
||||
# Use the helper for DB agents
|
||||
db_agent_ids = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
db_agent_ids = (
|
||||
await AgentRequestHandler._get_db_agent_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
)
|
||||
)
|
||||
agent_ids.update(db_agent_ids)
|
||||
|
||||
return list(agent_ids)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to get agents from access groups: {str(e)}"
|
||||
)
|
||||
verbose_logger.warning(f"Failed to get agents from access groups: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -326,11 +289,15 @@ class AgentRequestHandler:
|
|||
Get list of agent access groups for the given user/key based on permissions.
|
||||
"""
|
||||
access_groups: List[str] = []
|
||||
access_groups_for_key = await AgentRequestHandler._get_agent_access_groups_for_key(
|
||||
user_api_key_auth
|
||||
access_groups_for_key = (
|
||||
await AgentRequestHandler._get_agent_access_groups_for_key(
|
||||
user_api_key_auth
|
||||
)
|
||||
)
|
||||
access_groups_for_team = await AgentRequestHandler._get_agent_access_groups_for_team(
|
||||
user_api_key_auth
|
||||
access_groups_for_team = (
|
||||
await AgentRequestHandler._get_agent_access_groups_for_team(
|
||||
user_api_key_auth
|
||||
)
|
||||
)
|
||||
|
||||
# If team has access groups, then key must have a subset of the team's access groups
|
||||
|
|
@ -378,7 +345,9 @@ class AgentRequestHandler:
|
|||
|
||||
return key_object_permission.agent_access_groups or []
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get agent access groups for key: {str(e)}")
|
||||
verbose_logger.warning(
|
||||
f"Failed to get agent access groups for key: {str(e)}"
|
||||
)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -425,4 +394,3 @@ class AgentRequestHandler:
|
|||
f"Failed to get agent access groups for team: {str(e)}"
|
||||
)
|
||||
return []
|
||||
|
||||
|
|
|
|||
|
|
@ -1368,6 +1368,22 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
raise Exception
|
||||
|
||||
_response = LiteLLM_TeamTableCachedObj(**response.dict())
|
||||
|
||||
# Load object_permission if object_permission_id exists but object_permission is not loaded
|
||||
if _response.object_permission_id and not _response.object_permission:
|
||||
try:
|
||||
_response.object_permission = await get_object_permission(
|
||||
object_permission_id=_response.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Failed to load object_permission for team {team_id} with object_permission_id={_response.object_permission_id}: {e}"
|
||||
)
|
||||
|
||||
# save the team object to cache
|
||||
await _cache_team_object(
|
||||
team_id=team_id,
|
||||
|
|
@ -1550,6 +1566,21 @@ async def get_team_object_by_alias(
|
|||
team = teams[0]
|
||||
team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump())
|
||||
|
||||
# Load object_permission if object_permission_id exists but object_permission is not loaded
|
||||
if team_obj.object_permission_id and not team_obj.object_permission:
|
||||
try:
|
||||
team_obj.object_permission = await get_object_permission(
|
||||
object_permission_id=team_obj.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Failed to load object_permission for team {team_obj.team_id} with object_permission_id={team_obj.object_permission_id}: {e}"
|
||||
)
|
||||
|
||||
# Cache the result by both alias and team_id
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
|
|
@ -1838,6 +1869,21 @@ async def get_key_object(
|
|||
|
||||
_response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True))
|
||||
|
||||
# Load object_permission if object_permission_id exists but object_permission is not loaded
|
||||
if _response.object_permission_id and not _response.object_permission:
|
||||
try:
|
||||
_response.object_permission = await get_object_permission(
|
||||
object_permission_id=_response.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Failed to load object_permission for key with object_permission_id={_response.object_permission_id}: {e}"
|
||||
)
|
||||
|
||||
# save the key object to cache
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
|
|
|
|||
|
|
@ -331,11 +331,7 @@ class TestMCPRequestHandler:
|
|||
# Create an async mock for user_api_key_auth
|
||||
async def mock_user_api_key_auth(api_key, request):
|
||||
return UserAPIKeyAuth(
|
||||
token=(
|
||||
"test-token-sha256-empty-hash"
|
||||
if api_key
|
||||
else None
|
||||
),
|
||||
token=("test-token-sha256-empty-hash" if api_key else None),
|
||||
api_key=api_key,
|
||||
user_id="test-user-id" if api_key else None,
|
||||
team_id="test-team-id" if api_key else None,
|
||||
|
|
@ -632,8 +628,7 @@ class TestMCPOAuth2AuthFlow:
|
|||
|
||||
# OAuth2 headers should still contain the Authorization token
|
||||
assert (
|
||||
oauth2_headers.get("Authorization")
|
||||
== "Bearer atlassian-oauth2-token"
|
||||
oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-token"
|
||||
)
|
||||
|
||||
async def test_litellm_key_in_authorization_backward_compat(self):
|
||||
|
|
@ -1291,21 +1286,21 @@ async def test_get_team_object_permission_with_already_loaded_permission():
|
|||
mcp_access_groups=["group1"],
|
||||
vector_stores=["store1"],
|
||||
)
|
||||
|
||||
|
||||
# Create mock team object with object_permission already loaded
|
||||
mock_team_obj = LiteLLM_TeamTable(
|
||||
team_id="team-123",
|
||||
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="team-123",
|
||||
)
|
||||
|
||||
|
||||
# Mock get_team_object to return our team with loaded permission
|
||||
# Also need to mock prisma_client from proxy_server
|
||||
mock_prisma = MagicMock()
|
||||
|
|
@ -1313,96 +1308,81 @@ async def test_get_team_object_permission_with_already_loaded_permission():
|
|||
"litellm.proxy.proxy_server.prisma_client",
|
||||
mock_prisma,
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object"
|
||||
) as mock_get_team:
|
||||
with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team:
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission"
|
||||
) as mock_get_perm:
|
||||
mock_get_team.return_value = mock_team_obj
|
||||
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_team_object_permission(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
|
||||
# Assert we got the object permission
|
||||
assert result == mock_object_permission
|
||||
assert result.mcp_servers == ["server1", "server2"]
|
||||
|
||||
|
||||
# Verify get_team_object was called
|
||||
mock_get_team.assert_called_once()
|
||||
|
||||
|
||||
# Verify get_object_permission was NOT called (since it was already loaded)
|
||||
mock_get_perm.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_permission_fetches_from_db_when_not_loaded():
|
||||
async def test_get_team_object_permission_with_core_auth_auto_loading():
|
||||
"""
|
||||
Test that _get_team_object_permission fetches from DB when object_permission
|
||||
is not loaded but object_permission_id exists.
|
||||
Test that _get_team_object_permission returns the object_permission that was
|
||||
automatically loaded by get_team_object() in the core auth flow.
|
||||
|
||||
Note: After migrating permission loading to core auth (get_team_object in auth_checks.py),
|
||||
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
|
||||
|
||||
# Create mock object permission (to be returned from DB)
|
||||
# Create mock object permission
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-456",
|
||||
mcp_servers=["server3", "server4"],
|
||||
mcp_access_groups=["group2"],
|
||||
vector_stores=["store2"],
|
||||
)
|
||||
|
||||
# Create mock team object WITHOUT object_permission loaded (but has ID)
|
||||
|
||||
# Create mock team object WITH object_permission already loaded
|
||||
# (This is what get_team_object() returns after the core auth migration)
|
||||
mock_team_obj = LiteLLM_TeamTable(
|
||||
team_id="team-456",
|
||||
object_permission=None,
|
||||
object_permission=mock_object_permission, # Already loaded by core auth
|
||||
object_permission_id="perm-456",
|
||||
)
|
||||
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="team-456",
|
||||
)
|
||||
|
||||
|
||||
# Mock the methods
|
||||
# Also need to mock prisma_client from proxy_server
|
||||
mock_prisma = MagicMock()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
mock_prisma,
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object"
|
||||
) as mock_get_team:
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission"
|
||||
) as mock_get_perm:
|
||||
mock_get_team.return_value = mock_team_obj
|
||||
mock_get_perm.return_value = mock_object_permission
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_team_object_permission(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
# Assert we got the object permission
|
||||
assert result == mock_object_permission
|
||||
assert result.mcp_servers == ["server3", "server4"]
|
||||
|
||||
# Verify get_team_object was called
|
||||
mock_get_team.assert_called_once()
|
||||
|
||||
# Verify get_object_permission WAS called (since it wasn't loaded)
|
||||
mock_get_perm.assert_called_once_with(
|
||||
object_permission_id="perm-456",
|
||||
prisma_client=mock.ANY,
|
||||
user_api_key_cache=mock.ANY,
|
||||
parent_otel_span=mock_user_auth.parent_otel_span,
|
||||
proxy_logging_obj=mock.ANY,
|
||||
)
|
||||
with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team:
|
||||
mock_get_team.return_value = mock_team_obj
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_team_object_permission(mock_user_auth)
|
||||
|
||||
# Assert we got the object permission (already loaded by core auth)
|
||||
assert result == mock_object_permission
|
||||
assert result.mcp_servers == ["server3", "server4"]
|
||||
|
||||
# Verify get_team_object was called
|
||||
mock_get_team.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1420,14 +1400,14 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
|
|||
mcp_access_groups=["dev-group"],
|
||||
vector_stores=[],
|
||||
)
|
||||
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="team-789",
|
||||
)
|
||||
|
||||
|
||||
# Mock the helper methods
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_team_object_permission"
|
||||
|
|
@ -1437,13 +1417,16 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
|
|||
) as mock_get_access_group_servers:
|
||||
# Configure mocks
|
||||
mock_get_team_perm.return_value = mock_object_permission
|
||||
mock_get_access_group_servers.return_value = ["group-server1", "group-server2"]
|
||||
|
||||
mock_get_access_group_servers.return_value = [
|
||||
"group-server1",
|
||||
"group-server2",
|
||||
]
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
|
||||
# Assert the result contains both direct and access group servers
|
||||
assert set(result) == {
|
||||
"direct-server1",
|
||||
|
|
@ -1451,10 +1434,10 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
|
|||
"group-server1",
|
||||
"group-server2",
|
||||
}
|
||||
|
||||
|
||||
# Verify _get_team_object_permission was called (the helper we fixed)
|
||||
mock_get_team_perm.assert_called_once_with(mock_user_auth)
|
||||
|
||||
|
||||
# Verify access groups were resolved
|
||||
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
|
||||
|
||||
|
|
@ -1471,21 +1454,21 @@ async def test_get_allowed_mcp_servers_for_team_with_no_object_permission():
|
|||
user_id="test-user",
|
||||
team_id="team-no-perm",
|
||||
)
|
||||
|
||||
|
||||
# Mock the helper to return None (no object permission)
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_team_object_permission"
|
||||
) as mock_get_team_perm:
|
||||
mock_get_team_perm.return_value = None
|
||||
|
||||
|
||||
# Call the method
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
|
||||
mock_user_auth
|
||||
)
|
||||
|
||||
|
||||
# Assert empty list is returned
|
||||
assert result == []
|
||||
|
||||
|
||||
# Verify the helper was called
|
||||
mock_get_team_perm.assert_called_once_with(mock_user_auth)
|
||||
|
||||
|
|
@ -1509,9 +1492,7 @@ async def test_get_allowed_mcp_servers_for_team_without_team_id_returns_empty():
|
|||
team_id=None,
|
||||
)
|
||||
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
|
||||
mock_user_auth
|
||||
)
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(mock_user_auth)
|
||||
|
||||
assert result == []
|
||||
|
||||
|
|
@ -1546,9 +1527,7 @@ async def test_get_allowed_mcp_servers_for_key_guard_conditions(
|
|||
"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
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma_client_value):
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
|
@ -1569,9 +1548,7 @@ async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_non
|
|||
|
||||
mock_prisma = object()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma
|
||||
), patch(
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -70,7 +70,9 @@ def _route_has_dependency(route, dependency) -> bool:
|
|||
dependant = getattr(route, "dependant", None)
|
||||
if dependant is None:
|
||||
return False
|
||||
return any(getattr(dep, "call", None) == dependency for dep in dependant.dependencies)
|
||||
return any(
|
||||
getattr(dep, "call", None) == dependency for dep in dependant.dependencies
|
||||
)
|
||||
|
||||
|
||||
class TestExecuteWithMcpClient:
|
||||
|
|
@ -481,7 +483,9 @@ class TestListToolsRestAPI:
|
|||
|
||||
captured = {"called": False}
|
||||
|
||||
async def fake_get_tools(server, server_auth_header, raw_headers=None):
|
||||
async def fake_get_tools(
|
||||
server, server_auth_header, raw_headers=None, user_api_key_auth=None
|
||||
):
|
||||
captured["called"] = True
|
||||
captured["server"] = server
|
||||
captured["auth_header"] = server_auth_header
|
||||
|
|
@ -659,3 +663,293 @@ class TestCallToolRestAPI:
|
|||
assert captured["name"] == "demo-tool"
|
||||
assert captured["arguments"] == {"foo": "bar"}
|
||||
assert captured["allowed_mcp_servers"] == [stub_server]
|
||||
|
||||
|
||||
class TestGetToolsForSingleServer:
|
||||
"""Test _get_tools_for_single_server with object_permission filtering"""
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
async def test_filters_tools_by_object_permission_mcp_tool_permissions(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Test that tools are filtered by user_api_key_auth.object_permission.mcp_tool_permissions"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
# Create mock tools
|
||||
class MockTool:
|
||||
def __init__(self, name, description):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [
|
||||
MockTool("tool1", "First tool"),
|
||||
MockTool("tool2", "Second tool"),
|
||||
MockTool("tool3", "Third tool"),
|
||||
]
|
||||
|
||||
# Mock _get_tools_from_server to return all tools
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
# Create server
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=None, # No server-level filtering
|
||||
)
|
||||
|
||||
# Create UserAPIKeyAuth with object_permission
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test-permission-id",
|
||||
mcp_tool_permissions={"test-server-id": ["tool1", "tool3"]},
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
# Call the function
|
||||
result = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify only allowed tools are returned
|
||||
assert len(result) == 2
|
||||
tool_names = [tool.name for tool in result]
|
||||
assert "tool1" in tool_names
|
||||
assert "tool3" in tool_names
|
||||
assert "tool2" not in tool_names
|
||||
|
||||
async def test_no_filtering_when_object_permission_is_none(self, monkeypatch):
|
||||
"""Test that all tools are returned when object_permission is None"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name, description):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [
|
||||
MockTool("tool1", "First tool"),
|
||||
MockTool("tool2", "Second tool"),
|
||||
]
|
||||
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=None,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
|
||||
# All tools should be returned
|
||||
assert len(result) == 2
|
||||
|
||||
async def test_no_filtering_when_mcp_tool_permissions_is_none(self, monkeypatch):
|
||||
"""Test that all tools are returned when mcp_tool_permissions is None"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name, description):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [
|
||||
MockTool("tool1", "First tool"),
|
||||
MockTool("tool2", "Second tool"),
|
||||
]
|
||||
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test-permission-id",
|
||||
mcp_tool_permissions=None, # No tool permissions set
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
|
||||
# All tools should be returned
|
||||
assert len(result) == 2
|
||||
|
||||
async def test_no_filtering_when_server_not_in_mcp_tool_permissions(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Test that all tools are returned when server is not in mcp_tool_permissions"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name, description):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [
|
||||
MockTool("tool1", "First tool"),
|
||||
MockTool("tool2", "Second tool"),
|
||||
]
|
||||
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test-permission-id",
|
||||
mcp_tool_permissions={"other-server-id": ["tool1"]}, # Different server
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
|
||||
# All tools should be returned since server is not in permissions
|
||||
assert len(result) == 2
|
||||
|
||||
async def test_combines_server_allowed_tools_and_object_permission_filters(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Test that both server.allowed_tools and object_permission.mcp_tool_permissions filters are applied"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name, description):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [
|
||||
MockTool("tool1", "First tool"),
|
||||
MockTool("tool2", "Second tool"),
|
||||
MockTool("tool3", "Third tool"),
|
||||
MockTool("tool4", "Fourth tool"),
|
||||
]
|
||||
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
# Server allows tool1, tool2, tool3
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=["tool1", "tool2", "tool3"],
|
||||
)
|
||||
|
||||
# Object permission allows tool2, tool3, tool4
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test-permission-id",
|
||||
mcp_tool_permissions={"test-server-id": ["tool2", "tool3", "tool4"]},
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Only tools in both lists should be returned (intersection): tool2, tool3
|
||||
assert len(result) == 2
|
||||
tool_names = [tool.name for tool in result]
|
||||
assert "tool2" in tool_names
|
||||
assert "tool3" in tool_names
|
||||
assert "tool1" not in tool_names
|
||||
assert "tool4" not in tool_names
|
||||
|
|
|
|||
151
tests/test_litellm/proxy/auth/test_object_permission_loading.py
Normal file
151
tests/test_litellm/proxy/auth/test_object_permission_loading.py
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
"""
|
||||
Test that object_permission is automatically loaded when fetching keys and teams.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import get_key_object, get_team_object
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_loads_object_permission():
|
||||
"""
|
||||
Test that get_key_object automatically loads object_permission when object_permission_id exists.
|
||||
"""
|
||||
# Mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache
|
||||
|
||||
# Mock the DB response with object_permission_id but no object_permission
|
||||
mock_token_data = MagicMock()
|
||||
mock_token_data.model_dump.return_value = {
|
||||
"token": "test_token_hash",
|
||||
"user_id": "test_user",
|
||||
"object_permission_id": "test_perm_id",
|
||||
"object_permission": None,
|
||||
}
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data)
|
||||
|
||||
# Mock the object_permission that should be loaded
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test_perm_id",
|
||||
mcp_servers=["server1", "server2"],
|
||||
vector_stores=["store1"],
|
||||
)
|
||||
|
||||
# Mock get_object_permission to return the permission
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
AsyncMock(return_value=mock_object_permission)
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks._cache_key_object",
|
||||
AsyncMock()
|
||||
):
|
||||
result = await get_key_object(
|
||||
hashed_token="test_token_hash",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
# Verify that object_permission was loaded
|
||||
assert result.object_permission is not None
|
||||
assert result.object_permission.object_permission_id == "test_perm_id"
|
||||
assert result.object_permission.mcp_servers == ["server1", "server2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_no_permission_id():
|
||||
"""
|
||||
Test that get_key_object works correctly when no object_permission_id exists.
|
||||
"""
|
||||
# Mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache
|
||||
|
||||
# Mock the DB response without object_permission_id
|
||||
mock_token_data = MagicMock()
|
||||
mock_token_data.model_dump.return_value = {
|
||||
"token": "test_token_hash",
|
||||
"user_id": "test_user",
|
||||
"object_permission_id": None,
|
||||
"object_permission": None,
|
||||
}
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._cache_key_object",
|
||||
AsyncMock()
|
||||
):
|
||||
result = await get_key_object(
|
||||
hashed_token="test_token_hash",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
# Verify that object_permission is None
|
||||
assert result.object_permission is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_loads_object_permission():
|
||||
"""
|
||||
Test that get_team_object automatically loads object_permission when object_permission_id exists.
|
||||
"""
|
||||
# Mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache
|
||||
|
||||
# Mock team data with object_permission_id
|
||||
mock_team = MagicMock()
|
||||
mock_team.dict.return_value = {
|
||||
"team_id": "test_team",
|
||||
"team_alias": "Test Team",
|
||||
"object_permission_id": "test_perm_id",
|
||||
"object_permission": None,
|
||||
}
|
||||
|
||||
# Mock the object_permission that should be loaded
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test_perm_id",
|
||||
mcp_servers=["team_server1"],
|
||||
vector_stores=["team_store1"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_team_db_check",
|
||||
AsyncMock(return_value=mock_team)
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
AsyncMock(return_value=mock_object_permission)
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks._cache_team_object",
|
||||
AsyncMock()
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks._should_check_db",
|
||||
return_value=True
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks._update_last_db_access_time"
|
||||
):
|
||||
result = await get_team_object(
|
||||
team_id="test_team",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
# Verify that object_permission was loaded
|
||||
assert result.object_permission is not None
|
||||
assert result.object_permission.object_permission_id == "test_perm_id"
|
||||
assert result.object_permission.mcp_servers == ["team_server1"]
|
||||
|
|
@ -1,14 +1,15 @@
|
|||
import os
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy.management_endpoints import (
|
||||
mcp_management_endpoints as mgmt_endpoints,
|
||||
|
|
@ -624,6 +625,72 @@ class TestListMCPServers:
|
|||
assert server.alias == "Allowed Zapier MCP"
|
||||
assert server.url == "https://actions.zapier.com/mcp/sse"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_user_with_object_permission_respects_mcp_servers(self):
|
||||
"""
|
||||
Test that admin users with explicit object_permission.mcp_servers
|
||||
only see the servers specified in object_permission.
|
||||
|
||||
Scenario: Admin user has object_permission.mcp_servers set to specific servers
|
||||
Expected: Only those servers are returned, not all servers in the registry
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
||||
# Create mock object permission with specific servers
|
||||
mock_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="test-obj-perm-id",
|
||||
mcp_servers=["server-1", "server-2"], # Only these two servers
|
||||
mcp_access_groups=[],
|
||||
mcp_tool_permissions={},
|
||||
vector_stores=[],
|
||||
agents=[],
|
||||
agent_access_groups=[],
|
||||
)
|
||||
|
||||
# Create admin user with object permission
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin_user_id",
|
||||
api_key="admin_api_key",
|
||||
object_permission=mock_object_permission,
|
||||
object_permission_id="test-obj-perm-id",
|
||||
)
|
||||
|
||||
# Mock servers that the user should see
|
||||
server_1 = generate_mock_mcp_server_db_record(
|
||||
server_id="server-1", alias="Server 1", url="https://server1.example.com"
|
||||
)
|
||||
server_2 = generate_mock_mcp_server_db_record(
|
||||
server_id="server-2", alias="Server 2", url="https://server2.example.com"
|
||||
)
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_all_allowed_mcp_servers = AsyncMock(
|
||||
return_value=[server_1, server_2]
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
|
||||
AsyncMock(return_value=[mock_user_auth]),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_all_mcp_servers,
|
||||
)
|
||||
|
||||
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
|
||||
|
||||
# Verify results - should only return the 2 servers in object_permission
|
||||
assert len(result) == 2
|
||||
server_ids = {server.server_id for server in result}
|
||||
assert server_ids == {"server-1", "server-2"}
|
||||
|
||||
# Verify credentials are redacted
|
||||
assert all(server.credentials is None for server in result)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_redacts_credentials(self):
|
||||
mock_server = generate_mock_mcp_server_db_record(
|
||||
|
|
@ -1251,7 +1318,9 @@ class TestMCPRegistryEndpoint:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_registry.return_value = {mock_server.server_id: mock_server}
|
||||
# The registry endpoint uses get_filtered_registry (filters by client IP)
|
||||
mock_manager.get_filtered_registry.return_value = {mock_server.server_id: mock_server}
|
||||
mock_manager.get_filtered_registry.return_value = {
|
||||
mock_server.server_id: mock_server
|
||||
}
|
||||
|
||||
with (
|
||||
patch_proxy_general_settings({"enable_mcp_registry": True}),
|
||||
|
|
|
|||
|
|
@ -52,7 +52,7 @@ import type { KeyResponse, Team } from "./key_team_helpers/key_list";
|
|||
import MCPServerSelector from "./mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "./mcp_server_management/MCPToolPermissions";
|
||||
import NotificationsManager from "./molecules/notifications_manager";
|
||||
import { Organization, fetchMCPAccessGroups, getGuardrailsList, teamDeleteCall } from "./networking";
|
||||
import { Organization, fetchMCPAccessGroups, getGuardrailsList, getPoliciesList, teamDeleteCall } from "./networking";
|
||||
import NumericalInput from "./shared/numerical_input";
|
||||
import VectorStoreSelector from "./vector_store_management/VectorStoreSelector";
|
||||
|
||||
|
|
@ -223,6 +223,7 @@ const Teams: React.FC<TeamProps> = ({
|
|||
const [isTeamDeleting, setIsTeamDeleting] = useState(false);
|
||||
// Add this state near the other useState declarations
|
||||
const [guardrailsList, setGuardrailsList] = useState<string[]>([]);
|
||||
const [policiesList, setPoliciesList] = useState<string[]>([]);
|
||||
const [expandedAccordions, setExpandedAccordions] = useState<Record<string, boolean>>({});
|
||||
const [loggingSettings, setLoggingSettings] = useState<any[]>([]);
|
||||
const [mcpAccessGroups, setMcpAccessGroups] = useState<string[]>([]);
|
||||
|
|
@ -273,7 +274,22 @@ const Teams: React.FC<TeamProps> = ({
|
|||
}
|
||||
};
|
||||
|
||||
const fetchPolicies = async () => {
|
||||
try {
|
||||
if (accessToken == null) {
|
||||
return;
|
||||
}
|
||||
|
||||
const response = await getPoliciesList(accessToken);
|
||||
const policyNames = response.policies.map((p: { policy_name: string }) => p.policy_name);
|
||||
setPoliciesList(policyNames);
|
||||
} catch (error) {
|
||||
console.error("Failed to fetch policies:", error);
|
||||
}
|
||||
};
|
||||
|
||||
fetchGuardrails();
|
||||
fetchPolicies();
|
||||
}, [accessToken]);
|
||||
|
||||
const fetchMcpAccessGroups = async () => {
|
||||
|
|
@ -1330,6 +1346,36 @@ const Teams: React.FC<TeamProps> = ({
|
|||
}
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Policies{" "}
|
||||
<Tooltip title="Apply policies to this team to control guardrails and other settings">
|
||||
<a
|
||||
href="https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies"
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</a>
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="policies"
|
||||
className="mt-8"
|
||||
help="Select existing policies or enter new ones"
|
||||
>
|
||||
<Select2
|
||||
mode="tags"
|
||||
style={{ width: "100%" }}
|
||||
placeholder="Select or enter policies"
|
||||
options={policiesList.map((name) => ({
|
||||
value: name,
|
||||
label: name,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue