mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
add endpoint for bulk key updates for team
This commit is contained in:
parent
b318231fe9
commit
ed06f3bcc0
4 changed files with 752 additions and 27 deletions
|
|
@ -239,6 +239,7 @@ class KeyManagementRoutes(str, enum.Enum):
|
|||
KEY_BLOCK = "/key/block"
|
||||
KEY_UNBLOCK = "/key/unblock"
|
||||
KEY_BULK_UPDATE = "/key/bulk_update"
|
||||
TEAM_KEY_BULK_UPDATE = "/team/key/bulk_update"
|
||||
KEY_RESET_SPEND = "/key/{key_id}/reset_spend"
|
||||
|
||||
# info and health routes
|
||||
|
|
@ -538,6 +539,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
KeyManagementRoutes.KEY_BLOCK.value,
|
||||
KeyManagementRoutes.KEY_UNBLOCK.value,
|
||||
KeyManagementRoutes.KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
|
||||
KeyManagementRoutes.SPEND_LOGS.value,
|
||||
KeyManagementRoutes.KEY_RESET_SPEND.value,
|
||||
|
|
|
|||
|
|
@ -88,8 +88,8 @@ from litellm.router import Router
|
|||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateKeyRequest,
|
||||
BulkUpdateKeyRequestItem,
|
||||
BulkUpdateKeyResponse,
|
||||
BulkUpdateTeamKeysRequest,
|
||||
FailedKeyUpdate,
|
||||
SuccessfulKeyUpdate,
|
||||
)
|
||||
|
|
@ -1881,7 +1881,7 @@ async def _get_and_validate_existing_key(
|
|||
|
||||
|
||||
async def _process_single_key_update(
|
||||
key_update_item: BulkUpdateKeyRequestItem,
|
||||
update_key_request: UpdateKeyRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_changed_by: Optional[str],
|
||||
prisma_client: Optional[PrismaClient],
|
||||
|
|
@ -1897,7 +1897,7 @@ async def _process_single_key_update(
|
|||
including validation, permission checks, team checks, and database updates.
|
||||
|
||||
Args:
|
||||
key_update_item: The key update request item
|
||||
update_key_request: Fully-constructed UpdateKeyRequest for the target key
|
||||
user_api_key_dict: The authenticated user's API key info
|
||||
litellm_changed_by: Optional header for tracking who made the change
|
||||
prisma_client: Prisma client instance
|
||||
|
|
@ -1912,11 +1912,11 @@ async def _process_single_key_update(
|
|||
HTTPException: For various validation and permission errors
|
||||
"""
|
||||
# Validate max_budget
|
||||
_validate_max_budget(key_update_item.max_budget)
|
||||
_validate_max_budget(update_key_request.max_budget)
|
||||
|
||||
# Get and validate existing key
|
||||
existing_key_row = await _get_and_validate_existing_key(
|
||||
token=key_update_item.key,
|
||||
token=update_key_request.key,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
|
|
@ -1930,15 +1930,6 @@ async def _process_single_key_update(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
# Create UpdateKeyRequest from BulkUpdateKeyRequestItem
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key=key_update_item.key,
|
||||
budget_id=key_update_item.budget_id,
|
||||
max_budget=key_update_item.max_budget,
|
||||
team_id=key_update_item.team_id,
|
||||
tags=key_update_item.tags,
|
||||
)
|
||||
|
||||
# Custom key update hook
|
||||
if user_custom_key_update is not None:
|
||||
if inspect.iscoroutinefunction(user_custom_key_update):
|
||||
|
|
@ -2003,12 +1994,12 @@ async def _process_single_key_update(
|
|||
detail={"error": "Database not connected"},
|
||||
)
|
||||
|
||||
_data = {**non_default_values, "token": key_update_item.key}
|
||||
response = await prisma_client.update_data(token=key_update_item.key, data=_data)
|
||||
_data = {**non_default_values, "token": update_key_request.key}
|
||||
response = await prisma_client.update_data(token=update_key_request.key, data=_data)
|
||||
|
||||
# Delete cache
|
||||
await _delete_cache_key_object(
|
||||
hashed_token=_hash_token_if_needed(key_update_item.key),
|
||||
hashed_token=_hash_token_if_needed(update_key_request.key),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -2598,9 +2589,15 @@ async def bulk_update_keys(
|
|||
|
||||
for key_update_item in data.keys:
|
||||
try:
|
||||
# Process single key update using reusable function
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key=key_update_item.key,
|
||||
budget_id=key_update_item.budget_id,
|
||||
max_budget=key_update_item.max_budget,
|
||||
team_id=key_update_item.team_id,
|
||||
tags=key_update_item.tags,
|
||||
)
|
||||
updated_key_info = await _process_single_key_update(
|
||||
key_update_item=key_update_item,
|
||||
update_key_request=update_key_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -2665,6 +2662,201 @@ async def bulk_update_keys(
|
|||
)
|
||||
|
||||
|
||||
def _build_failed_team_key_update(
|
||||
token: str,
|
||||
exception: Exception,
|
||||
existing_key_row: Optional[LiteLLM_VerificationToken],
|
||||
) -> FailedKeyUpdate:
|
||||
"""Normalize an exception from the per-key update loop into a FailedKeyUpdate."""
|
||||
if isinstance(exception, HTTPException):
|
||||
detail = exception.detail
|
||||
if isinstance(detail, dict):
|
||||
error_message = detail.get("error", str(exception))
|
||||
else:
|
||||
error_message = str(detail)
|
||||
elif isinstance(exception, ProxyException):
|
||||
error_message = exception.message
|
||||
else:
|
||||
error_message = str(exception)
|
||||
|
||||
key_info: Optional[Dict[str, Any]] = None
|
||||
if existing_key_row is not None:
|
||||
if hasattr(existing_key_row, "model_dump"):
|
||||
key_info = existing_key_row.model_dump()
|
||||
elif hasattr(existing_key_row, "dict"):
|
||||
key_info = existing_key_row.dict()
|
||||
if key_info:
|
||||
key_info.pop("token", None)
|
||||
|
||||
return FailedKeyUpdate(key=token, key_info=key_info, failed_reason=error_message)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/team/key/bulk_update",
|
||||
tags=["key management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=BulkUpdateKeyResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def bulk_update_team_keys(
|
||||
data: BulkUpdateTeamKeysRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
):
|
||||
"""
|
||||
Apply one update payload to many keys inside a single team.
|
||||
|
||||
Scoped to a single team — pass `team_id` plus either `key_ids` (specific
|
||||
keys in that team) or `all_keys_in_team=True`. The same `update_fields`
|
||||
payload is applied to every selected key.
|
||||
|
||||
Callable by proxy admins, or by team admins with `KEY_UPDATE` permission
|
||||
for the target team.
|
||||
|
||||
Each key update is processed independently — partial failures are returned
|
||||
in `failed_updates` rather than aborting the whole batch.
|
||||
|
||||
Example request:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/team/key/bulk_update' \\
|
||||
--header 'Authorization: Bearer sk-1234' \\
|
||||
--header 'Content-Type: application/json' \\
|
||||
--data '{
|
||||
"team_id": "team-123",
|
||||
"all_keys_in_team": true,
|
||||
"update_fields": {"max_budget": 50.0, "budget_duration": "30d"}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
llm_router,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
user_custom_key_update,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "Database not connected"},
|
||||
)
|
||||
|
||||
if not data.team_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "team_id is required"},
|
||||
)
|
||||
|
||||
MAX_BATCH_SIZE = 500
|
||||
if data.key_ids is not None and len(data.key_ids) > MAX_BATCH_SIZE:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Maximum {MAX_BATCH_SIZE} keys can be updated at once. Found {len(data.key_ids)} key_ids."
|
||||
},
|
||||
)
|
||||
|
||||
# Resolve which keys to update (single batched query, scoped to team).
|
||||
if data.all_keys_in_team:
|
||||
existing_keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": data.team_id},
|
||||
order={"token": "asc"},
|
||||
take=MAX_BATCH_SIZE + 1,
|
||||
)
|
||||
if len(existing_keys) > MAX_BATCH_SIZE:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}."
|
||||
},
|
||||
)
|
||||
requested_tokens = [row.token for row in existing_keys]
|
||||
else:
|
||||
# validator guarantees key_ids is set and non-empty when all_keys_in_team=False
|
||||
assert data.key_ids is not None
|
||||
existing_keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": data.team_id, "token": {"in": data.key_ids}}
|
||||
)
|
||||
requested_tokens = list(data.key_ids)
|
||||
|
||||
if not requested_tokens:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"No keys found for team {data.team_id}"},
|
||||
)
|
||||
|
||||
# Fail-fast auth: if caller is not proxy admin, verify they have KEY_UPDATE
|
||||
# permission for this team once (using any in-team key as the auth subject).
|
||||
# Per-key checks inside _process_single_key_update will short-circuit on the
|
||||
# cached team object after this.
|
||||
if (
|
||||
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
and existing_keys
|
||||
):
|
||||
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route=KeyManagementRoutes.KEY_UPDATE,
|
||||
prisma_client=prisma_client,
|
||||
existing_key_row=existing_keys[0],
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
existing_by_token = {row.token: row for row in existing_keys}
|
||||
update_field_dict = data.update_fields.model_dump(exclude_unset=True)
|
||||
|
||||
successful_updates: List[SuccessfulKeyUpdate] = []
|
||||
failed_updates: List[FailedKeyUpdate] = []
|
||||
|
||||
for token in requested_tokens:
|
||||
try:
|
||||
if token not in existing_by_token:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Key not found in team {data.team_id}"},
|
||||
)
|
||||
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key=token,
|
||||
**update_field_dict,
|
||||
)
|
||||
updated_key_info = await _process_single_key_update(
|
||||
update_key_request=update_key_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
user_custom_key_update=user_custom_key_update,
|
||||
)
|
||||
|
||||
successful_updates.append(
|
||||
SuccessfulKeyUpdate(key=token, key_info=updated_key_info)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"Failed to update key {token} in team {data.team_id}: {e}"
|
||||
)
|
||||
failed_updates.append(
|
||||
_build_failed_team_key_update(
|
||||
token=token,
|
||||
exception=e,
|
||||
existing_key_row=existing_by_token.get(token),
|
||||
)
|
||||
)
|
||||
|
||||
return BulkUpdateKeyResponse(
|
||||
total_requested=len(requested_tokens),
|
||||
successful_updates=successful_updates,
|
||||
failed_updates=failed_updates,
|
||||
)
|
||||
|
||||
|
||||
async def validate_key_team_change(
|
||||
key: LiteLLM_VerificationToken,
|
||||
team: LiteLLM_TeamTable,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, model_validator
|
||||
|
||||
from litellm.proxy._types import KeyRequestBase
|
||||
|
||||
|
||||
class BulkUpdateKeyRequestItem(BaseModel):
|
||||
|
|
@ -40,3 +43,71 @@ class BulkUpdateKeyResponse(BaseModel):
|
|||
total_requested: int
|
||||
successful_updates: List[SuccessfulKeyUpdate]
|
||||
failed_updates: List[FailedKeyUpdate]
|
||||
|
||||
|
||||
class KeyUpdateFields(KeyRequestBase):
|
||||
"""
|
||||
Mirror of UpdateKeyRequest minus per-key identifiers (`key`, `key_alias`)
|
||||
and the scope guard (`team_id`). Used as the broadcast payload in
|
||||
BulkUpdateTeamKeysRequest — one set of fields applied to many keys.
|
||||
"""
|
||||
|
||||
duration: Optional[str] = None
|
||||
spend: Optional[float] = None
|
||||
metadata: Optional[dict] = None
|
||||
temp_budget_increase: Optional[float] = None
|
||||
temp_budget_expiry: Optional[datetime] = None
|
||||
auto_rotate: Optional[bool] = None
|
||||
rotation_interval: Optional[str] = None
|
||||
organization_id: Optional[str] = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_temp_budget(self) -> "KeyUpdateFields":
|
||||
if self.temp_budget_increase is not None or self.temp_budget_expiry is not None:
|
||||
if self.temp_budget_increase is None or self.temp_budget_expiry is None:
|
||||
raise ValueError(
|
||||
"temp_budget_increase and temp_budget_expiry must be set together"
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def reject_per_key_or_scope_fields(self) -> "KeyUpdateFields":
|
||||
forbidden = [
|
||||
name
|
||||
for name in ("key", "key_alias", "team_id")
|
||||
if getattr(self, name, None) is not None
|
||||
]
|
||||
if forbidden:
|
||||
raise ValueError(
|
||||
f"Fields not allowed in update_fields for bulk team key updates: "
|
||||
f"{forbidden}. `key`/`key_alias` are per-key identifiers; `team_id` "
|
||||
f"is the scope guard set at the request top level."
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
class BulkUpdateTeamKeysRequest(BaseModel):
|
||||
"""
|
||||
Request for applying one update payload to many keys inside a single team.
|
||||
|
||||
Exactly one of `key_ids` (specific keys in the team) or `all_keys_in_team`
|
||||
(every key in the team) must be provided.
|
||||
"""
|
||||
|
||||
team_id: str
|
||||
key_ids: Optional[List[str]] = None
|
||||
all_keys_in_team: bool = False
|
||||
update_fields: KeyUpdateFields
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_selection(self) -> "BulkUpdateTeamKeysRequest":
|
||||
has_key_ids = self.key_ids is not None and len(self.key_ids) > 0
|
||||
if has_key_ids and self.all_keys_in_team:
|
||||
raise ValueError(
|
||||
"Provide either `key_ids` or `all_keys_in_team=True`, not both."
|
||||
)
|
||||
if not has_key_ids and not self.all_keys_in_team:
|
||||
raise ValueError(
|
||||
"Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`."
|
||||
)
|
||||
return self
|
||||
|
|
|
|||
|
|
@ -5689,7 +5689,7 @@ async def test_process_single_key_update():
|
|||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
):
|
||||
# Create update request
|
||||
key_update_item = BulkUpdateKeyRequestItem(
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key="test-key-123",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
|
|
@ -5703,7 +5703,7 @@ async def test_process_single_key_update():
|
|||
|
||||
# Call the function
|
||||
result = await _process_single_key_update(
|
||||
key_update_item=key_update_item,
|
||||
update_key_request=update_key_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
|
|
@ -9855,9 +9855,6 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
|
|||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_process_single_key_update,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateKeyRequestItem,
|
||||
)
|
||||
|
||||
token_hash = "abc123def456"
|
||||
|
||||
|
|
@ -9900,7 +9897,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
|
|||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
key_update_item = BulkUpdateKeyRequestItem(
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key=token_hash,
|
||||
max_budget=100.0,
|
||||
)
|
||||
|
|
@ -9912,7 +9909,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
|
|||
)
|
||||
|
||||
await _process_single_key_update(
|
||||
key_update_item=key_update_item,
|
||||
update_key_request=update_key_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
|
|
@ -10019,3 +10016,466 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha
|
|||
call_kwargs = mock_delete_cache.call_args.kwargs
|
||||
# The token hash should be passed as-is, NOT double-hashed
|
||||
assert call_kwargs["hashed_token"] == token_hash
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /team/key/bulk_update tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_team_key(token: str, team_id: str = "team-abc") -> LiteLLM_VerificationToken:
|
||||
return LiteLLM_VerificationToken(
|
||||
token=token,
|
||||
user_id="user-123",
|
||||
models=[],
|
||||
team_id=team_id,
|
||||
max_budget=None,
|
||||
)
|
||||
|
||||
|
||||
def _patch_team_keys_helpers(monkeypatch, *, hash_lookup):
|
||||
"""Common monkeypatching for the bulk_update_team_keys path."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data",
|
||||
AsyncMock(return_value={"max_budget": 50.0}),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
lambda token: hash_lookup[token],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
AsyncMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook",
|
||||
AsyncMock(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_success_with_key_ids(monkeypatch):
|
||||
"""Proxy admin updates two explicitly listed keys; both succeed."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
key_a = _make_team_key("tok-a")
|
||||
key_b = _make_team_key("tok-b")
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[key_a, key_b]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
side_effect=[key_a, key_b]
|
||||
)
|
||||
updated_a = MagicMock()
|
||||
updated_a.model_dump.return_value = {"max_budget": 50.0, "user_id": "user-123"}
|
||||
updated_b = MagicMock()
|
||||
updated_b.model_dump.return_value = {"max_budget": 50.0, "user_id": "user-123"}
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
side_effect=[{"data": updated_a}, {"data": updated_b}]
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
_patch_team_keys_helpers(
|
||||
monkeypatch, hash_lookup={"tok-a": "hashed-a", "tok-b": "hashed-b"}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
AsyncMock(),
|
||||
)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
key_ids=["tok-a", "tok-b"],
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
response = await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert response.total_requested == 2
|
||||
assert len(response.successful_updates) == 2
|
||||
assert len(response.failed_updates) == 0
|
||||
assert {u.key for u in response.successful_updates} == {"tok-a", "tok-b"}
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_awaited_once()
|
||||
where_arg = (
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.await_args.kwargs[
|
||||
"where"
|
||||
]
|
||||
)
|
||||
assert where_arg["team_id"] == "team-abc"
|
||||
assert where_arg["token"] == {"in": ["tok-a", "tok-b"]}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_success_all_keys_in_team(monkeypatch):
|
||||
"""`all_keys_in_team=True` resolves via find_many and updates every key."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
keys = [_make_team_key(f"tok-{i}") for i in range(3)]
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=keys
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
side_effect=keys
|
||||
)
|
||||
updated_objs = [MagicMock() for _ in keys]
|
||||
for obj in updated_objs:
|
||||
obj.model_dump.return_value = {"max_budget": 50.0}
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
side_effect=[{"data": obj} for obj in updated_objs]
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
_patch_team_keys_helpers(
|
||||
monkeypatch, hash_lookup={f"tok-{i}": f"hashed-{i}" for i in range(3)}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
AsyncMock(),
|
||||
)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
response = await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert response.total_requested == 3
|
||||
assert len(response.successful_updates) == 3
|
||||
assert len(response.failed_updates) == 0
|
||||
where_kwargs = (
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.await_args.kwargs
|
||||
)
|
||||
assert where_kwargs["take"] == 501
|
||||
assert where_kwargs["where"] == {"team_id": "team-abc"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_key_not_in_team(monkeypatch):
|
||||
"""key_ids includes a token not in the team — that token fails, others succeed."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
in_team = _make_team_key("tok-a")
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[in_team]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=in_team
|
||||
)
|
||||
updated = MagicMock()
|
||||
updated.model_dump.return_value = {"max_budget": 50.0}
|
||||
mock_prisma_client.update_data = AsyncMock(return_value={"data": updated})
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
_patch_team_keys_helpers(monkeypatch, hash_lookup={"tok-a": "hashed-a"})
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
AsyncMock(),
|
||||
)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
key_ids=["tok-a", "tok-foreign"],
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
response = await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert response.total_requested == 2
|
||||
assert [u.key for u in response.successful_updates] == ["tok-a"]
|
||||
assert [u.key for u in response.failed_updates] == ["tok-foreign"]
|
||||
assert "not found in team" in response.failed_updates[0].failed_reason
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_team_admin_authorized(monkeypatch):
|
||||
"""Non-admin caller passes the upfront team-permission check and succeeds."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
key_a = _make_team_key("tok-a")
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[key_a]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=key_a
|
||||
)
|
||||
updated = MagicMock()
|
||||
updated.model_dump.return_value = {"max_budget": 50.0}
|
||||
mock_prisma_client.update_data = AsyncMock(return_value={"data": updated})
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
_patch_team_keys_helpers(monkeypatch, hash_lookup={"tok-a": "hashed-a"})
|
||||
|
||||
auth_check = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
auth_check,
|
||||
)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
response = await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
api_key="sk-team-admin",
|
||||
user_id="team-admin-user",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert len(response.successful_updates) == 1
|
||||
# upfront check (1) + per-key check inside _process_single_key_update (1)
|
||||
assert auth_check.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_team_admin_no_permission(monkeypatch):
|
||||
"""Non-admin without KEY_UPDATE permission raises ProxyException upfront (no partial result)."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
key_a = _make_team_key("tok-a")
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[key_a]
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
AsyncMock(
|
||||
side_effect=ProxyException(
|
||||
message="not in team",
|
||||
type="team_member_permission_error",
|
||||
param="/key/update",
|
||||
code=401,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException):
|
||||
await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
api_key="sk-outsider",
|
||||
user_id="outsider",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
mock_prisma_client.update_data.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_xor_validation():
|
||||
"""Both selection modes set, or neither set, raises at request construction."""
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
BulkUpdateTeamKeysRequest(
|
||||
team_id="t",
|
||||
key_ids=["k"],
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=10.0),
|
||||
)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
BulkUpdateTeamKeysRequest(
|
||||
team_id="t",
|
||||
update_fields=KeyUpdateFields(max_budget=10.0),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_forbidden_fields():
|
||||
"""update_fields.team_id / .key / .key_alias are rejected on construction."""
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
KeyUpdateFields(team_id="other-team")
|
||||
with pytest.raises(Exception):
|
||||
KeyUpdateFields(key="sk-abc")
|
||||
with pytest.raises(Exception):
|
||||
KeyUpdateFields(key_alias="alias")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_batch_size_cap(monkeypatch):
|
||||
"""find_many returning > MAX_BATCH_SIZE rows raises 400."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
too_many = [_make_team_key(f"tok-{i}") for i in range(501)]
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=too_many
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-abc",
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 400
|
||||
assert "more than 500" in exc.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_team_keys_no_keys_found(monkeypatch):
|
||||
"""find_many returns nothing → 404 (clear signal, not silent empty success)."""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
bulk_update_team_keys,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateTeamKeysRequest,
|
||||
KeyUpdateFields,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
|
||||
|
||||
request = BulkUpdateTeamKeysRequest(
|
||||
team_id="team-empty",
|
||||
all_keys_in_team=True,
|
||||
update_fields=KeyUpdateFields(max_budget=50.0),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await bulk_update_team_keys(
|
||||
data=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 404
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue