mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: model tpm/rpm, softbudget, cache invalidation
This commit is contained in:
parent
edbb8ce360
commit
0d1d8bb598
6 changed files with 372 additions and 62 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
52
tests/proxy_unit_tests/test_end_user_soft_budget_check.py
Normal file
52
tests/proxy_unit_tests/test_end_user_soft_budget_check.py
Normal 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
|
||||
|
|
@ -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]
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue