mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(alerting): keep BaseBudgetAlertType.get_event_message zero-arg
Requiring user_info broke existing callers and out-of-tree subclasses. The team member label now comes from SlackAlerting.budget_alerts, so the interface and its Readme are unchanged from main.
This commit is contained in:
parent
98eb84d7d3
commit
b9335b6839
5 changed files with 48 additions and 26 deletions
|
|
@ -15,7 +15,7 @@ The `budget_alert_types.py` module provides a flexible framework for handling di
|
|||
|
||||
- `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(user_info)`: Returns the message prefix 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:
|
||||
|
|
@ -36,7 +36,7 @@ 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(user_info) # Returns "User Budget: "
|
||||
event_message = budget_alert_class.get_event_message() # Returns "User Budget: "
|
||||
cache_id = budget_alert_class.get_id(user_info) # Returns user_id
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ class BaseBudgetAlertType(ABC):
|
|||
"""Base class for different budget alert types"""
|
||||
|
||||
@abstractmethod
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
"""Return the event message for this alert type"""
|
||||
|
||||
@abstractmethod
|
||||
|
|
@ -17,7 +17,7 @@ class BaseBudgetAlertType(ABC):
|
|||
|
||||
|
||||
class ProxyBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Proxy Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -25,7 +25,7 @@ class ProxyBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class SoftBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Soft Budget Crossed: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -35,7 +35,7 @@ class SoftBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class UserBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "User Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -43,7 +43,7 @@ class UserBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class TeamBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Team Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -51,7 +51,7 @@ class TeamBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class OrganizationBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Organization Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -59,9 +59,7 @@ class OrganizationBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class TokenBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
return "Team Member Budget: "
|
||||
def get_event_message(self) -> str:
|
||||
return "Key Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -71,7 +69,7 @@ class TokenBudgetAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class ProjectedLimitExceededAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Key Budget: Projected Limit Exceeded"
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
@ -79,7 +77,7 @@ class ProjectedLimitExceededAlert(BaseBudgetAlertType):
|
|||
|
||||
|
||||
class ProjectBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self, user_info: CallInfo) -> str:
|
||||
def get_event_message(self) -> str:
|
||||
return "Project Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
|
|
|
|||
|
|
@ -551,7 +551,12 @@ class SlackAlerting(CustomBatchLogger):
|
|||
budget_alert_class: Final = get_budget_alert_type(type)
|
||||
_id: Final = budget_alert_class.get_id(user_info)
|
||||
user_info_str: Final = self._get_user_info_str(user_info)
|
||||
event_message = budget_alert_class.get_event_message(user_info)
|
||||
# Team member max-budget alerts ride the key-budget alert type; label them by what they measure.
|
||||
event_message = (
|
||||
"Team Member Budget: "
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER
|
||||
else budget_alert_class.get_event_message()
|
||||
)
|
||||
|
||||
# Set default event unless we're in projected_limit_exceeded
|
||||
event: (
|
||||
|
|
|
|||
|
|
@ -94,12 +94,4 @@ class TestTokenBudgetAlert:
|
|||
)
|
||||
|
||||
assert alert.get_id(user_info) == "hashed_key"
|
||||
assert alert.get_event_message(user_info) == "Key Budget: "
|
||||
|
||||
def test_get_event_message_labels_team_member_alerts(self):
|
||||
alert = TokenBudgetAlert()
|
||||
user_info = CallInfo(
|
||||
spend=8.0, max_budget=10.0, user_id="member_1", team_id="team_a", event_group=Litellm_EntityType.TEAM_MEMBER
|
||||
)
|
||||
|
||||
assert alert.get_event_message(user_info) == "Team Member Budget: "
|
||||
assert alert.get_event_message() == "Key Budget: "
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ class TestSlackAlerting(unittest.TestCase):
|
|||
|
||||
def test_get_event_and_event_message_max_budget(self):
|
||||
event = None
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message()
|
||||
|
||||
# Test case 1: When spend exceeds max_budget
|
||||
user_info = CallInfo(
|
||||
|
|
@ -75,33 +76,32 @@ class TestSlackAlerting(unittest.TestCase):
|
|||
soft_budget=None,
|
||||
event_group=Litellm_EntityType.KEY,
|
||||
)
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message(user_info)
|
||||
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)
|
||||
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message()
|
||||
user_info = CallInfo(
|
||||
max_budget=100.0,
|
||||
spend=95.0,
|
||||
soft_budget=None,
|
||||
event_group=Litellm_EntityType.KEY,
|
||||
)
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message(user_info)
|
||||
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.assertEqual(event_message, "User Budget: 5% or less of budget remaining")
|
||||
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message()
|
||||
user_info = CallInfo(
|
||||
max_budget=100.0,
|
||||
spend=85.0,
|
||||
soft_budget=None,
|
||||
event_group=Litellm_EntityType.KEY,
|
||||
)
|
||||
event_message = get_budget_alert_type("user_budget").get_event_message(user_info)
|
||||
event, event_message = self.slack_alerting._get_event_and_event_message(
|
||||
user_info=user_info, event=event, event_message=event_message
|
||||
)
|
||||
|
|
@ -393,6 +393,33 @@ def _slack_alerting_with_env_resolution() -> SlackAlerting:
|
|||
return slack_alerting
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"event_group, expected_prefix",
|
||||
[
|
||||
(Litellm_EntityType.TEAM_MEMBER, "Team Member Budget: Budget Crossed"),
|
||||
(Litellm_EntityType.KEY, "Key Budget: Budget Crossed"),
|
||||
],
|
||||
)
|
||||
async def test_max_budget_alert_labels_team_member_budget(event_group, expected_prefix):
|
||||
slack_alerting: Final = _slack_alerting_with_env_resolution()
|
||||
slack_alerting.send_alert = AsyncMock()
|
||||
|
||||
await slack_alerting.budget_alerts(
|
||||
type="max_budget_alert",
|
||||
user_info=CallInfo(
|
||||
spend=10.5,
|
||||
max_budget=10.0,
|
||||
token="hashed_key",
|
||||
user_id="member_1",
|
||||
team_id="team_a",
|
||||
event_group=event_group,
|
||||
),
|
||||
)
|
||||
|
||||
assert slack_alerting.send_alert.await_args.kwargs["message"].startswith(expected_prefix)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_alert_falls_back_to_alerting_webhook_url_env(monkeypatch):
|
||||
monkeypatch.delenv("SLACK_WEBHOOK_URL", raising=False)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue