This commit is contained in:
leilei3167 2026-09-27 20:03:03 +08:00 • committed by GitHub
commit 0956e5e472
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 587 additions and 8 deletions

View file

@ -1,7 +1,7 @@
import enum
import json
import os
from collections.abc import Callable, Mapping
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, TypeAlias
@ -1262,6 +1262,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"]]
@ -1341,6 +1352,7 @@ class GenerateKeyResponse(KeyRequestBase):
updated_by: str | None = None
created_at: datetime | None = None
updated_at: datetime | None = None
warnings: Sequence[KeyTeamLimitWarning] | None = None
@model_validator(mode="before")
@classmethod

View file

@ -1713,6 +1713,101 @@ def check_team_key_rpm_tpm_limits(
)
_KEY_TEAM_LIMIT_WARNING_FIELDS: Final[tuple[str, ...]] = (
"rpm_limit",
"tpm_limit",
"max_parallel_requests",
"max_budget",
)
def _update_request_with_retained_team_limits(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
) -> UpdateKeyRequest:
"""Fill omitted limit fields from the existing key for team-cap warnings.
/key/update often changes only team_id (or a subset of limits). Retained
rpm/tpm/concurrency/budget must still be compared against the effective
team caps so reassignment cannot silently keep an over-cap value.
"""
retained: Final = MappingProxyType(
{
field_name: getattr(existing_key_row, field_name, None)
for field_name in _KEY_TEAM_LIMIT_WARNING_FIELDS
if field_name not in data.model_fields_set
}
)
if not retained:
return data
return data.model_copy(update=retained)
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
)
def _maybe_add_key_team_limit_warnings(
payload: Mapping[str, object],
data: GenerateKeyRequest | UpdateKeyRequest,
team_table: LiteLLM_TeamTable | LiteLLM_TeamTableCachedObj | None,
) -> Mapping[str, object]:
"""Attach team-limit warnings to a key update payload when caps are exceeded."""
if team_table is None:
return payload
warnings = _collect_key_team_limit_warnings(data=data, team_table=team_table)
if not warnings:
return payload
return MappingProxyType(
{**payload, "warnings": list(warnings)}
) # mutable-ok: GenerateKeyResponse/tests expect list warnings
async def _soft_resolve_existing_key_team_for_warnings(
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
) -> LiteLLM_TeamTableCachedObj | None:
"""Resolve the existing key's team for warning payloads only; ignore missing teams."""
team_id = getattr(existing_key_row, "team_id", None)
if team_id is None or prisma_client is None:
return None
try:
return await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
except HTTPException as e:
if e.status_code != status.HTTP_404_NOT_FOUND:
raise
return None
async def _check_team_key_limits(
team_table: LiteLLM_TeamTableCachedObj,
data: GenerateKeyRequest | UpdateKeyRequest,
@ -2138,12 +2233,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) # mutable-ok: GenerateKeyResponse/tests expect list warnings
return response
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.generate_key_fn(): Exception occured - %s", e)
@ -2314,12 +2415,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) # mutable-ok: GenerateKeyResponse/tests expect list warnings
return response
def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_metadata: dict) -> dict:
@ -2780,7 +2887,9 @@ async def _process_single_key_update(
# Enforce upperbound key params on update (don't fill defaults)
_enforce_upperbound_key_params(update_key_request, fill_defaults=False)
# Get team object and check team limits if team_id is provided
# Get team object and check team limits if team_id is provided on the request.
# Existing-key team is soft-resolved for warnings only — a missing team must not
# block an otherwise valid update (custom key policy runs later).
team_obj: LiteLLM_TeamTableCachedObj | None = None
if update_key_request.team_id is not None:
team_obj = await get_team_object(
@ -2796,6 +2905,12 @@ async def _process_single_key_update(
data=update_key_request,
prisma_client=prisma_client,
)
else:
team_obj = await _soft_resolve_existing_key_team_for_warnings(
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
# Validate team change if team is being changed
if is_different_team(data=update_key_request, existing_key_row=existing_key_row):
@ -2906,7 +3021,11 @@ async def _process_single_key_update(
updated_key_info.pop("token", None)
return updated_key_info
return _maybe_add_key_team_limit_warnings(
updated_key_info,
update_key_request,
team_obj,
)
async def _with_validated_object_permission(
@ -3071,7 +3190,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)
@ -3353,6 +3472,13 @@ 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=_update_request_with_retained_team_limits(data=data, existing_key_row=existing_key_row),
team_table=team_obj,
)
@router.post("/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)])
@management_endpoint_wrapper
@ -3476,7 +3602,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,
@ -3598,7 +3724,12 @@ 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[Mapping[str, object]] = MappingProxyType({"key": key, **response["data"]})
if team_limit_warnings:
return MappingProxyType(
{**updated_key_info, "warnings": list(team_limit_warnings)}
) # mutable-ok: key/update response tests expect list 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)

View file

@ -48,6 +48,8 @@ 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,
_maybe_add_key_team_limit_warnings,
_common_key_generation_helper,
_effective_key_after_update,
_effective_key_for_generate,
@ -70,6 +72,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
check_team_key_model_specific_limits,
delete_verification_tokens,
generate_key_fn,
generate_service_account_key_fn,
generate_key_helper_fn,
key_aliases,
key_generation_check,
@ -3577,6 +3580,420 @@ 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,
},
)
async def test_validate_update_key_data_warns_on_retained_limits_team_change(monkeypatch):
"""Team reassignment without limit fields still warns on retained over-cap values."""
existing = LiteLLM_VerificationToken(
token="hashed-token",
team_id="team-old",
user_id="user-1",
models=[],
max_budget=100.0,
max_parallel_requests=50,
)
new_team = LiteLLM_TeamTableCachedObj(
team_id="team-new",
team_alias="team-new",
max_budget=10.0,
max_parallel_requests=5,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
data = UpdateKeyRequest(key="sk-test-key-123456", team_id="team-new")
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=new_team),
)
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(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_team_change",
AsyncMock(),
)
warnings = await _validate_update_key_data(
data=data,
existing_key_row=existing,
user_api_key_dict=user_api_key_dict,
llm_router=MagicMock(),
premium_user=True,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert warnings == (
{
"field": "max_parallel_requests",
"requested": 50,
"effective_team_cap": 5,
},
{
"field": "max_budget",
"requested": 100.0,
"effective_team_cap": 10.0,
},
)
def test_maybe_add_key_team_limit_warnings_passthrough_and_attach():
"""Attach warnings to update payloads only when caps are exceeded."""
team_table = LiteLLM_TeamTableCachedObj(
team_id="warn-team",
team_alias="warn-team",
rpm_limit=60,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
payload = {"key": "sk-test-key-123456", "rpm_limit": 600}
assert (
_maybe_add_key_team_limit_warnings(
payload,
UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600),
None,
)
is payload
)
assert (
_maybe_add_key_team_limit_warnings(
payload,
UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=30),
team_table,
)
is payload
)
with_warnings = _maybe_add_key_team_limit_warnings(
payload,
UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600),
team_table,
)
assert with_warnings is not payload
assert with_warnings["key"] == "sk-test-key-123456"
assert with_warnings["warnings"] == [
{
"field": "rpm_limit",
"requested": 600,
"effective_team_cap": 60,
}
]
@pytest.mark.asyncio
async def test_generate_service_account_key_fn_attaches_team_limit_warnings(monkeypatch):
"""Service-account generate also surfaces team-limit warnings."""
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="sa-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.check_org_admin_can_generate_keys",
AsyncMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_team_id_used_in_service_account_request",
AsyncMock(),
)
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.key_generation_check",
MagicMock(),
)
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-service-account-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_service_account_key_fn(
data=data, user_api_key_dict=user_api_key_dict
)
assert result.key == "sk-service-account-123456"
assert result.warnings == [
{
"field": "rpm_limit",
"requested": 600,
"effective_team_cap": 60,
}
]
@pytest.mark.asyncio
async def test_check_team_key_limits_no_existing_keys():
"""

View file

@ -30920,6 +30920,8 @@ export interface components {
updated_by?: string | null;
/** User Id */
user_id?: string | null;
/** Warnings */
warnings?: components["schemas"]["KeyTeamLimitWarning"][] | null;
};
/** GenericGuardrailAPIInputs */
GenericGuardrailAPIInputs: {
@ -31659,6 +31661,21 @@ export interface components {
/** Keys */
keys?: string[] | null;
};
/**
* KeyTeamLimitWarning
* @description Non-blocking warning when a key limit exceeds its team's effective cap.
*/
KeyTeamLimitWarning: {
/** Effective Team Cap */
effective_team_cap: number;
/**
* Field
* @enum {string}
*/
field: "rpm_limit" | "tpm_limit" | "max_parallel_requests" | "max_budget";
/** Requested */
requested: number;
};
/**
* KeyUpdateFields
* @description Allowlist of bulk-broadcastable fields for /team/key/bulk_update; `extra="forbid"` blocks RBAC/ownership/scope mutations even by team admins.
@ -37050,6 +37067,8 @@ export interface components {
user_id?: string | null;
/** User Role */
user_role?: ("proxy_admin" | "proxy_admin_viewer" | "internal_user" | "internal_user_viewer") | null;
/** Warnings */
warnings?: components["schemas"]["KeyTeamLimitWarning"][] | null;
};
/**
* OAuth2SecurityScheme