[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
This commit is contained in:
Ishaan Jaff 2025-05-02 07:06:07 -07:00 • committed by GitHub
parent cb177dbd7a
commit 96e75628d6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 574 additions and 100 deletions

View file

@ -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)

View file

@ -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()

View file

@ -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:

View file

@ -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):

View file

@ -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(

View file

@ -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(

View file

@ -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.)",

View file

@ -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),

View file

@ -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:

View file

@ -19,4 +19,6 @@ vector_stores:
source: "https://www.litellm.com/docs"
general_settings:
alerting: ["webhook"]

View file

@ -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(

View file

@ -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)

View file

@ -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