From 0d1d8bb598d8e2331af433197359b24a9f6d6c9f Mon Sep 17 00:00:00 2001 From: Harshit28j Date: Wed, 4 Mar 2026 21:06:29 +0530 Subject: [PATCH] fix: model tpm/rpm, softbudget, cache invalidation --- litellm/proxy/auth/auth_checks.py | 62 ++++++++- .../proxy/hooks/model_max_budget_limiter.py | 126 ++++++++++++++++++ .../budget_management_endpoints.py | 4 + .../customer_endpoints.py | 96 ++++++++----- .../test_end_user_soft_budget_check.py | 52 ++++++++ .../test_customer_endpoints.py | 94 +++++++++---- 6 files changed, 372 insertions(+), 62 deletions(-) create mode 100644 tests/proxy_unit_tests/test_end_user_soft_budget_check.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 41b0a0bc38f..b57b767cc92 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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], diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 5e48ef2879e..accb3890576 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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( diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 20c7f9ec412..0705225bb9b 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index bfafc943c45..d9c14ae1692 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -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, - ) \ No newline at end of file + ) diff --git a/tests/proxy_unit_tests/test_end_user_soft_budget_check.py b/tests/proxy_unit_tests/test_end_user_soft_budget_check.py new file mode 100644 index 00000000000..84f37438b34 --- /dev/null +++ b/tests/proxy_unit_tests/test_end_user_soft_budget_check.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 25ff6f89427..dd565df0bc4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -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] )