From fa6c9bf42ef8aba7b7d31d0e2384b0fcf29407a3 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 20 Aug 2024 14:01:12 -0700 Subject: [PATCH 1/5] feat(user_api_key_auth.py): allow team admin to add new members to team --- litellm/proxy/_types.py | 10 ++ litellm/proxy/auth/user_api_key_auth.py | 18 +- .../management_endpoints/team_endpoints.py | 1 + litellm/proxy/utils.py | 29 ++++ litellm/tests/test_proxy_server.py | 163 ++++++++++++++++++ 5 files changed, 220 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 75934ee1f15..fa778093683 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -21,6 +21,13 @@ else: Span = Any +class LiteLLMTeamRoles(enum.Enum): + # team admin + TEAM_ADMIN = "admin" + # team member + TEAM_MEMBER = "user" + + class LitellmUserRoles(str, enum.Enum): """ Admin Roles: @@ -335,6 +342,8 @@ class LiteLLMRoutes(enum.Enum): + sso_only_routes ) + team_admin_routes: List = ["/team/member_add"] + internal_user_routes + # class LiteLLMAllowedRoutes(LiteLLMBase): # """ @@ -1308,6 +1317,7 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): soft_budget: Optional[float] = None team_model_aliases: Optional[Dict] = None team_member_spend: Optional[float] = None + team_member: Optional[Member] = None team_metadata: Optional[Dict] = None # End User Params diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index c980f47b58f..d20ab54bd8f 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -975,7 +975,7 @@ async def user_api_key_auth( if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin if is_llm_api_route(route=route): pass - elif is_llm_api_route(route=request["route"].name): + elif is_llm_api_route(route=route): pass elif ( route in LiteLLMRoutes.info_routes.value @@ -1046,11 +1046,17 @@ async def user_api_key_auth( status_code=status.HTTP_403_FORBIDDEN, detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}", ) + elif ( _user_role == LitellmUserRoles.INTERNAL_USER.value and route in LiteLLMRoutes.internal_user_routes.value ): pass + elif ( + _is_user_team_admin(user_api_key_dict=valid_token) + and route in LiteLLMRoutes.team_admin_routes.value + ): + pass else: user_role = "unknown" user_id = "unknown" @@ -1326,3 +1332,13 @@ def get_api_key_from_custom_header( f"No LiteLLM Virtual Key pass. Please set header={custom_litellm_key_header_name}: Bearer " ) return api_key + + +def _is_user_team_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: + if user_api_key_dict.team_member is None: + return False + + if user_api_key_dict.team_member.role == LiteLLMTeamRoles.TEAM_ADMIN.value: + return True + + return False diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 815ab308c1b..4b7af502d8f 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -417,6 +417,7 @@ async def team_member_add( If user doesn't exist, new user row will also be added to User Table + Only proxy_admin or admin of team, allowed to access this endpoint. ``` curl -X POST 'http://0.0.0.0:4000/team/member_add' \ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a2b09b4e697..a7701771791 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -44,6 +44,7 @@ from litellm.proxy._types import ( DynamoDBArgs, LiteLLM_VerificationTokenView, LitellmUserRoles, + Member, ResetTeamBudgetRequest, SpendLogsMetadata, SpendLogsPayload, @@ -1395,6 +1396,7 @@ class PrismaClient: t.blocked AS team_blocked, t.team_alias AS team_alias, t.metadata AS team_metadata, + t.members_with_roles AS team_members_with_roles, tm.spend AS team_member_spend, m.aliases as team_model_aliases FROM "LiteLLM_VerificationToken" AS v @@ -1412,6 +1414,33 @@ class PrismaClient: response["team_models"] = [] if response["team_blocked"] is None: response["team_blocked"] = False + + team_member: Optional[Member] = None + if ( + response["team_members_with_roles"] is not None + and response["user_id"] is not None + ): + ## find the team member corresponding to user id + """ + [ + { + "role": "admin", + "user_id": "default_user_id", + "user_email": null + }, + { + "role": "user", + "user_id": null, + "user_email": "test@email.com" + } + ] + """ + for tm in response["team_members_with_roles"]: + if tm.get("user_id") is not None and response[ + "user_id" + ] == tm.get("user_id"): + team_member = Member(**tm) + response["team_member"] = team_member response = LiteLLM_VerificationTokenView( **response, last_refreshed_at=time.time() ) diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index 28f3aad6322..7caf1ecbc13 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -930,6 +930,169 @@ async def test_create_team_member_add(prisma_client, new_member_method): ) +@pytest.mark.parametrize("team_member_role", ["admin", "user"]) +@pytest.mark.asyncio +async def test_create_team_member_add_team_admin_user_api_key_auth( + prisma_client, team_member_role +): + import time + + from fastapi import Request + + from litellm.proxy._types import LiteLLM_TeamTableCachedObj, Member + from litellm.proxy.proxy_server import ( + ProxyException, + hash_token, + user_api_key_auth, + user_api_key_cache, + ) + + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm, "max_internal_user_budget", 10) + setattr(litellm, "internal_user_budget_duration", "5m") + await litellm.proxy.proxy_server.prisma_client.connect() + user = f"ishaan {uuid.uuid4().hex}" + _team_id = "litellm-test-client-id-new" + user_key = "sk-12345678" + + valid_token = UserAPIKeyAuth( + team_id=_team_id, + token=hash_token(user_key), + team_member=Member(role=team_member_role, user_id=user), + last_refreshed_at=time.time(), + ) + user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) + + team_obj = LiteLLM_TeamTableCachedObj( + team_id=_team_id, + blocked=False, + last_refreshed_at=time.time(), + metadata={"guardrails": {"modify_guardrails": False}}, + ) + + user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) + + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + + ## TEST IF TEAM ADMIN ALLOWED TO CALL /MEMBER_ADD ENDPOINT + import json + + from starlette.datastructures import URL + + request = Request(scope={"type": "http"}) + request._url = URL(url="/team/member_add") + + body = {} + json_bytes = json.dumps(body).encode("utf-8") + + request._body = json_bytes + + try: + await user_api_key_auth(request=request, api_key="Bearer " + user_key) + if team_member_role == "user": + pytest.fail( + "Expected this call to fail. User not allowed to access this route." + ) + except ProxyException: + if team_member_role == "admin": + pytest.fail( + "Expected this call to succeed. Team admin allowed to access /team/member_add" + ) + + +@pytest.mark.parametrize("new_member_method", ["user_id", "user_email"]) +@pytest.mark.asyncio +async def test_create_team_member_add_team_admin(prisma_client, new_member_method): + """ + Relevant issue - https://github.com/BerriAI/litellm/issues/5300 + + Allow team admins to: + - Add and remove team members + - raise error if team member not an existing 'internal_user' + """ + import time + + from fastapi import Request + + from litellm.proxy._types import LiteLLM_TeamTableCachedObj, Member + from litellm.proxy.proxy_server import ( + ProxyException, + hash_token, + user_api_key_auth, + user_api_key_cache, + ) + + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm, "max_internal_user_budget", 10) + setattr(litellm, "internal_user_budget_duration", "5m") + await litellm.proxy.proxy_server.prisma_client.connect() + user = f"ishaan {uuid.uuid4().hex}" + _team_id = "litellm-test-client-id-new" + user_key = "sk-12345678" + + valid_token = UserAPIKeyAuth( + team_id=_team_id, + token=hash_token(user_key), + team_member=Member(role="admin", user_id=user), + last_refreshed_at=time.time(), + ) + user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) + + team_obj = LiteLLM_TeamTableCachedObj( + team_id=_team_id, + blocked=False, + last_refreshed_at=time.time(), + metadata={"guardrails": {"modify_guardrails": False}}, + ) + + user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) + + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + if new_member_method == "user_id": + data = { + "team_id": _team_id, + "member": [{"role": "user", "user_id": user}], + } + elif new_member_method == "user_email": + data = { + "team_id": _team_id, + "member": [{"role": "user", "user_email": user}], + } + team_member_add_request = TeamMemberAddRequest(**data) + + with patch( + "litellm.proxy.proxy_server.prisma_client.db.litellm_usertable", + new_callable=AsyncMock, + ) as mock_litellm_usertable: + mock_client = AsyncMock() + mock_litellm_usertable.upsert = mock_client + mock_litellm_usertable.find_many = AsyncMock(return_value=None) + + await team_member_add( + data=team_member_add_request, + user_api_key_dict=valid_token, + http_request=Request( + scope={"type": "http", "path": "/user/new"}, + ), + ) + + mock_client.assert_called() + + print(f"mock_client.call_args: {mock_client.call_args}") + print("mock_client.call_args.kwargs: {}".format(mock_client.call_args.kwargs)) + + assert ( + mock_client.call_args.kwargs["data"]["create"]["max_budget"] + == litellm.max_internal_user_budget + ) + assert ( + mock_client.call_args.kwargs["data"]["create"]["budget_duration"] + == litellm.internal_user_budget_duration + ) + + @pytest.mark.asyncio async def test_user_info_team_list(prisma_client): """Assert user_info for admin calls team_list function""" From 19083a4d3193cd5f5634ff9a7b4005971c487fc1 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 20 Aug 2024 16:25:13 -0700 Subject: [PATCH 2/5] feat(_types.py): allow team admin to delete member from team --- litellm/proxy/_types.py | 5 ++++- litellm/tests/test_proxy_server.py | 5 +++-- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index fa778093683..25cf2f56d3f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -342,7 +342,10 @@ class LiteLLMRoutes(enum.Enum): + sso_only_routes ) - team_admin_routes: List = ["/team/member_add"] + internal_user_routes + team_admin_routes: List = [ + "/team/member_add", + "/team/member_delete", + ] + internal_user_routes # class LiteLLMAllowedRoutes(LiteLLMBase): diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index 7caf1ecbc13..a1d6d9dee1f 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -931,9 +931,10 @@ async def test_create_team_member_add(prisma_client, new_member_method): @pytest.mark.parametrize("team_member_role", ["admin", "user"]) +@pytest.mark.parametrize("team_route", ["/team/member_add", "/team/member_delete"]) @pytest.mark.asyncio async def test_create_team_member_add_team_admin_user_api_key_auth( - prisma_client, team_member_role + prisma_client, team_member_role, team_route ): import time @@ -981,7 +982,7 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( from starlette.datastructures import URL request = Request(scope={"type": "http"}) - request._url = URL(url="/team/member_add") + request._url = URL(url=team_route) body = {} json_bytes = json.dumps(body).encode("utf-8") From a61f3e7656bee7d70948860edf2bcce0bc45140f Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 20 Aug 2024 16:57:18 -0700 Subject: [PATCH 3/5] refactor(team_endpoints.py): refactor auth checks for team member endpoints to ui team admin to manage it --- litellm/proxy/_types.py | 4 +- litellm/proxy/auth/user_api_key_auth.py | 17 +------ .../key_management_endpoints.py | 2 +- .../management_endpoints/team_endpoints.py | 46 ++++++++++++++++++- litellm/tests/test_proxy_server.py | 42 +++++++++-------- 5 files changed, 72 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 25cf2f56d3f..0177c219074 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -342,10 +342,10 @@ class LiteLLMRoutes(enum.Enum): + sso_only_routes ) - team_admin_routes: List = [ + self_managed_routes: List = [ "/team/member_add", "/team/member_delete", - ] + internal_user_routes + ] # routes that manage their own allowed/disallowed logic # class LiteLLMAllowedRoutes(LiteLLMBase): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index d20ab54bd8f..378dd845255 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -975,8 +975,6 @@ async def user_api_key_auth( if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin if is_llm_api_route(route=route): pass - elif is_llm_api_route(route=route): - pass elif ( route in LiteLLMRoutes.info_routes.value ): # check if user allowed to call an info route @@ -1053,9 +1051,8 @@ async def user_api_key_auth( ): pass elif ( - _is_user_team_admin(user_api_key_dict=valid_token) - and route in LiteLLMRoutes.team_admin_routes.value - ): + route in LiteLLMRoutes.self_managed_routes.value + ): # routes that manage their own allowed/disallowed logic pass else: user_role = "unknown" @@ -1332,13 +1329,3 @@ def get_api_key_from_custom_header( f"No LiteLLM Virtual Key pass. Please set header={custom_litellm_key_header_name}: Bearer " ) return api_key - - -def _is_user_team_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: - if user_api_key_dict.team_member is None: - return False - - if user_api_key_dict.team_member.role == LiteLLMTeamRoles.TEAM_ADMIN.value: - return True - - return False diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1758b416dda..2e16b533c87 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -849,7 +849,7 @@ async def generate_key_helper_fn( } if ( - litellm.get_secret("DISABLE_KEY_NAME", False) == True + litellm.get_secret("DISABLE_KEY_NAME", False) is True ): # allow user to disable storing abbreviated key name (shown in UI, to help figure out which key spent how much) pass else: diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4b7af502d8f..8f8ad610a80 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -30,7 +30,7 @@ from litellm.proxy._types import ( UpdateTeamRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.auth.user_api_key_auth import _is_user_proxy_admin, user_api_key_auth from litellm.proxy.management_helpers.utils import ( add_new_member, management_endpoint_wrapper, @@ -39,6 +39,16 @@ from litellm.proxy.management_helpers.utils import ( router = APIRouter() +def _is_user_team_admin( + user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable +) -> bool: + for member in team_obj.members_with_roles: + if member.user_id is not None and member.user_id == user_api_key_dict.user_id: + return True + + return False + + #### TEAM MANAGEMENT #### @router.post( "/team/new", @@ -466,6 +476,23 @@ async def team_member_add( complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN + + if ( + user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + and not _is_user_team_admin( + user_api_key_dict=user_api_key_dict, team_obj=complete_team_data + ) + ): + raise HTTPException( + status_code=403, + detail={ + "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( + "/team/member_add", complete_team_data.team_id + ) + }, + ) + if isinstance(data.member, Member): # add to team db new_member = data.member @@ -570,6 +597,23 @@ async def team_member_delete( ) existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN + + if ( + user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + and not _is_user_team_admin( + user_api_key_dict=user_api_key_dict, team_obj=existing_team_row + ) + ): + raise HTTPException( + status_code=403, + detail={ + "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( + "/team/member_delete", existing_team_row.team_id + ) + }, + ) + ## DELETE MEMBER FROM TEAM new_team_members: List[Member] = [] for m in existing_team_row.members_with_roles: diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index a1d6d9dee1f..79da5bf7c55 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -989,22 +989,16 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( request._body = json_bytes - try: - await user_api_key_auth(request=request, api_key="Bearer " + user_key) - if team_member_role == "user": - pytest.fail( - "Expected this call to fail. User not allowed to access this route." - ) - except ProxyException: - if team_member_role == "admin": - pytest.fail( - "Expected this call to succeed. Team admin allowed to access /team/member_add" - ) + ## ALLOWED BY USER_API_KEY_AUTH + await user_api_key_auth(request=request, api_key="Bearer " + user_key) @pytest.mark.parametrize("new_member_method", ["user_id", "user_email"]) +@pytest.mark.parametrize("user_role", ["admin", "user"]) @pytest.mark.asyncio -async def test_create_team_member_add_team_admin(prisma_client, new_member_method): +async def test_create_team_member_add_team_admin( + prisma_client, new_member_method, user_role +): """ Relevant issue - https://github.com/BerriAI/litellm/issues/5300 @@ -1018,6 +1012,7 @@ async def test_create_team_member_add_team_admin(prisma_client, new_member_metho from litellm.proxy._types import LiteLLM_TeamTableCachedObj, Member from litellm.proxy.proxy_server import ( + HTTPException, ProxyException, hash_token, user_api_key_auth, @@ -1035,8 +1030,8 @@ async def test_create_team_member_add_team_admin(prisma_client, new_member_metho valid_token = UserAPIKeyAuth( team_id=_team_id, + user_id=user, token=hash_token(user_key), - team_member=Member(role="admin", user_id=user), last_refreshed_at=time.time(), ) user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) @@ -1045,6 +1040,7 @@ async def test_create_team_member_add_team_admin(prisma_client, new_member_metho team_id=_team_id, blocked=False, last_refreshed_at=time.time(), + members_with_roles=[Member(role=user_role, user_id=user)], metadata={"guardrails": {"modify_guardrails": False}}, ) @@ -1071,13 +1067,19 @@ async def test_create_team_member_add_team_admin(prisma_client, new_member_metho mock_litellm_usertable.upsert = mock_client mock_litellm_usertable.find_many = AsyncMock(return_value=None) - await team_member_add( - data=team_member_add_request, - user_api_key_dict=valid_token, - http_request=Request( - scope={"type": "http", "path": "/user/new"}, - ), - ) + try: + await team_member_add( + data=team_member_add_request, + user_api_key_dict=valid_token, + http_request=Request( + scope={"type": "http", "path": "/user/new"}, + ), + ) + except HTTPException as e: + if user_role == "user": + assert e.status_code == 403 + else: + raise e mock_client.assert_called() From 5ba517819cef4bf315117fd1e9d3ba115bfefc15 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 21 Aug 2024 08:37:04 -0700 Subject: [PATCH 4/5] test(test_proxy_server.py): fix test to specify user role --- litellm/proxy/management_endpoints/team_endpoints.py | 3 ++- litellm/tests/test_proxy_server.py | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8f8ad610a80..d3c2e3e839b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -488,7 +488,8 @@ async def team_member_add( status_code=403, detail={ "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( - "/team/member_add", complete_team_data.team_id + "/team/member_add", + complete_team_data.team_id, ) }, ) diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index 79da5bf7c55..c4d9afbbdb3 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -909,7 +909,7 @@ async def test_create_team_member_add(prisma_client, new_member_method): await team_member_add( data=team_member_add_request, - user_api_key_dict=UserAPIKeyAuth(), + user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"), http_request=Request( scope={"type": "http", "path": "/user/new"}, ), From 83bed56b6620b075c73e944a4d1bc013b61ebbf5 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 21 Aug 2024 12:46:43 -0700 Subject: [PATCH 5/5] fix(internal_user_endpoints.py): pass in user api key dict value --- litellm/proxy/management_endpoints/internal_user_endpoints.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index a0e020b11fc..b5701711829 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -119,6 +119,7 @@ async def new_user( http_request=Request( scope={"type": "http", "path": "/user/new"}, ), + user_api_key_dict=user_api_key_dict, ) if data.send_invite_email is True: