feat: cleanup for customer_endpoints support of object permission id

This commit is contained in:
Krrish Dholakia 2026-02-14 19:27:51 -08:00
parent 3e06e27293
commit f365402cca
2 changed files with 126 additions and 35 deletions

View file

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

View file

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