fix(key_management_endpoints.py): allow team/org admins to update team member keys

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-08-23 20:32:59 +00:00
parent f005afa146
commit 9d06b8566b
2 changed files with 153 additions and 16 deletions

View file

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

View file

@ -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."""