mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
feat: cleanup for customer_endpoints support of object permission id
This commit is contained in:
parent
3e06e27293
commit
f365402cca
2 changed files with 126 additions and 35 deletions
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue