diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 5842ab2fe96..67cb783bbba 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -332,72 +332,23 @@ class MCPRequestHandler: user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[str]: """ - Get list of allowed MCP servers for the given user/key based on permissions + 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) Returns: List[str]: List of allowed MCP servers by server id """ from typing import List - from litellm.proxy.proxy_server import general_settings, prisma_client + from litellm.proxy.proxy_server import general_settings try: - # Check if end user MCP access enforcement is enabled - require_end_user_mcp_access = general_settings.get( - "require_end_user_mcp_access_defined", False - ) - - # If flag is enabled and this is an end_user request, check for explicit permissions - if ( - require_end_user_mcp_access - and user_api_key_auth - and user_api_key_auth.end_user_id - and prisma_client - ): - try: - # Fetch end user object with object_permission - end_user_obj = await prisma_client.db.litellm_endusertable.find_unique( - where={"user_id": user_api_key_auth.end_user_id}, - include={"object_permission": True}, - ) - - # If end user exists but has no object_permission defined, block all MCP access - if end_user_obj and end_user_obj.object_permission is None: - verbose_logger.debug( - f"require_end_user_mcp_access_defined=True and end_user {user_api_key_auth.end_user_id} has no object_permission - blocking MCP access" - ) - return [] - - # If end user has object_permission, check their allowed MCP servers - if end_user_obj and end_user_obj.object_permission: - end_user_mcp_servers = end_user_obj.object_permission.mcp_servers or [] - end_user_access_groups = end_user_obj.object_permission.mcp_access_groups or [] - - # Get servers from access groups - access_group_servers = ( - await MCPRequestHandler._get_mcp_servers_from_access_groups( - end_user_access_groups - ) - ) - - # Combine direct servers and access group servers for end user - end_user_allowed = list(set(end_user_mcp_servers + access_group_servers)) - - # If end user has explicit permissions, use only those - if len(end_user_allowed) > 0: - verbose_logger.debug( - f"require_end_user_mcp_access_defined=True - using end_user explicit permissions: {end_user_allowed}" - ) - return end_user_allowed - - except Exception as e: - verbose_logger.warning( - f"Failed to check end_user MCP permissions: {str(e)}" - ) - # On error, block access if flag is enabled - return [] - - allowed_mcp_servers: List[str] = [] + # Get allowed servers from key and team allowed_mcp_servers_for_key = ( await MCPRequestHandler._get_allowed_mcp_servers_for_key( user_api_key_auth @@ -410,8 +361,9 @@ class MCPRequestHandler: ) ######################################################### - # If team has mcp_servers, handle inheritance and intersection logic + # Calculate key/team allowed servers using inheritance and intersection logic ######################################################### + allowed_mcp_servers: List[str] = [] if len(allowed_mcp_servers_for_team) > 0: if len(allowed_mcp_servers_for_key) > 0: # Key has its own MCP permissions - use intersection with team permissions @@ -424,6 +376,51 @@ class MCPRequestHandler: else: allowed_mcp_servers = allowed_mcp_servers_for_key + ######################################################### + # Check end_user permissions if end_user_id is set + ######################################################### + if user_api_key_auth and user_api_key_auth.end_user_id: + allowed_mcp_servers_for_end_user = ( + await MCPRequestHandler._get_allowed_mcp_servers_for_end_user( + user_api_key_auth + ) + ) + + + # If end_user has explicit MCP server permissions, apply intersection + 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: + verbose_logger.debug( + "require_end_user_mcp_access_defined=True - using only end_user permissions" + ) + allowed_mcp_servers = allowed_mcp_servers_for_end_user + else: + # Otherwise, apply intersection: key/team AND end_user + 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}" + ) + # 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: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") @@ -664,6 +661,64 @@ class MCPRequestHandler: ) return [] + @staticmethod + async def _get_allowed_mcp_servers_for_end_user( + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + ) -> List[str]: + """ + Get allowed MCP servers for an end user. + + Returns the MCP servers from the end_user's object_permission. + """ + from litellm.proxy.auth.auth_checks import get_end_user_object + from litellm.proxy.proxy_server import (prisma_client, + proxy_logging_obj, + user_api_key_cache) + + if not user_api_key_auth or not user_api_key_auth.end_user_id: + return [] + + if prisma_client is None: + + verbose_logger.debug("prisma_client is None") + return [] + + try: + # Use optimized get_end_user_object function with caching + end_user_obj = await get_end_user_object( + end_user_id=user_api_key_auth.end_user_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, + route="/mcp", + ) + + + if end_user_obj is None or end_user_obj.object_permission is None: + return [] + + # Get direct MCP servers + direct_mcp_servers = end_user_obj.object_permission.mcp_servers or [] + + + + # Get MCP servers from access groups + access_group_servers = ( + await MCPRequestHandler._get_mcp_servers_from_access_groups( + end_user_obj.object_permission.mcp_access_groups or [] + ) + ) + + # Combine both lists + all_servers = direct_mcp_servers + access_group_servers + return list(set(all_servers)) + except Exception as e: + verbose_logger.warning( + f"Failed to get allowed MCP servers for end_user: {str(e)}" + ) + return [] + @staticmethod def _get_config_server_ids_for_access_groups( config_mcp_servers, access_groups: List[str] diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3eb6f28ddfc..5a8a9f4e4df 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -11,7 +11,8 @@ Run checks for: import asyncio import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, + cast) from fastapi import HTTPException, Request, status from pydantic import BaseModel @@ -20,41 +21,27 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.caching.dual_cache import LimitedSizeOrderedDict -from litellm.constants import ( - CLI_JWT_EXPIRATION_HOURS, - CLI_JWT_TOKEN_NAME, - DEFAULT_ACCESS_GROUP_CACHE_TTL, - DEFAULT_IN_MEMORY_TTL, - DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, - DEFAULT_MAX_RECURSE_DEPTH, - EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE, -) +from litellm.constants import (CLI_JWT_EXPIRATION_HOURS, CLI_JWT_TOKEN_NAME, + DEFAULT_ACCESS_GROUP_CACHE_TTL, + DEFAULT_IN_MEMORY_TTL, + DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, + DEFAULT_MAX_RECURSE_DEPTH, + EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE) from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider -from litellm.proxy._types import ( - RBAC_ROLES, - CallInfo, - LiteLLM_AccessGroupTable, - LiteLLM_BudgetTable, - LiteLLM_EndUserTable, - Litellm_EntityType, - LiteLLM_JWTAuth, - LiteLLM_ObjectPermissionTable, - LiteLLM_OrganizationMembershipTable, - LiteLLM_OrganizationTable, - LiteLLM_TagTable, - LiteLLM_TeamMembership, - LiteLLM_TeamTable, - LiteLLM_TeamTableCachedObj, - LiteLLM_UserTable, - LiteLLMRoutes, - LitellmUserRoles, - NewTeamRequest, - ProxyErrorTypes, - ProxyException, - RoleBasedPermissions, - SpecialModelNames, - UserAPIKeyAuth, -) +from litellm.proxy._types import (RBAC_ROLES, CallInfo, + LiteLLM_AccessGroupTable, + LiteLLM_BudgetTable, LiteLLM_EndUserTable, + Litellm_EntityType, LiteLLM_JWTAuth, + LiteLLM_ObjectPermissionTable, + LiteLLM_OrganizationMembershipTable, + LiteLLM_OrganizationTable, LiteLLM_TagTable, + LiteLLM_TeamMembership, LiteLLM_TeamTable, + LiteLLM_TeamTableCachedObj, + LiteLLM_UserTable, LiteLLMRoutes, + LitellmUserRoles, NewTeamRequest, + ProxyErrorTypes, ProxyException, + RoleBasedPermissions, SpecialModelNames, + UserAPIKeyAuth) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics @@ -366,7 +353,8 @@ async def common_checks( _request_metadata: dict = request_body.get("metadata", {}) or {} if _request_metadata.get("guardrails"): # check if team allowed to modify guardrails - from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails + from litellm.proxy.guardrails.guardrail_helpers import \ + can_modify_guardrails can_modify: bool = can_modify_guardrails(team_object) if can_modify is False: @@ -792,7 +780,7 @@ async def get_end_user_object( try: response = await prisma_client.db.litellm_endusertable.find_unique( where={"user_id": end_user_id}, - include={"litellm_budget_table": True}, + include={"litellm_budget_table": True, "object_permission": True}, ) if response is None: @@ -1812,9 +1800,8 @@ class ExperimentalUIJWTToken: def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str: from datetime import timedelta - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - encrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + encrypt_value_helper if user_info.user_role is None: raise Exception("User role is required for experimental UI login") @@ -1860,9 +1847,8 @@ class ExperimentalUIJWTToken: """ from datetime import timedelta - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - encrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + encrypt_value_helper if user_info.user_role is None: raise Exception("User role is required for CLI JWT login") @@ -1901,9 +1887,8 @@ class ExperimentalUIJWTToken: import json from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - decrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + decrypt_value_helper decrypted_token = decrypt_value_helper( hashed_token, key="ui_hash_key", exception_type="debug" @@ -2150,11 +2135,11 @@ async def _get_resources_from_access_groups( # Lazy import to avoid circular imports if prisma_client is None or user_api_key_cache is None: - from litellm.proxy.proxy_server import ( - prisma_client as _prisma_client, - proxy_logging_obj as _proxy_logging_obj, - user_api_key_cache as _user_api_key_cache, - ) + from litellm.proxy.proxy_server import prisma_client as _prisma_client + from litellm.proxy.proxy_server import \ + proxy_logging_obj as _proxy_logging_obj + from litellm.proxy.proxy_server import \ + user_api_key_cache as _user_api_key_cache prisma_client = prisma_client or _prisma_client user_api_key_cache = user_api_key_cache or _user_api_key_cache @@ -2936,7 +2921,8 @@ async def _tag_max_budget_check( BudgetExceededError if any tag is over its max budget. Triggers a budget alert if any tag is over its max budget. """ - from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import \ + get_tags_from_request_body if prisma_client is None: return diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index b9ac225c9eb..cd4aa9cbe91 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -22,8 +22,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_daily_activity import \ get_daily_activity from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, attach_object_permission_to_dict, - handle_update_object_permission_common) + _set_object_permission, handle_update_object_permission_common) from litellm.proxy.utils import handle_exception_on_proxy from litellm.types.proxy.management_endpoints.common_daily_activity import \ SpendAnalyticsPaginatedResponse @@ -166,6 +165,38 @@ def new_budget_request(data: NewCustomerRequest) -> Optional[BudgetNewRequest]: return None +async def _handle_customer_object_permission_update( + non_default_values: dict, + end_user_table_data_typed: Optional[LiteLLM_EndUserTable], + update_end_user_table_data: dict, + prisma_client, +) -> None: + """ + Handle object permission updates for customer endpoints. + + Updates the update_end_user_table_data dict in place with the new object_permission_id. + + Args: + non_default_values: Dictionary containing the update values including object_permission + end_user_table_data_typed: Existing end user table data + update_end_user_table_data: Dictionary to update with new object_permission_id + prisma_client: Prisma database client + """ + if "object_permission" in non_default_values: + existing_object_permission_id = ( + end_user_table_data_typed.object_permission_id + if end_user_table_data_typed is not None + else None + ) + object_permission_id = await handle_update_object_permission_common( + data_json=non_default_values, + existing_object_permission_id=existing_object_permission_id, + prisma_client=prisma_client, + ) + if object_permission_id is not None: + update_end_user_table_data["object_permission_id"] = object_permission_id + + @router.post( "/end_user/new", tags=["Customer Management"], @@ -499,19 +530,12 @@ async def update_end_user( update_end_user_table_data[k] = v ## Handle object permission updates (MCP servers, vector stores, etc.) - if "object_permission" in non_default_values: - existing_object_permission_id = ( - end_user_table_data_typed.object_permission_id - if end_user_table_data_typed is not None - else None - ) - object_permission_id = await handle_update_object_permission_common( - data_json=non_default_values, - existing_object_permission_id=existing_object_permission_id, - prisma_client=prisma_client, - ) - if object_permission_id is not None: - update_end_user_table_data["object_permission_id"] = object_permission_id + await _handle_customer_object_permission_update( + non_default_values=non_default_values, + end_user_table_data_typed=end_user_table_data_typed, + update_end_user_table_data=update_end_user_table_data, + prisma_client=prisma_client, + ) ## Check if we need to create a new budget (only if budget fields are provided, not just budget_id) ## if budget_table_data: