From f365402ccaa3f570c8cb68787bdee091e4ea7970 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 14 Feb 2026 19:27:51 -0800 Subject: [PATCH] feat: cleanup for customer_endpoints support of object permission id --- .../mcp_server/auth/user_api_key_auth_mcp.py | 105 +++++++++++++----- .../customer_endpoints.py | 56 +++++++++- 2 files changed, 126 insertions(+), 35 deletions(-) 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 ed4fb133478..5842ab2fe96 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 @@ -6,12 +6,8 @@ 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 @@ -343,7 +339,64 @@ class MCPRequestHandler: """ from typing import List + from litellm.proxy.proxy_server import general_settings, prisma_client + 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] = [] allowed_mcp_servers_for_key = ( await MCPRequestHandler._get_allowed_mcp_servers_for_key( @@ -402,11 +455,9 @@ class MCPRequestHandler: 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, - user_api_key_cache, - ) + from litellm.proxy.proxy_server import (prisma_client, + proxy_logging_obj, + user_api_key_cache) verbose_logger.debug( f"MCP team permission lookup: team_id={user_api_key_auth.team_id if user_api_key_auth else None}" @@ -541,12 +592,11 @@ class MCPRequestHandler: user_api_key_auth ) if key_object_permission is None and user_api_key_auth and user_api_key_auth.object_permission_id: - 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, - ) + 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) if prisma_client is not None: key_object_permission = await get_object_permission( object_permission_id=user_api_key_auth.object_permission_id, @@ -660,9 +710,8 @@ class MCPRequestHandler: try: # Import here to avoid circular import - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import \ + global_mcp_server_manager # Use the new helper for config-loaded servers server_ids = MCPRequestHandler._get_config_server_ids_for_access_groups( @@ -718,11 +767,9 @@ class MCPRequestHandler: user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[str]: 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, - ) + from litellm.proxy.proxy_server import (prisma_client, + proxy_logging_obj, + user_api_key_cache) if user_api_key_auth is None: return [] @@ -758,11 +805,9 @@ class MCPRequestHandler: Get MCP access groups for the team """ from litellm.proxy.auth.auth_checks import get_team_object - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + from litellm.proxy.proxy_server import (prisma_client, + proxy_logging_obj, + user_api_key_cache) if user_api_key_auth is None: return [] diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 025b8194803..b9ac225c9eb 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -22,7 +22,8 @@ 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, handle_update_object_permission_common) + _set_object_permission, attach_object_permission_to_dict, + 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 @@ -294,13 +295,28 @@ async def new_end_user( prisma_client=prisma_client, ) + # Ensure object_permission is not in the data being sent to create + # It should have been converted to object_permission_id by _set_object_permission + if "object_permission" in new_end_user_obj: + verbose_proxy_logger.warning( + f"object_permission still in new_end_user_obj after _set_object_permission: {new_end_user_obj.get('object_permission')}" + ) + new_end_user_obj.pop("object_permission", None) + ## WRITE TO DB ## end_user_record = await prisma_client.db.litellm_endusertable.create( data=new_end_user_obj, # type: ignore include={"litellm_budget_table": True, "object_permission": True}, ) - return end_user_record + # Convert to dict and clean up recursive fields + response_dict = end_user_record.model_dump() + if response_dict.get("object_permission"): + # Remove reverse relations from object_permission + for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]: + response_dict["object_permission"].pop(field, None) + + return response_dict except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.management_endpoints.customer_endpoints.new_end_user(): Exception occured - {}".format( @@ -366,7 +382,15 @@ async def end_user_info( code=404, param="end_user_id", ) - return user_info.model_dump(exclude_none=True) + + # Convert to dict and clean up recursive fields + response_dict = user_info.model_dump(exclude_none=True) + if response_dict.get("object_permission"): + # Remove reverse relations from object_permission + for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]: + response_dict["object_permission"].pop(field, None) + + return response_dict except Exception as e: verbose_proxy_logger.exception( @@ -520,6 +544,15 @@ async def update_end_user( ## Update user table, with update params + new budget id (if set) ## verbose_proxy_logger.debug("/customer/update: Received data = %s", data) + + # Ensure object_permission is not in the update data + # It should have been converted to object_permission_id by handle_update_object_permission_common + if "object_permission" in update_end_user_table_data: + verbose_proxy_logger.warning( + f"object_permission still in update_end_user_table_data: {update_end_user_table_data.get('object_permission')}" + ) + update_end_user_table_data.pop("object_permission", None) + if data.user_id is not None and len(data.user_id) > 0: update_end_user_table_data["user_id"] = data.user_id # type: ignore verbose_proxy_logger.debug("In update customer, user_id condition block.") @@ -533,7 +566,15 @@ async def update_end_user( verbose_proxy_logger.debug( f"received response from updating prisma client. response={response}" ) - return response + + # Convert to dict and clean up recursive fields + response_dict = response.model_dump() + if response_dict.get("object_permission"): + # Remove reverse relations from object_permission + for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]: + response_dict["object_permission"].pop(field, None) + + return response_dict else: raise ValueError(f"user_id is required, passed user_id = {data.user_id}") @@ -690,7 +731,12 @@ async def list_end_user( returned_response: List[LiteLLM_EndUserTable] = [] for item in response: - returned_response.append(LiteLLM_EndUserTable(**item.model_dump())) + item_dict = item.model_dump() + # Remove reverse relations from object_permission + if item_dict.get("object_permission"): + for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]: + item_dict["object_permission"].pop(field, None) + returned_response.append(LiteLLM_EndUserTable(**item_dict)) return returned_response except Exception as e: