mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): warn when key limits exceed team caps on generate/update
Compare rpm/tpm/max_parallel_requests/max_budget against the team on /key/generate and /key/update, and return field-level warnings when the key asks higher than the team cap, without rejecting the write yet
This commit is contained in:
parent
30f33a949b
commit
9c890f658d
3 changed files with 288 additions and 7 deletions
|
|
@ -1190,6 +1190,17 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
|
|||
return v
|
||||
|
||||
|
||||
KeyTeamLimitField = Literal["rpm_limit", "tpm_limit", "max_parallel_requests", "max_budget"]
|
||||
|
||||
|
||||
class KeyTeamLimitWarning(TypedDict):
|
||||
"""Non-blocking warning when a key limit exceeds its team's effective cap."""
|
||||
|
||||
field: ReadOnly[KeyTeamLimitField]
|
||||
requested: ReadOnly[float | int]
|
||||
effective_team_cap: ReadOnly[float | int]
|
||||
|
||||
|
||||
class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase):
|
||||
index_name: str
|
||||
index_permissions: list[Literal["read", "write"]]
|
||||
|
|
@ -1267,6 +1278,7 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
updated_by: str | None = None
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
warnings: list[KeyTeamLimitWarning] | None = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -1501,6 +1501,33 @@ def check_team_key_rpm_tpm_limits(
|
|||
)
|
||||
|
||||
|
||||
def _collect_key_team_limit_warnings(
|
||||
data: GenerateKeyRequest | UpdateKeyRequest,
|
||||
team_table: LiteLLM_TeamTable | LiteLLM_TeamTableCachedObj,
|
||||
) -> tuple[KeyTeamLimitWarning, ...]:
|
||||
"""
|
||||
Compare key rpm/tpm/max_parallel_requests/max_budget against the team's caps.
|
||||
|
||||
Returns warnings when the key requests a higher value than the team allows.
|
||||
Does not reject; runtime still applies the stricter team limit.
|
||||
"""
|
||||
comparisons: Final[tuple[tuple[KeyTeamLimitField, float | int | None, float | int | None], ...]] = (
|
||||
("rpm_limit", data.rpm_limit, team_table.rpm_limit),
|
||||
("tpm_limit", data.tpm_limit, team_table.tpm_limit),
|
||||
("max_parallel_requests", data.max_parallel_requests, team_table.max_parallel_requests),
|
||||
("max_budget", data.max_budget, team_table.max_budget),
|
||||
)
|
||||
return tuple(
|
||||
KeyTeamLimitWarning(
|
||||
field=field_name,
|
||||
requested=requested,
|
||||
effective_team_cap=team_cap,
|
||||
)
|
||||
for field_name, requested, team_cap in comparisons
|
||||
if requested is not None and team_cap is not None and requested > team_cap
|
||||
)
|
||||
|
||||
|
||||
async def _check_team_key_limits(
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: GenerateKeyRequest | UpdateKeyRequest,
|
||||
|
|
@ -1924,12 +1951,18 @@ async def generate_key_fn(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
return await _common_key_generation_helper(
|
||||
team_limit_warnings: Final = (
|
||||
_collect_key_team_limit_warnings(data=data, team_table=team_table) if team_table is not None else ()
|
||||
)
|
||||
response: Final = await _common_key_generation_helper(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
team_table=team_table,
|
||||
)
|
||||
if team_limit_warnings:
|
||||
response.warnings = list(team_limit_warnings)
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.generate_key_fn(): Exception occured - %s", e)
|
||||
|
|
@ -2090,12 +2123,18 @@ async def generate_service_account_key_fn(
|
|||
|
||||
data.user_id = None # do not allow user_id to be set for service account keys
|
||||
|
||||
return await _common_key_generation_helper(
|
||||
team_limit_warnings: Final = (
|
||||
_collect_key_team_limit_warnings(data=data, team_table=team_table) if team_table is not None else ()
|
||||
)
|
||||
response: Final = await _common_key_generation_helper(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
team_table=team_table,
|
||||
)
|
||||
if team_limit_warnings:
|
||||
response.warnings = list(team_limit_warnings)
|
||||
return response
|
||||
|
||||
|
||||
def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_metadata: dict) -> dict:
|
||||
|
|
@ -2538,9 +2577,10 @@ async def _process_single_key_update(
|
|||
|
||||
# Get team object and check team limits if team_id is provided
|
||||
team_obj: LiteLLM_TeamTableCachedObj | None = None
|
||||
if update_key_request.team_id is not None:
|
||||
_team_id_to_check: Final = update_key_request.team_id or getattr(existing_key_row, "team_id", None)
|
||||
if _team_id_to_check is not None:
|
||||
team_obj = await get_team_object(
|
||||
team_id=update_key_request.team_id,
|
||||
team_id=_team_id_to_check,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_db_only=True,
|
||||
|
|
@ -2631,6 +2671,11 @@ async def _process_single_key_update(
|
|||
|
||||
updated_key_info.pop("token", None)
|
||||
|
||||
team_limit_warnings: Final = (
|
||||
_collect_key_team_limit_warnings(data=update_key_request, team_table=team_obj) if team_obj is not None else ()
|
||||
)
|
||||
if team_limit_warnings:
|
||||
return {**updated_key_info, "warnings": list(team_limit_warnings)}
|
||||
return updated_key_info
|
||||
|
||||
|
||||
|
|
@ -2688,7 +2733,7 @@ async def _validate_update_key_data(
|
|||
premium_user: bool,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> None:
|
||||
) -> tuple[KeyTeamLimitWarning, ...]:
|
||||
"""Validate permissions and constraints for key update."""
|
||||
checked_prisma_client: Final = _require_prisma_client(prisma_client)
|
||||
|
||||
|
|
@ -2946,6 +2991,10 @@ async def _validate_update_key_data(
|
|||
if normalized_object_permission is not None:
|
||||
data.object_permission = LiteLLM_ObjectPermissionBase(**normalized_object_permission)
|
||||
|
||||
if team_obj is None:
|
||||
return ()
|
||||
return _collect_key_team_limit_warnings(data=data, team_table=team_obj)
|
||||
|
||||
|
||||
@router.post("/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)])
|
||||
@management_endpoint_wrapper
|
||||
|
|
@ -3067,7 +3116,7 @@ async def update_key_fn(
|
|||
key: Final = _resolve_token_to_update(data=data, existing_key_row=existing_key_row)
|
||||
data.key = key
|
||||
|
||||
await _validate_update_key_data(
|
||||
team_limit_warnings: Final = await _validate_update_key_data(
|
||||
data=data,
|
||||
existing_key_row=existing_key_row,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -3177,7 +3226,10 @@ async def update_key_fn(
|
|||
if response is None:
|
||||
raise ValueError("Failed to update key got response = None")
|
||||
|
||||
return {"key": key, **response["data"]}
|
||||
updated_key_info: Final[dict[str, object]] = {"key": key, **response["data"]}
|
||||
if team_limit_warnings:
|
||||
return {**updated_key_info, "warnings": list(team_limit_warnings)}
|
||||
return updated_key_info
|
||||
# update based on remaining passed in values
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_key_fn(): Exception occured - %s", e)
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
_check_org_key_limits,
|
||||
_check_project_key_limits,
|
||||
_check_team_key_limits,
|
||||
_collect_key_team_limit_warnings,
|
||||
_common_key_generation_helper,
|
||||
_enforce_upperbound_key_params,
|
||||
_get_and_validate_existing_key,
|
||||
|
|
@ -3177,6 +3178,222 @@ async def test_update_key_fn_auto_rotate_disable():
|
|||
assert result["auto_rotate"] is False
|
||||
|
||||
|
||||
def test_collect_key_team_limit_warnings_above_team_caps():
|
||||
"""Key limits above team caps produce field-level warnings without rejection."""
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="capped-team",
|
||||
team_alias="capped-team",
|
||||
rpm_limit=60,
|
||||
tpm_limit=1000,
|
||||
max_parallel_requests=5,
|
||||
max_budget=10.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
data = GenerateKeyRequest(
|
||||
team_id="capped-team",
|
||||
rpm_limit=600,
|
||||
tpm_limit=5000,
|
||||
max_parallel_requests=20,
|
||||
max_budget=100.0,
|
||||
)
|
||||
|
||||
warnings = _collect_key_team_limit_warnings(data=data, team_table=team_table)
|
||||
|
||||
assert warnings == (
|
||||
{
|
||||
"field": "rpm_limit",
|
||||
"requested": 600,
|
||||
"effective_team_cap": 60,
|
||||
},
|
||||
{
|
||||
"field": "tpm_limit",
|
||||
"requested": 5000,
|
||||
"effective_team_cap": 1000,
|
||||
},
|
||||
{
|
||||
"field": "max_parallel_requests",
|
||||
"requested": 20,
|
||||
"effective_team_cap": 5,
|
||||
},
|
||||
{
|
||||
"field": "max_budget",
|
||||
"requested": 100.0,
|
||||
"effective_team_cap": 10.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_collect_key_team_limit_warnings_within_or_unset_caps():
|
||||
"""No warning when key is within team caps, team has no cap, or key omits the field."""
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="partial-caps",
|
||||
team_alias="partial-caps",
|
||||
rpm_limit=60,
|
||||
tpm_limit=None,
|
||||
max_parallel_requests=5,
|
||||
max_budget=None,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
data = UpdateKeyRequest(
|
||||
key="sk-test-key-123456",
|
||||
rpm_limit=60,
|
||||
tpm_limit=999999,
|
||||
max_parallel_requests=3,
|
||||
max_budget=50.0,
|
||||
)
|
||||
|
||||
warnings = _collect_key_team_limit_warnings(data=data, team_table=team_table)
|
||||
|
||||
assert warnings == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_fn_attaches_team_limit_warnings(monkeypatch):
|
||||
"""/key/generate succeeds and returns warnings when key rpm exceeds team rpm."""
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="warn-team",
|
||||
team_alias="warn-team",
|
||||
rpm_limit=60,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
data = GenerateKeyRequest(team_id="warn-team", rpm_limit=600, key_alias="over-cap")
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-1234",
|
||||
user_id="admin",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=team_table),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.key_generation_check",
|
||||
MagicMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits",
|
||||
AsyncMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.check_org_admin_can_generate_keys",
|
||||
AsyncMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
MagicMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache",
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
from litellm.proxy._types import GenerateKeyResponse
|
||||
|
||||
generated = GenerateKeyResponse(
|
||||
key="sk-generated-key-123456",
|
||||
token_id="hashed",
|
||||
team_id="warn-team",
|
||||
rpm_limit=600,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._common_key_generation_helper",
|
||||
AsyncMock(return_value=generated),
|
||||
)
|
||||
|
||||
result = await generate_key_fn(data=data, user_api_key_dict=user_api_key_dict)
|
||||
|
||||
assert result.key == "sk-generated-key-123456"
|
||||
assert result.rpm_limit == 600
|
||||
assert result.warnings == [
|
||||
{
|
||||
"field": "rpm_limit",
|
||||
"requested": 600,
|
||||
"effective_team_cap": 60,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_update_key_data_returns_team_limit_warnings(monkeypatch):
|
||||
"""/key/update validation returns warnings when updated limits exceed team caps."""
|
||||
existing_key = LiteLLM_VerificationToken(
|
||||
token="hashed-token",
|
||||
team_id="warn-team",
|
||||
user_id="user-1",
|
||||
models=[],
|
||||
)
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="warn-team",
|
||||
team_alias="warn-team",
|
||||
rpm_limit=60,
|
||||
max_budget=10.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
data = UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600, max_budget=100.0)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-1234",
|
||||
user_id="admin",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=team_table),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits",
|
||||
AsyncMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
|
||||
AsyncMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.common_key_access_checks",
|
||||
MagicMock(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.enforce_member_can_assign_access_groups",
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
warnings = await _validate_update_key_data(
|
||||
data=data,
|
||||
existing_key_row=existing_key,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=None,
|
||||
premium_user=True,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
)
|
||||
|
||||
assert warnings == (
|
||||
{
|
||||
"field": "rpm_limit",
|
||||
"requested": 600,
|
||||
"effective_team_cap": 60,
|
||||
},
|
||||
{
|
||||
"field": "max_budget",
|
||||
"requested": 100.0,
|
||||
"effective_team_cap": 10.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_no_existing_keys():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue