From 96e75628d68d0fd9e8d38952ff20c8f41840e3f9 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 2 May 2025 07:06:07 -0700 Subject: [PATCH] [Fix + Refactor] Trigger Soft Budget Webhooks When Key Crosses Threshold (#10491) * fix slack alerting with webhooks * emit correct event group/entity on webhooks * refactor to use a common class of alerts with abc methods * fixes for tests * refactor to use a common class of alerts with abc methods * Send a budget alert on slack or webhook * unit test slack alerting * fix code qa --- litellm/integrations/SlackAlerting/Readme.md | 33 ++++ .../SlackAlerting/budget_alert_types.py | 93 +++++++++ .../SlackAlerting/slack_alerting.py | 183 ++++++++++++------ litellm/proxy/_types.py | 6 +- litellm/proxy/auth/auth_checks.py | 4 + litellm/proxy/auth/user_api_key_auth.py | 26 +-- .../health_endpoints/_health_endpoints.py | 4 +- .../proxy/hooks/key_management_event_hooks.py | 3 +- .../internal_user_endpoints.py | 36 ++-- litellm/proxy/proxy_config.yaml | 2 + litellm/proxy/proxy_server.py | 1 + .../SlackAlerting/test_slack_alerting.py | 163 ++++++++++++++++ tests/logging_callback_tests/test_alerting.py | 120 +++++++++++- 13 files changed, 574 insertions(+), 100 deletions(-) create mode 100644 litellm/integrations/SlackAlerting/budget_alert_types.py create mode 100644 tests/litellm/integrations/SlackAlerting/test_slack_alerting.py diff --git a/litellm/integrations/SlackAlerting/Readme.md b/litellm/integrations/SlackAlerting/Readme.md index f28f71500cc..1719941d0da 100644 --- a/litellm/integrations/SlackAlerting/Readme.md +++ b/litellm/integrations/SlackAlerting/Readme.md @@ -9,5 +9,38 @@ This folder contains the Slack Alerting integration for LiteLLM Gateway. - `types.py`: This file contains the AlertType enum which is used to define the different types of alerts that can be sent to Slack. - `utils.py`: This file contains common utils used specifically for slack alerting +## Budget Alert Types + +The `budget_alert_types.py` module provides a flexible framework for handling different types of budget alerts: + +- `BaseBudgetAlertType`: An abstract base class with abstract methods that all alert types must implement: + - `get_event_group()`: Returns the Litellm_EntityType for the alert + - `get_event_message()`: Returns the message prefix for the alert + - `get_id(user_info)`: Returns the ID to use for caching/tracking the alert + +Concrete implementations include: +- `ProxyBudgetAlert`: Alerting for proxy-level budget concerns +- `SoftBudgetAlert`: Alerting when soft budgets are crossed +- `UserBudgetAlert`: Alerting for user-level budget concerns +- `TeamBudgetAlert`: Alerting for team-level budget concerns +- `TokenBudgetAlert`: Alerting for API key budget concerns +- `ProjectedLimitExceededAlert`: Alerting when projected spend will exceed budget + +Use the `get_budget_alert_type()` factory function to get the appropriate alert type class for a given alert type string: + +```python +from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type + +# Get the appropriate handler +budget_alert_class = get_budget_alert_type("user_budget") + +# Use the handler methods +event_group = budget_alert_class.get_event_group() # Returns Litellm_EntityType.USER +event_message = budget_alert_class.get_event_message() # Returns "User Budget: " +cache_id = budget_alert_class.get_id(user_info) # Returns user_id +``` + +To add a new budget alert type, simply create a new class that extends `BaseBudgetAlertType` and implements all the required methods, then add it to the dictionary in the `get_budget_alert_type()` function. + ## Further Reading - [Doc setting up Alerting on LiteLLM Proxy (Gateway)](https://docs.litellm.ai/docs/proxy/alerting) \ No newline at end of file diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py new file mode 100644 index 00000000000..beebee8b6bf --- /dev/null +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -0,0 +1,93 @@ +from abc import ABC, abstractmethod +from typing import Literal + +from litellm.proxy._types import CallInfo + + +class BaseBudgetAlertType(ABC): + """Base class for different budget alert types""" + + @abstractmethod + def get_event_message(self) -> str: + """Return the event message for this alert type""" + pass + + @abstractmethod + def get_id(self, user_info: CallInfo) -> str: + """Return the ID to use for caching/tracking this alert""" + pass + + +class ProxyBudgetAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "Proxy Budget: " + + def get_id(self, user_info: CallInfo) -> str: + return "default_id" + + +class SoftBudgetAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "Soft Budget Crossed: " + + def get_id(self, user_info: CallInfo) -> str: + return "default_id" + + +class UserBudgetAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "User Budget: " + + def get_id(self, user_info: CallInfo) -> str: + return user_info.user_id or "default_id" + + +class TeamBudgetAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "Team Budget: " + + def get_id(self, user_info: CallInfo) -> str: + return user_info.team_id or "default_id" + + +class TokenBudgetAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "Key Budget: " + + def get_id(self, user_info: CallInfo) -> str: + return user_info.token or "default_id" + + +class ProjectedLimitExceededAlert(BaseBudgetAlertType): + def get_event_message(self) -> str: + return "Key Budget: Projected Limit Exceeded" + + def get_id(self, user_info: CallInfo) -> str: + return user_info.token or "default_id" + + +def get_budget_alert_type( + type: Literal[ + "token_budget", + "soft_budget", + "user_budget", + "team_budget", + "proxy_budget", + "projected_limit_exceeded", + ], +) -> BaseBudgetAlertType: + """Factory function to get the appropriate budget alert type class""" + + alert_types = { + "proxy_budget": ProxyBudgetAlert(), + "soft_budget": SoftBudgetAlert(), + "user_budget": UserBudgetAlert(), + "team_budget": TeamBudgetAlert(), + "token_budget": TokenBudgetAlert(), + "projected_limit_exceeded": ProjectedLimitExceededAlert(), + } + + if type in alert_types: + return alert_types[type] + else: + return ProxyBudgetAlert() diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 9fde042ae79..7e7aa4d370e 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -6,7 +6,7 @@ import os import random import time from datetime import timedelta -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union from openai import APIError @@ -18,6 +18,7 @@ from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.constants import HOURS_IN_A_DAY from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.litellm_core_utils.exception_mapping_utils import ( _add_key_name_and_team_to_alert, @@ -26,7 +27,13 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._types import AlertType, CallInfo, VirtualKeyEvent, WebhookEvent +from litellm.proxy._types import ( + AlertType, + CallInfo, + Litellm_EntityType, + VirtualKeyEvent, + WebhookEvent, +) from litellm.types.integrations.slack_alerting import * from ..email_templates.templates import * @@ -570,7 +577,7 @@ class SlackAlerting(CustomBatchLogger): ttl=self.alerting_args.budget_alert_ttl, ) - async def budget_alerts( # noqa: PLR0915 + async def budget_alerts( self, type: Literal[ "token_budget", @@ -582,6 +589,13 @@ class SlackAlerting(CustomBatchLogger): ], user_info: CallInfo, ): + """ + Send a budget alert on slack or webhook + + Args: + type: The type of budget alert to send + user_info: The user info to send the alert for + """ ## PREVENTITIVE ALERTING ## - https://github.com/BerriAI/litellm/issues/2727 # - Alert once within 24hr period # - Cache this information @@ -593,9 +607,15 @@ class SlackAlerting(CustomBatchLogger): return if "budget_alerts" not in self.alert_types: return - _id: Optional[str] = "default_id" # used for caching + + # Get the appropriate budget alert type handler + budget_alert_class = get_budget_alert_type(type) + _id = budget_alert_class.get_id(user_info) user_info_json = user_info.model_dump(exclude_none=True) user_info_str = self._get_user_info_str(user_info) + event_message = budget_alert_class.get_event_message() + + # Set default event unless we're in projected_limit_exceeded event: Optional[ Literal[ "budget_crossed", @@ -603,69 +623,30 @@ class SlackAlerting(CustomBatchLogger): "projected_limit_exceeded", "soft_budget_crossed", ] - ] = None - event_group: Optional[ - Literal["internal_user", "team", "key", "proxy", "customer"] - ] = None - event_message: str = "" + ] = ( + "projected_limit_exceeded" if type == "projected_limit_exceeded" else None + ) + webhook_event: Optional[WebhookEvent] = None - if type == "proxy_budget": - event_group = "proxy" - event_message += "Proxy Budget: " - elif type == "soft_budget": - event_group = "proxy" - event_message += "Soft Budget Crossed: " - elif type == "user_budget": - event_group = "internal_user" - event_message += "User Budget: " - _id = user_info.user_id or _id - elif type == "team_budget": - event_group = "team" - event_message += "Team Budget: " - _id = user_info.team_id or _id - elif type == "token_budget": - event_group = "key" - event_message += "Key Budget: " - _id = user_info.token - elif type == "projected_limit_exceeded": - event_group = "key" - event_message += "Key Budget: Projected Limit Exceeded" - event = "projected_limit_exceeded" - _id = user_info.token # percent of max_budget left to spend if user_info.max_budget is None and user_info.soft_budget is None: return - percent_left: float = 0 - if user_info.max_budget is not None: - if user_info.max_budget > 0: - percent_left = ( - user_info.max_budget - user_info.spend - ) / user_info.max_budget # check if crossed budget - if user_info.max_budget is not None: - if user_info.spend >= user_info.max_budget: - event = "budget_crossed" - event_message += ( - f"Budget Crossed\n Total Budget:`{user_info.max_budget}`" - ) - elif percent_left <= SLACK_ALERTING_THRESHOLD_5_PERCENT: - event = "threshold_crossed" - event_message += "5% Threshold Crossed " - elif percent_left <= SLACK_ALERTING_THRESHOLD_15_PERCENT: - event = "threshold_crossed" - event_message += "15% Threshold Crossed" - elif user_info.soft_budget is not None: - if user_info.spend >= user_info.soft_budget: - event = "soft_budget_crossed" - if event is not None and event_group is not None: + event, event_message = self._get_event_and_event_message( + event=event, + user_info=user_info, + event_message=event_message, + ) + + # send alert + if event is not None and user_info.event_group is not None: _cache_key = "budget_alerts:{}:{}".format(event, _id) result = await _cache.async_get_cache(key=_cache_key) if result is None: webhook_event = WebhookEvent( event=event, - event_group=event_group, event_message=event_message, **user_info_json, ) @@ -685,6 +666,82 @@ class SlackAlerting(CustomBatchLogger): return return + def _get_event_and_event_message( + self, + user_info: CallInfo, + event: Optional[ + Literal[ + "budget_crossed", + "threshold_crossed", + "soft_budget_crossed", + "projected_limit_exceeded", + ] + ], + event_message: str, + ) -> Tuple[ + Optional[ + Literal[ + "budget_crossed", + "threshold_crossed", + "soft_budget_crossed", + "projected_limit_exceeded", + ] + ], + str, + ]: + """ + Get the event and event message for a budget alert + + This will append any new information to the event_message + + Handles Max Budget and Soft Budget Alerts + """ + percent_left: float = self._get_percent_of_max_budget_left(user_info=user_info) + + ##################################################################### + # SOFT BUDGET CHECK + # Check if the key/team/user has a soft budget set and they have crossed it + ##################################################################### + if user_info.soft_budget is not None: + if user_info.spend >= user_info.soft_budget: + event = "soft_budget_crossed" + event_message += f"Total Soft Budget:`{user_info.soft_budget}`" + + ##################################################################### + # MAX BUDGET CHECK + # Check if the key/team/user has a max budget set and they have either + ## a. Crossed their max budget + ## b. Either 5% or 15% of their max budget is left + ##################################################################### + if user_info.max_budget is not None: + if user_info.spend >= user_info.max_budget: + event = "budget_crossed" + event_message += ( + f"Budget Crossed\n Total Budget:`{user_info.max_budget}`" + ) + elif percent_left <= SLACK_ALERTING_THRESHOLD_5_PERCENT: + event = "threshold_crossed" + event_message += "5% Threshold Crossed " + elif percent_left <= SLACK_ALERTING_THRESHOLD_15_PERCENT: + event = "threshold_crossed" + event_message += "15% Threshold Crossed" + + return event, event_message + + def _get_percent_of_max_budget_left(self, user_info: CallInfo) -> float: + """ + Get the percent of the max budget that is left + """ + percent_left: float = 0.0 + current_spend: float = user_info.spend + max_budget: Optional[float] = user_info.max_budget + if max_budget is None: + return percent_left + if max_budget <= 0: + return percent_left + percent_left = (max_budget - current_spend) / max_budget + return percent_left + def _get_user_info_str(self, user_info: CallInfo) -> str: """ Create a standard message for a budget alert @@ -693,6 +750,8 @@ class SlackAlerting(CustomBatchLogger): _all_fields_as_dict.pop("token") msg = "" for k, v in _all_fields_as_dict.items(): + if isinstance(v, Litellm_EntityType): + v = v.value msg += f"*{k}:* `{v}`\n" return msg @@ -725,7 +784,7 @@ class SlackAlerting(CustomBatchLogger): projected_exceeded_date=None, projected_spend=None, event="spend_tracked", - event_group="customer", + event_group=Litellm_EntityType.END_USER, event_message="Customer spend tracked. Customer={}, spend={}".format( end_user_id, response_cost ), @@ -823,9 +882,9 @@ class SlackAlerting(CustomBatchLogger): ### UNIQUE CACHE KEY ### cache_key = provider + region_name - outage_value: Optional[ - ProviderRegionOutageModel - ] = await self.internal_usage_cache.async_get_cache(key=cache_key) + outage_value: Optional[ProviderRegionOutageModel] = ( + await self.internal_usage_cache.async_get_cache(key=cache_key) + ) if ( getattr(exception, "status_code", None) is None @@ -1406,9 +1465,9 @@ Model Info: self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url ): - slack_webhook_url: Optional[ - Union[str, List[str]] - ] = self.alert_to_webhook_url[alert_type] + slack_webhook_url: Optional[Union[str, List[str]]] = ( + self.alert_to_webhook_url[alert_type] + ) elif self.default_webhook_url is not None: slack_webhook_url = self.default_webhook_url else: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ed80c72dec4..beabe11434a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -161,6 +161,9 @@ class Litellm_EntityType(enum.Enum): TEAM_MEMBER = "team_member" ORGANIZATION = "organization" + # global proxy level entity + PROXY = "proxy" + def hash_token(token: str): import hashlib @@ -1823,6 +1826,7 @@ class CallInfo(LiteLLMPydanticObjectBase): key_alias: Optional[str] = None projected_exceeded_date: Optional[str] = None projected_spend: Optional[float] = None + event_group: Litellm_EntityType class WebhookEvent(CallInfo): @@ -1835,8 +1839,8 @@ class WebhookEvent(CallInfo): "internal_user_created", "spend_tracked", ] - event_group: Literal["internal_user", "key", "team", "proxy", "customer"] event_message: str # human-readable description of event + event_group: Litellm_EntityType class SpecialModelNames(enum.Enum): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index dac817e8116..05eaf8e5a2b 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -26,6 +26,7 @@ from litellm.proxy._types import ( RBAC_ROLES, CallInfo, LiteLLM_EndUserTable, + Litellm_EntityType, LiteLLM_JWTAuth, LiteLLM_OrganizationMembershipTable, LiteLLM_OrganizationTable, @@ -1338,6 +1339,7 @@ async def _virtual_key_max_budget_check( team_id=valid_token.team_id, user_email=user_email, key_alias=valid_token.key_alias, + event_group=Litellm_EntityType.KEY, ) asyncio.create_task( proxy_logging_obj.budget_alerts( @@ -1383,6 +1385,7 @@ async def _virtual_key_soft_budget_check( team_alias=valid_token.team_alias, user_email=None, key_alias=valid_token.key_alias, + event_group=Litellm_EntityType.KEY, ) asyncio.create_task( proxy_logging_obj.budget_alerts( @@ -1418,6 +1421,7 @@ async def _team_max_budget_check( user_id=valid_token.user_id, team_id=valid_token.team_id, team_alias=valid_token.team_alias, + event_group=Litellm_EntityType.TEAM, ) asyncio.create_task( proxy_logging_obj.budget_alerts( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 40921ac851b..60b7474bd6c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -192,6 +192,7 @@ async def get_global_proxy_spend( max_budget=litellm.max_budget, spend=global_proxy_spend, token=token, + event_group=Litellm_EntityType.PROXY, ) asyncio.create_task( proxy_logging_obj.budget_alerts( @@ -520,23 +521,23 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 proxy_logging_obj=proxy_logging_obj, ) if _end_user_object is not None: - end_user_params[ - "allowed_model_region" - ] = _end_user_object.allowed_model_region + end_user_params["allowed_model_region"] = ( + _end_user_object.allowed_model_region + ) if _end_user_object.litellm_budget_table is not None: budget_info = _end_user_object.litellm_budget_table if budget_info.tpm_limit is not None: - end_user_params[ - "end_user_tpm_limit" - ] = budget_info.tpm_limit + end_user_params["end_user_tpm_limit"] = ( + budget_info.tpm_limit + ) if budget_info.rpm_limit is not None: - end_user_params[ - "end_user_rpm_limit" - ] = budget_info.rpm_limit + end_user_params["end_user_rpm_limit"] = ( + budget_info.rpm_limit + ) if budget_info.max_budget is not None: - end_user_params[ - "end_user_max_budget" - ] = budget_info.max_budget + end_user_params["end_user_max_budget"] = ( + budget_info.max_budget + ) except Exception as e: if isinstance(e, litellm.BudgetExceededError): raise e @@ -952,6 +953,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 max_budget=litellm.max_budget, user_id=litellm_proxy_admin_name, team_id=valid_token.team_id, + event_group=Litellm_EntityType.PROXY, ) asyncio.create_task( proxy_logging_obj.budget_alerts( diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 9de845397aa..fc3a2e0650f 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -14,6 +14,7 @@ from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS from litellm.proxy._types import ( AlertType, CallInfo, + Litellm_EntityType, ProxyErrorTypes, ProxyException, UserAPIKeyAuth, @@ -168,6 +169,7 @@ async def health_services_endpoint( # noqa: PLR0915 user_id=user_api_key_dict.user_id, key_alias=user_api_key_dict.key_alias, team_id=user_api_key_dict.team_id, + event_group=Litellm_EntityType.KEY, ) await proxy_logging_obj.budget_alerts( type="user_budget", @@ -251,7 +253,7 @@ async def health_services_endpoint( # noqa: PLR0915 if service == "email": webhook_event = WebhookEvent( event="key_created", - event_group="key", + event_group=Litellm_EntityType.KEY, event_message="Test Email Alert", token=user_api_key_dict.token or "", key_alias="Email Test key (This is only a test alert key. DO NOT USE THIS IN PRODUCTION.)", diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index c2c4f0669f3..cd5b1f353b7 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -13,6 +13,7 @@ from litellm.proxy._types import ( GenerateKeyResponse, KeyRequest, LiteLLM_AuditLogs, + Litellm_EntityType, LiteLLM_VerificationToken, LitellmTableNames, ProxyErrorTypes, @@ -305,7 +306,7 @@ class KeyManagementEventHooks: ) event = WebhookEvent( event="key_created", - event_group="key", + event_group=Litellm_EntityType.KEY, event_message="API Key Created", token=response.get("token", ""), spend=response.get("spend", 0.0), diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index e6ed4f5f103..f579fae25e4 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -94,9 +94,9 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d data_json["user_id"] = str(uuid.uuid4()) auto_create_key = data_json.pop("auto_create_key", True) if auto_create_key is False: - data_json[ - "table_name" - ] = "user" # only create a user, don't create key if 'auto_create_key' set to False + data_json["table_name"] = ( + "user" # only create a user, don't create key if 'auto_create_key' set to False + ) is_internal_user = False if data.user_role and data.user_role.is_internal_user_role: @@ -292,7 +292,7 @@ async def new_user( event = WebhookEvent( event="internal_user_created", - event_group="internal_user", + event_group=Litellm_EntityType.USER, event_message="Welcome to LiteLLM Proxy", token=response.get("token", ""), spend=response.get("spend", 0.0), @@ -651,6 +651,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di if "budget_duration" in non_default_values: from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + non_default_values["budget_reset_at"] = get_budget_reset_time( budget_duration=non_default_values["budget_duration"] ) @@ -665,10 +666,11 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di "budget_duration" not in non_default_values ): # applies internal user limits, if user role updated if is_internal_user and litellm.internal_user_budget_duration is not None: - non_default_values[ - "budget_duration" - ] = litellm.internal_user_budget_duration + non_default_values["budget_duration"] = ( + litellm.internal_user_budget_duration + ) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + non_default_values["budget_reset_at"] = get_budget_reset_time( budget_duration=non_default_values["budget_duration"] ) @@ -1054,9 +1056,9 @@ async def get_users( where=where_conditions, skip=skip, take=page_size, - order=order_by - if order_by - else {"created_at": "desc"}, # Default to created_at desc if no sort specified + order=( + order_by if order_by else {"created_at": "desc"} + ), # Default to created_at desc if no sort specified ) # Get total count of user rows @@ -1309,13 +1311,13 @@ async def ui_view_users( } # Query users with pagination and filters - users: Optional[ - List[BaseModel] - ] = await prisma_client.db.litellm_usertable.find_many( - where=where_conditions, - skip=skip, - take=page_size, - order={"created_at": "desc"}, + users: Optional[List[BaseModel]] = ( + await prisma_client.db.litellm_usertable.find_many( + where=where_conditions, + skip=skip, + take=page_size, + order={"created_at": "desc"}, + ) ) if not users: diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 7fa6492f6df..5d5ce4b77ae 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -19,4 +19,6 @@ vector_stores: source: "https://www.litellm.com/docs" +general_settings: + alerting: ["webhook"] diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c056147addd..c289b9d188c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1015,6 +1015,7 @@ async def update_cache( # noqa: PLR0915 user_id=existing_spend_obj.user_id, projected_spend=projected_spend, projected_exceeded_date=projected_exceeded_date, + event_group=Litellm_EntityType.KEY, ) # alert user asyncio.create_task( diff --git a/tests/litellm/integrations/SlackAlerting/test_slack_alerting.py b/tests/litellm/integrations/SlackAlerting/test_slack_alerting.py new file mode 100644 index 00000000000..425bf8b6a58 --- /dev/null +++ b/tests/litellm/integrations/SlackAlerting/test_slack_alerting.py @@ -0,0 +1,163 @@ +import datetime +import json +import os +import sys +import unittest +from typing import List, Optional, Tuple +from unittest.mock import ANY, MagicMock, Mock, patch + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system-path +import litellm +from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting +from litellm.proxy._types import CallInfo, Litellm_EntityType + + +class TestSlackAlerting(unittest.TestCase): + def setUp(self): + self.slack_alerting = SlackAlerting() + + def test_get_percent_of_max_budget_left(self): + # Test case 1: When max_budget is None + user_info = CallInfo( + max_budget=None, spend=50.0, event_group=Litellm_EntityType.KEY + ) + result = self.slack_alerting._get_percent_of_max_budget_left(user_info) + self.assertEqual(result, 0.0) + + # Test case 2: When max_budget is 0 + user_info = CallInfo( + max_budget=0.0, spend=50.0, event_group=Litellm_EntityType.KEY + ) + result = self.slack_alerting._get_percent_of_max_budget_left(user_info) + self.assertEqual(result, 0.0) + + # Test case 3: When spend is less than max_budget + user_info = CallInfo( + max_budget=100.0, spend=75.0, event_group=Litellm_EntityType.KEY + ) + result = self.slack_alerting._get_percent_of_max_budget_left(user_info) + self.assertEqual(result, 0.25) + + # Test case 4: When spend equals max_budget + user_info = CallInfo( + max_budget=100.0, spend=100.0, event_group=Litellm_EntityType.KEY + ) + result = self.slack_alerting._get_percent_of_max_budget_left(user_info) + self.assertEqual(result, 0.0) + + # Test case 5: When spend exceeds max_budget + user_info = CallInfo( + max_budget=100.0, spend=120.0, event_group=Litellm_EntityType.KEY + ) + result = self.slack_alerting._get_percent_of_max_budget_left(user_info) + self.assertEqual(result, -0.2) + + def test_get_event_and_event_message_max_budget(self): + # Initial setup with no event + event = None + event_message = "Test Message: " + + # Test case 1: When spend exceeds max_budget + user_info = CallInfo( + max_budget=100.0, + spend=120.0, + soft_budget=None, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + self.assertEqual(event, "budget_crossed") + self.assertTrue("Budget Crossed" in event_message) + + # Test case 2: When 5% of max_budget is left + user_info = CallInfo( + max_budget=100.0, + spend=95.0, + soft_budget=None, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + self.assertEqual(event, "threshold_crossed") + self.assertTrue("5% Threshold Crossed" in event_message) + + # Test case 3: When 15% of max_budget is left + user_info = CallInfo( + max_budget=100.0, + spend=85.0, + soft_budget=None, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + self.assertEqual(event, "threshold_crossed") + self.assertTrue("15% Threshold Crossed" in event_message) + + def test_get_event_and_event_message_soft_budget(self): + # Initial setup with no event + event = None + event_message = "Test Message: " + + # Test case 1: When spend exceeds soft_budget + user_info = CallInfo( + max_budget=None, + spend=120.0, + soft_budget=100.0, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + self.assertEqual(event, "soft_budget_crossed") + self.assertTrue("Total Soft Budget" in event_message) + + # Test case 2: When spend is less than soft_budget + user_info = CallInfo( + max_budget=None, + spend=90.0, + soft_budget=100.0, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=None, event_message=event_message + ) + print("got event", event) + print("got event_message", event_message) + self.assertEqual(event, None) # No event should be triggered + + def test_get_event_and_event_message_both_budgets(self): + # Initial setup with no event + event = None + event_message = "Test Message: " + + # Test case 1: When spend exceeds both max_budget and soft_budget + user_info = CallInfo( + max_budget=150.0, + spend=160.0, + soft_budget=100.0, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + # budget_crossed has higher priority + self.assertEqual(event, "budget_crossed") + self.assertTrue("Budget Crossed" in event_message) + + # Test case 2: When spend exceeds soft_budget but not max_budget + user_info = CallInfo( + max_budget=150.0, + spend=120.0, + soft_budget=100.0, + event_group=Litellm_EntityType.KEY, + ) + event, event_message = self.slack_alerting._get_event_and_event_message( + user_info=user_info, event=event, event_message=event_message + ) + self.assertEqual(event, "soft_budget_crossed") + self.assertTrue("Total Soft Budget" in event_message) diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index 26a5e0822fe..ef6b47154bf 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -33,7 +33,7 @@ from litellm.integrations.SlackAlerting.slack_alerting import ( DeploymentMetrics, SlackAlerting, ) -from litellm.proxy._types import CallInfo +from litellm.proxy._types import CallInfo, Litellm_EntityType, WebhookEvent from litellm.proxy.utils import ProxyLogging from litellm.router import AlertingConfig, Router from litellm.utils import get_api_base @@ -204,7 +204,10 @@ async def test_budget_alerts_crossed(slack_alerting): await slack_alerting.budget_alerts( "user_budget", user_info=CallInfo( - token="", spend=user_current_spend, max_budget=user_max_budget + token="", + spend=user_current_spend, + max_budget=user_max_budget, + event_group=Litellm_EntityType.USER, ), ) mock_send_alert.assert_awaited_once() @@ -219,7 +222,10 @@ async def test_budget_alerts_crossed_again(slack_alerting): await slack_alerting.budget_alerts( "user_budget", user_info=CallInfo( - token="", spend=user_current_spend, max_budget=user_max_budget + token="", + spend=user_current_spend, + max_budget=user_max_budget, + event_group=Litellm_EntityType.USER, ), ) mock_send_alert.assert_awaited_once() @@ -227,7 +233,10 @@ async def test_budget_alerts_crossed_again(slack_alerting): await slack_alerting.budget_alerts( "user_budget", user_info=CallInfo( - token="", spend=user_current_spend, max_budget=user_max_budget + token="", + spend=user_current_spend, + max_budget=user_max_budget, + event_group=Litellm_EntityType.USER, ), ) mock_send_alert.assert_not_awaited() @@ -502,6 +511,7 @@ async def test_send_token_budget_crossed_alerts(alerting_type): "key_alias": "my-test-key", "projected_exceeded_date": "10/20/2024", "projected_spend": 200, + "event_group": Litellm_EntityType.KEY, } user_info = CallInfo(**user_info) @@ -540,6 +550,7 @@ async def test_webhook_alerting(alerting_type): "key_alias": "my-test-key", "projected_exceeded_date": "10/20/2024", "projected_spend": 200, + "event_group": Litellm_EntityType.KEY, } user_info = CallInfo(**user_info) @@ -961,6 +972,7 @@ async def test_spend_report_cache(report_type): user_id="test@test.com", user_email="test@test.com", key_alias="test-key", + event_group=Litellm_EntityType.KEY, ) with patch.object( @@ -1000,6 +1012,7 @@ async def test_soft_budget_alerts(): user_id="test@test.com", user_email="test@test.com", key_alias="test-key", + event_group=Litellm_EntityType.KEY, ) await slack_alerting.budget_alerts( @@ -1010,14 +1023,109 @@ async def test_soft_budget_alerts(): # Verify alert message contains correct percentage alert_message = mock_send_alert.call_args[1]["message"] - print(alert_message) + + print("GOT MESSAGE\n\n", alert_message) expected_message = ( - "Soft Budget Crossed: \n\n" + "Soft Budget Crossed: Total Soft Budget:`80.0`\n" + "\n" "*spend:* `80.0`\n" "*soft_budget:* `80.0`\n" "*user_id:* `test@test.com`\n" "*user_email:* `test@test.com`\n" "*key_alias:* `test-key`\n" + "*event_group:* `key`\n" ) assert alert_message == expected_message + + +key_info = CallInfo( + token="test_token", + spend=81, + soft_budget=80, + max_budget=100, + user_id="test@test.com", + user_email="test@test.com", + key_alias="test-key", + event_group=Litellm_EntityType.KEY, +) + +team_info = CallInfo( + token="test_token", + spend=160, + soft_budget=150, + max_budget=200, + team_id="team-123", + team_alias="engineering-team", + event_group=Litellm_EntityType.TEAM, +) + +user_info = CallInfo( + token="test_token", + spend=45, + soft_budget=40, + max_budget=50, + user_id="user123", + event_group=Litellm_EntityType.USER, +) + +key_no_max_budget_info = CallInfo( + token="test_token", + spend=90, + soft_budget=85, + user_id="dev@test.com", + user_email="dev@test.com", + key_alias="dev-key", + event_group=Litellm_EntityType.KEY, +) + + +@pytest.mark.parametrize( + "entity_info", + [ + key_info, + team_info, + user_info, + key_no_max_budget_info, + ], +) +@pytest.mark.asyncio +async def test_soft_budget_alerts_webhook(entity_info): + """ + Tests that soft budget alerts are triggered for different entity types. + + Tests: + - Key with max budget + - Team + - User + - Key without max budget + """ + slack_alerting = SlackAlerting(alerting=["webhook"]) + + with patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert: + # Test entity hit soft budget limit + await slack_alerting.budget_alerts( + type="soft_budget", + user_info=entity_info, + ) + mock_send_alert.assert_called_once() + + # Verify the webhook event + call_args = mock_send_alert.call_args[1] + logged_webhook_event: WebhookEvent = call_args["user_info"] + + # Validate the webhook event has all expected fields + assert logged_webhook_event.spend == entity_info.spend + assert logged_webhook_event.soft_budget == entity_info.soft_budget + assert logged_webhook_event.max_budget == entity_info.max_budget + assert logged_webhook_event.user_id == entity_info.user_id + assert logged_webhook_event.user_email == entity_info.user_email + assert logged_webhook_event.key_alias == entity_info.key_alias + assert logged_webhook_event.event_group == entity_info.event_group + + + + + + + \ No newline at end of file