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:
ryan-crabbe-berri 2026-09-24 14:03:48 -07:00
parent 98eb84d7d3
commit b9335b6839
5 changed files with 48 additions and 26 deletions

View file

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

View file

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

View file

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

View file

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

View file

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