fix: model tpm/rpm, softbudget, cache invalidation

This commit is contained in:
Harshit28j 2026-03-04 21:06:29 +05:30
parent edbb8ce360
commit 0d1d8bb598
6 changed files with 372 additions and 62 deletions

View file

@ -415,6 +415,12 @@ async def common_checks( # noqa: PLR0915
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
await _end_user_soft_budget_check(
end_user_object=end_user_object,
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
# 6. [OPTIONAL] If 'enforce_user_param' enabled - did developer pass in 'user' param for openai endpoints
if (
general_settings.get("enforce_user_param", None) is not None
@ -2119,13 +2125,11 @@ async def get_key_object(
)
# else, check db
_valid_token: Optional[BaseModel] = (
await _fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
_valid_token: Optional[BaseModel] = await _fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if _valid_token is None:
@ -2947,6 +2951,50 @@ async def _team_max_budget_check(
)
async def _end_user_soft_budget_check(
end_user_object: Optional[LiteLLM_EndUserTable],
valid_token: Optional[UserAPIKeyAuth],
proxy_logging_obj: ProxyLogging,
):
"""
Triggers a budget alert if the end user is over its soft budget.
"""
if (
end_user_object is not None
and end_user_object.litellm_budget_table is not None
and end_user_object.litellm_budget_table.soft_budget is not None
and end_user_object.spend >= end_user_object.litellm_budget_table.soft_budget
):
verbose_proxy_logger.debug(
"Crossed Soft Budget for end_user %s, spend %s, soft_budget %s",
end_user_object.user_id,
end_user_object.spend,
end_user_object.litellm_budget_table.soft_budget,
)
if valid_token:
call_info = CallInfo(
token=valid_token.token,
spend=end_user_object.spend,
max_budget=end_user_object.litellm_budget_table.max_budget,
soft_budget=end_user_object.litellm_budget_table.soft_budget,
user_id=valid_token.user_id,
customer_id=end_user_object.user_id,
team_id=valid_token.team_id,
team_alias=valid_token.team_alias,
organization_id=valid_token.org_id,
user_email=None,
key_alias=valid_token.key_alias,
event_group=Litellm_EntityType.END_USER,
)
asyncio.create_task(
proxy_logging_obj.budget_alerts(
type="soft_budget",
user_info=call_info,
)
)
async def _team_soft_budget_check(
team_object: Optional[LiteLLM_TeamTable],
valid_token: Optional[UserAPIKeyAuth],

View file

@ -82,6 +82,43 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
max_budget=_current_model_budget_info.max_budget,
)
from datetime import datetime
from fastapi import HTTPException
current_minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
if (
_current_model_budget_info.tpm_limit is not None
and _current_model_budget_info.tpm_limit > 0
):
key_model_tpm_cache_key = (
f"key_model_tpm:{user_api_key_dict.token}:{model}:{current_minute}"
)
_current_tpm_spend = (
await self.dual_cache.async_get_cache(key=key_model_tpm_cache_key) or 0
)
if _current_tpm_spend >= _current_model_budget_info.tpm_limit:
raise HTTPException(
status_code=429,
detail=f"LiteLLM Rate Limit: Key={user_api_key_dict.token} exceeded tpm limit for model={model}. Current TPM: {_current_tpm_spend}, Limit: {_current_model_budget_info.tpm_limit}",
)
if (
_current_model_budget_info.rpm_limit is not None
and _current_model_budget_info.rpm_limit > 0
):
key_model_rpm_cache_key = (
f"key_model_rpm:{user_api_key_dict.token}:{model}:{current_minute}"
)
_current_rpm_spend = (
await self.dual_cache.async_get_cache(key=key_model_rpm_cache_key) or 0
)
if _current_rpm_spend >= _current_model_budget_info.rpm_limit:
raise HTTPException(
status_code=429,
detail=f"LiteLLM Rate Limit: Key={user_api_key_dict.token} exceeded rpm limit for model={model}. Current RPM: {_current_rpm_spend}, Limit: {_current_model_budget_info.rpm_limit}",
)
return True
async def is_end_user_within_model_budget(
@ -137,6 +174,45 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
max_budget=_current_model_budget_info.max_budget,
)
from datetime import datetime
from fastapi import HTTPException
current_minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
if (
_current_model_budget_info.tpm_limit is not None
and _current_model_budget_info.tpm_limit > 0
):
end_user_model_tpm_cache_key = (
f"end_user_model_tpm:{end_user_id}:{model}:{current_minute}"
)
_current_tpm_spend = (
await self.dual_cache.async_get_cache(key=end_user_model_tpm_cache_key)
or 0
)
if _current_tpm_spend >= _current_model_budget_info.tpm_limit:
raise HTTPException(
status_code=429,
detail=f"LiteLLM Rate Limit: End User={end_user_id} exceeded tpm limit for model={model}. Current TPM: {_current_tpm_spend}, Limit: {_current_model_budget_info.tpm_limit}",
)
if (
_current_model_budget_info.rpm_limit is not None
and _current_model_budget_info.rpm_limit > 0
):
end_user_model_rpm_cache_key = (
f"end_user_model_rpm:{end_user_id}:{model}:{current_minute}"
)
_current_rpm_spend = (
await self.dual_cache.async_get_cache(key=end_user_model_rpm_cache_key)
or 0
)
if _current_rpm_spend >= _current_model_budget_info.rpm_limit:
raise HTTPException(
status_code=429,
detail=f"LiteLLM Rate Limit: End User={end_user_id} exceeded rpm limit for model={model}. Current RPM: {_current_rpm_spend}, Limit: {_current_model_budget_info.rpm_limit}",
)
return True
async def _get_end_user_spend_for_model(
@ -268,6 +344,18 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
if model is None:
return
total_tokens = 0
if (
isinstance(response_obj, litellm.ModelResponse)
and hasattr(response_obj, "usage")
and response_obj.usage
):
total_tokens = getattr(response_obj.usage, "total_tokens", 0)
from datetime import datetime
current_minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
if (
virtual_key is not None
and user_api_key_model_max_budget is not None
@ -289,6 +377,22 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
response_cost=response_cost,
)
if key_budget_config is not None:
if (
key_budget_config.tpm_limit is not None
and key_budget_config.tpm_limit > 0
):
tpm_key = f"key_model_tpm:{virtual_key}:{model}:{current_minute}"
await self.dual_cache.async_increment_cache(
key=tpm_key, value=total_tokens
)
if (
key_budget_config.rpm_limit is not None
and key_budget_config.rpm_limit > 0
):
rpm_key = f"key_model_rpm:{virtual_key}:{model}:{current_minute}"
await self.dual_cache.async_increment_cache(key=rpm_key, value=1)
if (
end_user_id is not None
and user_api_key_end_user_model_max_budget is not None
@ -310,6 +414,28 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
response_cost=response_cost,
)
if key_budget_config is not None:
if (
key_budget_config.tpm_limit is not None
and key_budget_config.tpm_limit > 0
):
end_user_tpm_key = (
f"end_user_model_tpm:{end_user_id}:{model}:{current_minute}"
)
await self.dual_cache.async_increment_cache(
key=end_user_tpm_key, value=total_tokens
)
if (
key_budget_config.rpm_limit is not None
and key_budget_config.rpm_limit > 0
):
end_user_rpm_key = (
f"end_user_model_rpm:{end_user_id}:{model}:{current_minute}"
)
await self.dual_cache.async_increment_cache(
key=end_user_rpm_key, value=1
)
verbose_proxy_logger.debug(
"current state of in memory cache %s",
json.dumps(

View file

@ -169,6 +169,10 @@ async def update_budget(
}, # type: ignore
)
from litellm.proxy.proxy_server import user_api_key_cache
user_api_key_cache.delete_cache(key=f"budget_id:{budget_obj.budget_id}")
return response

View file

@ -19,13 +19,15 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_daily_activity import \
get_daily_activity
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission, handle_update_object_permission_common)
from litellm.proxy.utils import handle_exception_on_proxy
from litellm.types.proxy.management_endpoints.common_daily_activity import \
SpendAnalyticsPaginatedResponse
_set_object_permission,
handle_update_object_permission_common,
)
from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
router = APIRouter()
@ -109,8 +111,9 @@ async def unblock_user(data: BlockUsers):
```
"""
try:
from enterprise.enterprise_hooks.blocked_user_list import \
_ENTERPRISE_BlockedUserList
from enterprise.enterprise_hooks.blocked_user_list import (
_ENTERPRISE_BlockedUserList,
)
except ImportError:
raise HTTPException(
status_code=400,
@ -289,8 +292,11 @@ async def new_end_user(
- end-user object
- currently allowed models
"""
from litellm.proxy.proxy_server import (litellm_proxy_admin_name,
llm_router, prisma_client)
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
prisma_client,
)
if prisma_client is None:
raise HTTPException(
@ -321,14 +327,17 @@ async def new_end_user(
_new_budget = new_budget_request(data)
if _new_budget is not None:
try:
budget_record = await prisma_client.db.litellm_budgettable.create(
data={
budget_data = jsonify_object(
{
**_new_budget.model_dump(exclude_unset=True),
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name, # type: ignore
"updated_by": user_api_key_dict.user_id
or litellm_proxy_admin_name,
}
)
budget_record = await prisma_client.db.litellm_budgettable.create(
data=budget_data
)
except Exception as e:
raise HTTPException(status_code=422, detail={"error": str(e)})
@ -366,7 +375,13 @@ async def new_end_user(
response_dict = end_user_record.model_dump()
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
@ -425,7 +440,8 @@ async def end_user_info(
)
user_info = await prisma_client.db.litellm_endusertable.find_first(
where={"user_id": end_user_id}, include={"litellm_budget_table": True, "object_permission": True}
where={"user_id": end_user_id},
include={"litellm_budget_table": True, "object_permission": True},
)
if user_info is None:
@ -440,11 +456,17 @@ async def end_user_info(
response_dict = user_info.model_dump(exclude_none=True)
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.customer_endpoints.end_user_info(): Exception occured - {}".format(
@ -453,6 +475,7 @@ async def end_user_info(
)
raise handle_exception_on_proxy(e)
@router.post(
"/customer/update",
tags=["Customer Management"],
@ -520,8 +543,7 @@ async def update_end_user(
```
"""
from litellm.proxy.proxy_server import (litellm_proxy_admin_name,
prisma_client)
from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client
try:
data_json: dict = data.json()
@ -535,8 +557,9 @@ async def update_end_user(
if v is not None and v not in (
[],
{},
0,
): # models default to [], spend defaults to 0, we should not reset these values
): # models default to [], we should not reset these values
if v == 0 and k == "spend":
continue # spend defaults to 0, skip to avoid resetting
non_default_values[k] = v
## Get end user table data ##
@ -634,11 +657,21 @@ async def update_end_user(
f"received response from updating prisma client. response={response}"
)
from litellm.proxy.proxy_server import user_api_key_cache
user_api_key_cache.delete_cache(key=f"end_user_id:{data.user_id}")
# Convert to dict and clean up recursive fields
response_dict = response.model_dump()
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
@ -744,6 +777,7 @@ async def delete_end_user(
)
raise handle_exception_on_proxy(e)
@router.get(
"/customer/list",
tags=["Customer Management"],
@ -801,11 +835,17 @@ async def list_end_user(
item_dict = item.model_dump()
# Remove reverse relations from object_permission
if item_dict.get("object_permission"):
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
item_dict["object_permission"].pop(field, None)
returned_response.append(LiteLLM_EndUserTable(**item_dict))
return returned_response
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.customer_endpoints.list_end_user(): Exception occured - {}".format(
@ -814,6 +854,7 @@ async def list_end_user(
)
raise handle_exception_on_proxy(e)
@router.get(
"/customer/daily/activity",
tags=["Customer Management"],
@ -837,7 +878,6 @@ async def get_customer_daily_activity(
exclude_end_user_ids: Optional[str] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get daily activity for specific organizations or all accessible organizations.
"""
@ -857,7 +897,6 @@ async def get_customer_daily_activity(
exclude_end_user_ids.split(",") if exclude_end_user_ids else None
)
# Fetch organization aliases for metadata
where_condition = {}
if end_user_ids_list:
@ -865,10 +904,7 @@ async def get_customer_daily_activity(
end_user_aliases = await prisma_client.db.litellm_endusertable.find_many(
where=where_condition
)
end_user_alias_metadata = {
e.user_id: {"alias": e.alias}
for e in end_user_aliases
}
end_user_alias_metadata = {e.user_id: {"alias": e.alias} for e in end_user_aliases}
# Query daily activity for organizations
return await get_daily_activity(
@ -884,4 +920,4 @@ async def get_customer_daily_activity(
api_key=api_key,
page=page,
page_size=page_size,
)
)

View file

@ -0,0 +1,52 @@
import pytest
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.auth.auth_checks import _end_user_soft_budget_check
from litellm.proxy._types import (
LiteLLM_EndUserTable,
LiteLLM_BudgetTable,
UserAPIKeyAuth,
)
@pytest.mark.asyncio
async def test_end_user_soft_budget_alert_triggers():
# Setup
mock_proxy_logging = MagicMock()
mock_proxy_logging.budget_alerts = AsyncMock()
end_user_obj = LiteLLM_EndUserTable(
user_id="end-user-soft-budget",
spend=15.0,
blocked=False,
litellm_budget_table=LiteLLM_BudgetTable(
budget_id="budget-1", max_budget=20.0, soft_budget=10.0
),
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
team_alias="team-alias",
org_id="org-id",
key_alias="key-alias",
)
# Execute
await _end_user_soft_budget_check(
end_user_object=end_user_obj,
valid_token=valid_token,
proxy_logging_obj=mock_proxy_logging,
)
# Let the background task run
await asyncio.sleep(0.1)
# Verify
mock_proxy_logging.budget_alerts.assert_called_once()
args, kwargs = mock_proxy_logging.budget_alerts.call_args
assert kwargs["type"] == "soft_budget"
assert kwargs["user_info"].customer_id == "end-user-soft-budget"
assert kwargs["user_info"].spend == 15.0
assert kwargs["user_info"].soft_budget == 10.0

View file

@ -82,6 +82,41 @@ def test_update_customer_success(mock_prisma_client, mock_user_api_key_auth):
assert response.json()["alias"] == "Updated Test User"
@patch("litellm.proxy.proxy_server.user_api_key_cache")
def test_update_end_user_cache_invalidation(
mock_user_api_key_cache, mock_prisma_client, mock_user_api_key_auth
):
"""
Test that updating an end user invalidates their cache entry.
"""
# Mock the database responses
mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-cache", alias="Test User", blocked=False
)
updated_mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-cache", alias="Updated Test User", blocked=False
)
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=mock_end_user
)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=updated_mock_end_user
)
test_data = {"user_id": "test-user-cache", "alias": "Updated Test User"}
response = client.post(
"/customer/update", json=test_data, headers={"Authorization": "Bearer test-key"}
)
assert response.status_code == 200
# Verify cache invalidation was called with the correct key
mock_user_api_key_cache.delete_cache.assert_called_once_with(
key="end_user_id:test-user-cache"
)
def test_update_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
"""
Test that update_end_user raises a 404 ProxyException when user_id does not exist.
@ -103,7 +138,10 @@ def test_update_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
assert response.status_code == 404
response_json = response.json()
assert "error" in response_json
assert response_json["error"]["message"] == "End User Id=non-existent-user does not exist in db"
assert (
response_json["error"]["message"]
== "End User Id=non-existent-user does not exist in db"
)
assert response_json["error"]["type"] == "not_found"
assert response_json["error"]["param"] == "user_id"
assert response_json["error"]["code"] == "404"
@ -126,7 +164,10 @@ def test_info_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
assert response.status_code == 404
response_json = response.json()
assert "error" in response_json
assert response_json["error"]["message"] == "End User Id=non-existent-user does not exist in db"
assert (
response_json["error"]["message"]
== "End User Id=non-existent-user does not exist in db"
)
assert response_json["error"]["type"] == "not_found"
assert response_json["error"]["param"] == "end_user_id"
assert response_json["error"]["code"] == "404"
@ -165,7 +206,7 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
Test that all customer endpoints return the same error schema format.
All ProxyException errors should have: message, type, param, and code fields.
"""
def validate_error_schema(response_json):
assert "error" in response_json, "Response should have 'error' key"
error = response_json["error"]
@ -212,7 +253,7 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
# Test /customer/new - duplicate user error
from unittest.mock import MagicMock
mock_end_user = LiteLLM_EndUserTable(
user_id="existing-user", alias="Existing User", blocked=False
)
@ -229,33 +270,34 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
assert error["code"] == "400"
def test_customer_endpoints_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
def test_customer_endpoints_error_schema_consistency(
mock_prisma_client, mock_user_api_key_auth
):
"""
Test the exact scenarios from the curl examples provided.
Scenario 1: GET /end_user/info with non-existent user
OLD (incorrect): {"detail":{"error":"End User Id=... does not exist in db"}}
NEW (correct): {"error":{"message":"...","type":"not_found","param":"end_user_id","code":"404"}}
Scenario 2: POST /end_user/new with existing user
Expected: {"error":{"message":"...","type":"bad_request","param":"user_id","code":"400"}}
Both should use the same error format structure.
"""
# Scenario 1: GET /end_user/info with non-existent user
# Should return 404 with proper error schema
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None)
response1 = client.get(
"/end_user/info?end_user_id=fake-test-end-user-michaels-local-testng",
headers={"Authorization": "Bearer test-key"},
)
assert response1.status_code == 404, "Should return 404 for non-existent user"
response1_json = response1.json()
# Should have the correct format with {"error": {...}}
assert "error" in response1_json, "Should have top-level 'error' key"
error1 = response1_json["error"]
@ -266,22 +308,25 @@ def test_customer_endpoints_error_schema_consistency(mock_prisma_client, mock_us
assert error1["type"] == "not_found"
assert error1["code"] == "404"
assert "does not exist in db" in error1["message"]
# Scenario 2: POST /end_user/new with existing user
# Should return 400 with proper error schema
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(
side_effect=Exception("Unique constraint failed on the fields: (`user_id`)")
)
response2 = client.post(
"/end_user/new",
json={"user_id": "fake-test-end-user-michaels-local-testing", "budget_id": "Tier0"},
json={
"user_id": "fake-test-end-user-michaels-local-testing",
"budget_id": "Tier0",
},
headers={"Authorization": "Bearer test-key"},
)
assert response2.status_code == 400, "Should return 400 for duplicate user"
response2_json = response2.json()
# Should have the same error structure as Scenario 1
assert "error" in response2_json, "Should have top-level 'error' key"
error2 = response2_json["error"]
@ -292,11 +337,12 @@ def test_customer_endpoints_error_schema_consistency(mock_prisma_client, mock_us
assert error2["type"] == "bad_request"
assert error2["code"] == "400"
assert "Customer already exists" in error2["message"]
# Verify both errors have the same schema structure
assert set(error1.keys()) == set(error2.keys()), \
"Both errors should have the same top-level keys"
assert set(error1.keys()) == set(
error2.keys()
), "Both errors should have the same top-level keys"
# Both should have string values for all fields
for key in ["message", "type", "code"]:
assert isinstance(error1[key], str), f"error1[{key}] should be a string"
@ -312,9 +358,7 @@ async def test_get_customer_daily_activity_admin_param_passing(monkeypatch):
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(
return_value=[]
)
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse")
@ -367,7 +411,7 @@ async def test_get_customer_daily_activity_with_end_user_aliases(monkeypatch):
mock_end_user2 = MagicMock()
mock_end_user2.user_id = "end-user-2"
mock_end_user2.alias = "Customer Two"
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(
return_value=[mock_end_user1, mock_end_user2]
)