diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index db4467dda5b..4439a1f1267 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2680,6 +2680,52 @@ def _require_prisma_client(prisma_client: PrismaClient | None) -> PrismaClient: return prisma_client +async def _validate_update_key_caller_access( + data: UpdateKeyRequest, + existing_key_row: LiteLLM_VerificationToken, + user_api_key_dict: UserAPIKeyAuth, + llm_router: Router | None, + premium_user: bool, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + is_proxy_admin: bool, +) -> None: + _caller_is_key_owner: Final = ( + existing_key_row.user_id is None or existing_key_row.user_id == user_api_key_dict.user_id + ) + if is_proxy_admin or _caller_is_key_owner or existing_key_row.team_id is None: + 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, + ) + return + 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, + ) + if team_obj is None or not ( # pyright: ignore[reportUnnecessaryComparison] # get_team_object returns None when the team row is missing despite the non-Optional annotation + _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) + or await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj) + ): + raise HTTPException( + status_code=403, + detail={ # mutable-ok: FastAPI detail contract + "error": f"Only proxy admins, team admins, or org admins can call /key/update. " + f"user_role={user_api_key_dict.user_role}, user_id={user_api_key_dict.user_id}" + }, + ) + _check_model_access_group( + models=data.models, + llm_router=llm_router, + premium_user=premium_user, + ) + + async def _validate_update_key_data( data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken, @@ -2718,12 +2764,15 @@ async def _validate_update_key_data( user_api_key_dict=user_api_key_dict, ) - common_key_access_checks( - user_api_key_dict=user_api_key_dict, + await _validate_update_key_caller_access( data=data, - user_id=existing_key_row.user_id, + existing_key_row=existing_key_row, + user_api_key_dict=user_api_key_dict, llm_router=llm_router, premium_user=premium_user, + prisma_client=checked_prisma_client, + user_api_key_cache=user_api_key_cache, + is_proxy_admin=_is_proxy_admin, ) await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( diff --git a/tests/proxy_behavior/management/test_key_update.py b/tests/proxy_behavior/management/test_key_update.py index 7b7f6f5558b..2d8fe6b315f 100644 --- a/tests/proxy_behavior/management/test_key_update.py +++ b/tests/proxy_behavior/management/test_key_update.py @@ -12,8 +12,10 @@ pytestmark = pytest.mark.asyncio(loop_scope="session") # (id, actor, target_shape, expected_status). Pinned against current gating: # proxy_admin bypasses; org_admin is blocked by an early role gate (401); -# every other (INTERNAL_USER-roled) actor hits user_id-mismatch 403, no-team- -# admin 403, or team_member_permission 401 depending on target / membership. +# a team admin may update a key in their own team even when another member +# owns it (200); every other (INTERNAL_USER-roled) actor hits user_id-mismatch +# 403, no-team-admin 403, or team_member_permission 401 depending on target / +# membership. _SCENARIOS = [ ("self/proxy_admin", Actor.PROXY_ADMIN, "self", 200), ("self/org_admin", Actor.ORG_ADMIN, "self", 401), @@ -25,7 +27,7 @@ _SCENARIOS = [ ("self/service_account", Actor.SERVICE_ACCOUNT, "self", 403), ("owner_target/proxy_admin", Actor.PROXY_ADMIN, "owner", 200), ("owner_target/org_admin", Actor.ORG_ADMIN, "owner", 401), - ("owner_target/team_admin", Actor.TEAM_ADMIN, "owner", 403), + ("owner_target/team_admin", Actor.TEAM_ADMIN, "owner", 200), ("owner_target/internal_user", Actor.INTERNAL_USER, "owner", 403), ("owner_target/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "owner", 403), ("owner_target/cross_org_user", Actor.CROSS_ORG_USER, "owner", 403), 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 65cc23ea67f..f8e3955c3dd 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 @@ -10914,6 +10914,200 @@ async def test_update_key_team_member_cannot_change_budget(monkeypatch): assert str(exc.value.code) == "403" +@pytest.mark.asyncio +async def test_update_key_team_admin_can_change_budget_of_member_key(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "feedface" * 8 + team_id = "team-admin-budget" + admin_user_id = "team-admin" + member_user_id = "member-1" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = member_user_id + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "member-key" + mock_existing_key.models = [] + mock_existing_key.metadata = {} + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": member_user_id, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id=admin_user_id, role="admin"), + Member(user_id=member_user_id, role="user"), + ], + ) + + mock_updated_key = MagicMock() + mock_updated_key.token = test_hashed_token + mock_updated_key.max_budget = 500.0 + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=mock_existing_key) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token) + + async def mock_delete_cache_key_object(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + mock_delete_cache_key_object, + ) + + async def mock_enforce_unique_key_alias(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", + mock_enforce_unique_key_alias, + ) + + mock_request = MagicMock() + mock_request.query_params = {} + team_admin = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-admin", + user_id=admin_user_id, + team_id=team_id, + ) + + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, max_budget=500.0), + user_api_key_dict=team_admin, + litellm_changed_by=None, + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_update_key_non_admin_team_member_cannot_update_other_members_key( + monkeypatch, +): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "feedface" * 8 + team_id = "team-admin-budget" + admin_user_id = "team-admin" + member_user_id = "member-1" + caller_user_id = "member-2" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = member_user_id + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "member-key" + mock_existing_key.models = [] + mock_existing_key.metadata = {} + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": member_user_id, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id=admin_user_id, role="admin"), + Member(user_id=member_user_id, role="user"), + Member(user_id=caller_user_id, role="user"), + ], + team_member_permissions=["/key/update"], + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=mock_existing_key) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token) + + mock_request = MagicMock() + mock_request.query_params = {} + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-member-2", + user_id=caller_user_id, + team_id=team_id, + ) + + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, max_budget=500.0), + user_api_key_dict=team_member, + litellm_changed_by=None, + ) + assert str(exc.value.code) == "403" + + # ============================================================================ # LIT-1884: Internal users cannot create invalid keys # ============================================================================