use helper for _can_user_access_user_info

This commit is contained in:
Ishaan Jaff 2024-10-26 11:35:10 +04:00
parent efbfbbeb30
commit 2f8478769b
2 changed files with 71 additions and 15 deletions

View file

@ -2,6 +2,7 @@ import re
from typing import Optional
from fastapi import HTTPException, Request, status
from pydantic_core.core_schema import url_schema
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
@ -13,14 +14,14 @@ from litellm.proxy._types import (
)
from litellm.proxy.utils import hash_token
from .auth_checks_organization import _user_is_org_admin
from .auth_checks_organization import OrganizationRoleBasedAccessChecks
from .auth_utils import _has_user_setup_sso
class RouteChecks:
@staticmethod
def non_proxy_admin_allowed_routes_check(
async def non_proxy_admin_allowed_routes_check(
user_obj: Optional[LiteLLM_UserTable],
_user_role: Optional[LitellmUserRoles],
route: str,
@ -59,13 +60,11 @@ class RouteChecks:
verbose_proxy_logger.debug(
f"user_id: {user_id} & valid_token.user_id: {valid_token.user_id}"
)
if user_id and user_id != valid_token.user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="key not allowed to access this user's info. user_id={}, key's user_id={}".format(
user_id, valid_token.user_id
),
)
await RouteChecks._can_user_access_user_info(
valid_token=valid_token,
user_id=user_id,
user_obj=user_obj,
)
elif route == "/model/info":
# /model/info just shows models user has access to
pass
@ -111,7 +110,9 @@ class RouteChecks:
):
pass
elif (
_user_is_org_admin(request_data=request_data, user_object=user_obj)
OrganizationRoleBasedAccessChecks._user_is_org_admin_in_requested_org(
request_data=request_data, user_object=user_obj
)
and route in LiteLLMRoutes.org_admin_allowed_routes.value
):
pass
@ -213,3 +214,58 @@ class RouteChecks:
if re.match(pattern, route):
return True
return False
@staticmethod
async def _can_user_access_user_info(
valid_token: UserAPIKeyAuth,
user_id: Optional[str] = None,
user_obj: Optional[LiteLLM_UserTable] = None,
) -> bool:
"""
Returns True if valid_token is allowed to access `user_id` info
User is allowed to access /user/info for:
- their own user_id
- any user_id if they are an admin in that user's organization
Args:
valid_token (UserAPIKeyAuth): The valid token object
user_id (str): The user_id to check if the valid token is allowed to access
user_obj (LiteLLM_UserTable): The user object containing organization memberships
Returns:
bool: True if valid_token is allowed to access user_id info, False otherwise
Raises:
HTTPException: If valid_token is not allowed to access user_id info
"""
from litellm.proxy.proxy_server import prisma_client
if user_id is None:
return True
# user is trying to access their own info
if user_id == valid_token.user_id:
return True
# Check if Org Admin can access info for user_id
if prisma_client is not None:
# Org Admins can access info for users in their org
_queried_user_info = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
if _queried_user_info is not None:
for _membership in _queried_user_info.organization_memberships:
if OrganizationRoleBasedAccessChecks._user_is_admin_in_org(
user_object=user_obj,
organization_id=_membership.organization_id,
):
return True
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="key not allowed to access this user's info. user_id={}, key's user_id={}".format(
user_id, valid_token.user_id
),
)

View file

@ -138,7 +138,7 @@ def _is_ui_route_allowed(
)
def _is_api_route_allowed(
async def _is_api_route_allowed(
route: str,
request: Request,
request_data: dict,
@ -155,7 +155,7 @@ def _is_api_route_allowed(
raise Exception("Invalid proxy server token passed")
if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin
RouteChecks.non_proxy_admin_allowed_routes_check(
await RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=_user_role,
route=route,
@ -167,7 +167,7 @@ def _is_api_route_allowed(
return True
def _is_allowed_route(
async def _is_allowed_route(
route: str,
token_type: Literal["ui", "api"],
request: Request,
@ -182,7 +182,7 @@ def _is_allowed_route(
if token_type == "ui":
return _is_ui_route_allowed(route=route, user_obj=user_obj)
else:
return _is_api_route_allowed(
return await _is_api_route_allowed(
route=route,
request=request,
request_data=request_data,
@ -1137,7 +1137,7 @@ async def user_api_key_auth( # noqa: PLR0915
if token_team is not None and token_team == "litellm-dashboard"
else "api"
)
_is_route_allowed = _is_allowed_route(
_is_route_allowed = await _is_allowed_route(
route=route,
token_type=token_type,
user_obj=user_obj,