mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(key_management_endpoints.py): add guaranteed throughput support for model specific tpm / rpm limits
prevents admin from granting keys more tpm/rpm than created for a key
This commit is contained in:
parent
a83238a2db
commit
3ce074564e
1 changed files with 109 additions and 9 deletions
|
|
@ -27,6 +27,7 @@ from litellm.caching import DualCache
|
|||
from litellm.constants import LENGTH_OF_LITELLM_GENERATED_KEY, UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_key_object,
|
||||
_delete_cache_key_object,
|
||||
|
|
@ -633,20 +634,93 @@ async def _common_key_generation_helper( # noqa: PLR0915
|
|||
return response
|
||||
|
||||
|
||||
async def _check_team_key_limits(
|
||||
def check_team_key_model_specific_limits(
|
||||
keys: List[LiteLLM_VerificationToken],
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: GenerateKeyRequest,
|
||||
prisma_client: PrismaClient,
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating guaranteed throughput limits. If so, raise an error if we're overallocating.
|
||||
Check if the team key is allocating model specific limits. If so, raise an error if we're overallocating.
|
||||
"""
|
||||
if data.model_rpm_limit is None and data.model_tpm_limit is None:
|
||||
return
|
||||
# get total model specific tpm/rpm limit
|
||||
model_specific_rpm_limit = {}
|
||||
model_specific_tpm_limit = {}
|
||||
|
||||
for key in keys:
|
||||
if key.metadata.get("model_rpm_limit", None) is not None:
|
||||
for model, rpm_limit in key.metadata.get("model_rpm_limit", {}).items():
|
||||
model_specific_rpm_limit[model] = (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
)
|
||||
if key.metadata.get("model_tpm_limit", None) is not None:
|
||||
for model, tpm_limit in key.metadata.get("model_tpm_limit", {}).items():
|
||||
model_specific_tpm_limit[model] = (
|
||||
model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
)
|
||||
if data.model_rpm_limit is not None:
|
||||
for model, rpm_limit in data.model_rpm_limit.items():
|
||||
if (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
> team_table.rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated RPM limit={model_specific_rpm_limit.get(model, 0)} + Key RPM limit={rpm_limit} is greater than team RPM limit={team_table.rpm_limit}",
|
||||
)
|
||||
elif team_table.metadata and team_table.metadata.get("model_rpm_limit"):
|
||||
team_model_specific_rpm_limit_dict = team_table.metadata.get(
|
||||
"model_rpm_limit", {}
|
||||
)
|
||||
team_model_specific_rpm_limit = team_model_specific_rpm_limit_dict.get(
|
||||
model
|
||||
)
|
||||
if (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
> team_model_specific_rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated RPM limit={model_specific_rpm_limit.get(model, 0)} + Key RPM limit={rpm_limit} is greater than team RPM limit={team_model_specific_rpm_limit.get(model, 0)}",
|
||||
)
|
||||
if data.model_tpm_limit is not None:
|
||||
for model, tpm_limit in data.model_tpm_limit.items():
|
||||
if (
|
||||
team_table.tpm_limit is not None
|
||||
and model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
> team_table.tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated TPM limit={model_specific_tpm_limit.get(model, 0)} + Key TPM limit={tpm_limit} is greater than team TPM limit={team_table.tpm_limit}",
|
||||
)
|
||||
elif team_table.metadata and team_table.metadata.get("model_tpm_limit"):
|
||||
team_model_specific_tpm_limit_dict = team_table.metadata.get(
|
||||
"model_tpm_limit", {}
|
||||
)
|
||||
team_model_specific_tpm_limit = team_model_specific_tpm_limit_dict.get(
|
||||
model
|
||||
)
|
||||
if (
|
||||
team_model_specific_tpm_limit
|
||||
and model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
> team_model_specific_tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated TPM limit={model_specific_tpm_limit.get(model, 0)} + Key TPM limit={tpm_limit} is greater than team TPM limit={team_model_specific_tpm_limit}",
|
||||
)
|
||||
|
||||
|
||||
def check_team_key_rpm_tpm_limits(
|
||||
keys: List[LiteLLM_VerificationToken],
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: GenerateKeyRequest,
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating rpm/tpm limits. If so, raise an error if we're overallocating.
|
||||
"""
|
||||
# get all team keys
|
||||
# calculate allocated tpm/rpm limit
|
||||
# check if specified tpm/rpm limit is greater than allocated tpm/rpm limit
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": team_table.team_id},
|
||||
)
|
||||
if keys is not None and len(keys) > 0:
|
||||
allocated_tpm = sum(key.tpm_limit for key in keys if key.tpm_limit is not None)
|
||||
allocated_rpm = sum(key.rpm_limit for key in keys if key.rpm_limit is not None)
|
||||
|
|
@ -673,6 +747,32 @@ async def _check_team_key_limits(
|
|||
)
|
||||
|
||||
|
||||
async def _check_team_key_limits(
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: GenerateKeyRequest,
|
||||
prisma_client: PrismaClient,
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating guaranteed throughput limits. If so, raise an error if we're overallocating.
|
||||
"""
|
||||
# get all team keys
|
||||
# calculate allocated tpm/rpm limit
|
||||
# check if specified tpm/rpm limit is greater than allocated tpm/rpm limit
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": team_table.team_id},
|
||||
)
|
||||
check_team_key_model_specific_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
check_team_key_rpm_tpm_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/key/generate",
|
||||
tags=["key management"],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue