diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 54f567b7aa2..3f3b08d3d5e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -571,26 +571,28 @@ def common_key_access_checks( llm_router: Router | None, premium_user: bool, user_id: str | None = None, + enforce_self_key_restriction: bool = True, ) -> Literal[True]: """ Check if user is allowed to make a key request, for this key """ - try: - _is_allowed_to_make_key_request( - user_api_key_dict=user_api_key_dict, - user_id=user_id or data.user_id, - team_id=data.team_id, - ) - except AssertionError as e: - raise HTTPException( - status_code=403, - detail=str(e), - ) - except Exception as e: - raise HTTPException( - status_code=500, - detail=str(e), - ) + if enforce_self_key_restriction: + try: + _is_allowed_to_make_key_request( + user_api_key_dict=user_api_key_dict, + user_id=user_id or data.user_id, + team_id=data.team_id, + ) + except AssertionError as e: + raise HTTPException( + status_code=403, + detail=str(e), + ) + except Exception as e: + raise HTTPException( + status_code=500, + detail=str(e), + ) _check_model_access_group( models=data.models, @@ -2491,6 +2493,29 @@ async def _validate_mcp_servers_for_key_update( return normalized_object_permission +async def _caller_is_team_or_org_admin_for_key( + user_api_key_dict: UserAPIKeyAuth, + existing_key_row: LiteLLM_VerificationToken, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, +) -> bool: + """Team admins and org admins may update keys owned by other members of the key's team.""" + if existing_key_row.team_id is None or prisma_client is None: + return False + try: + team_obj: Final = await get_team_object( + team_id=existing_key_row.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, + ) + except HTTPException: + return False + if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): + return True + return await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj) + + async def _validate_update_key_data( data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken, @@ -2527,12 +2552,20 @@ async def _validate_update_key_data( user_api_key_dict=user_api_key_dict, ) + _caller_is_key_team_admin: Final = await _caller_is_team_or_org_admin_for_key( + user_api_key_dict=user_api_key_dict, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + common_key_access_checks( user_api_key_dict=user_api_key_dict, data=data, user_id=existing_key_row.user_id, llm_router=llm_router, premium_user=premium_user, + enforce_self_key_restriction=not _caller_is_key_team_admin, ) await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 0c615cbaa32..deab2ad1927 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -779,6 +779,110 @@ async def test_update_key_personal_non_admin_denied_access_groups( assert "Access groups" in str(exc.value.detail) +def _team_key_update_fixtures(): + team_obj = LiteLLM_TeamTableCachedObj( + team_id="team-a", + members_with_roles=[ + Member(user_id="user-a", role="user"), + Member(user_id="user-b", role="admin"), + ], + ) + existing_key_row = MagicMock( + token="hashed_user_a_team_key", + user_id="user-a", + team_id="team-a", + created_by="user-a", + max_budget=None, + organization_id=None, + project_id=None, + metadata=None, + object_permission_id=None, + models=[], + ) + return team_obj, existing_key_row + + +@pytest.mark.asyncio +async def test_update_key_team_admin_can_update_member_key(monkeypatch): + """A team admin must be able to update another member's key on their team + (e.g. change its expiration). Regression test for the + 'User can only create keys for themselves' 403 on /key/update.""" + team_obj, existing_key_row = _team_key_update_fixtures() + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock(return_value=team_obj), + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + AsyncMock(return_value=team_obj), + ) + + result = await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-user-a-team-key", duration="30d"), + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user-b", + user_id="user-b", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_update_key_non_admin_member_still_denied(monkeypatch): + """A regular team member (not team admin) must still be blocked from + updating another member's key.""" + team_obj, existing_key_row = _team_key_update_fixtures() + team_obj.team_member_permissions = ["/key/update"] + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock(return_value=team_obj), + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + AsyncMock(return_value=team_obj), + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-user-b-team-key", duration="30d"), + existing_key_row=MagicMock( + token="hashed_user_b_team_key", + user_id="user-b", + team_id="team-a", + created_by="user-b", + max_budget=None, + organization_id=None, + project_id=None, + metadata=None, + object_permission_id=None, + models=[], + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user-a", + user_id="user-a", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "User can only create keys for themselves" in str(exc.value.detail) + + @pytest.mark.asyncio async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch): """Ensure generate_key_helper_fn passes access_group_ids into the key insert payload."""