Merge pull request #40396 from jon-walton/litellm_user_budget_webhook_alerts

fix(proxy): emit internal user budget webhook alerts
This commit is contained in:
ryan-crabbe-berri 2026-09-11 18:03:35 -07:00 committed by GitHub
commit 50cd26cd9c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 174 additions and 6 deletions

View file

@ -673,7 +673,7 @@ class SlackAlerting(CustomBatchLogger):
Create a standard message for a budget alert
"""
_all_fields_as_dict: Final[dict[str, object]] = user_info.model_dump(exclude_none=True)
_all_fields_as_dict.pop("token")
_all_fields_as_dict.pop("token", None)
msg = ""
for k, v in _all_fields_as_dict.items():
if isinstance(v, Litellm_EntityType):

View file

@ -1022,6 +1022,19 @@ async def common_checks(
fallback_spend=user_object.spend or 0.0,
max_budget=user_budget,
)
call_info: Final = CallInfo(
spend=user_spend,
max_budget=user_budget,
user_id=user_object.user_id,
user_email=user_object.user_email,
event_group=Litellm_EntityType.USER,
)
asyncio.create_task(
proxy_logging_obj.budget_alerts(
type="user_budget",
user_info=call_info,
)
)
if math.isfinite(user_budget) and user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_spend,

View file

@ -45,6 +45,21 @@ class TestSlackAlerting(unittest.TestCase):
result = self.slack_alerting._get_percent_of_max_budget_left(user_info)
self.assertEqual(result, -0.2)
def test_get_user_info_str_omits_absent_token_for_user_alert(self):
user_info = CallInfo(
spend=85.0,
max_budget=100.0,
user_id="user-1",
user_email="person@example.com",
event_group=Litellm_EntityType.USER,
)
result = self.slack_alerting._get_user_info_str(user_info)
self.assertIn("*user_id:* `user-1`", result)
self.assertIn("*user_email:* `person@example.com`", result)
self.assertNotIn("*token:*", result)
def test_get_event_and_event_message_max_budget(self):
# Initial setup with no event
event = None

View file

@ -1,7 +1,7 @@
import asyncio
import json
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Optional
from typing import TYPE_CHECKING, Final, Literal, Optional
from unittest.mock import AsyncMock, MagicMock, patch
if TYPE_CHECKING:
@ -29,6 +29,7 @@ from litellm.proxy._types import (
ProxyException,
SSOUserDefinedValues,
UserAPIKeyAuth,
WebhookEvent,
)
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
@ -53,6 +54,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.constants import (
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
@ -5431,6 +5433,9 @@ async def test_common_checks_personal_user_budget_blocks_in_gather():
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
return 999.0 if counter_key == "spend:user:u1" else 0.0
proxy_logging_obj: Final = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_by_counter),
@ -5445,13 +5450,136 @@ async def test_common_checks_personal_user_budget_blocks_in_gather():
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
proxy_logging_obj=proxy_logging_obj,
valid_token=token,
request=MagicMock(spec=Request),
)
await asyncio.sleep(0)
assert "User=u1" in str(over.value)
async def _run_internal_user_budget_alert(
*,
spend: float,
) -> tuple[AsyncMock, litellm.BudgetExceededError | None]:
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
user: Final = LiteLLM_UserTable(
user_id="user-1",
user_email="person@example.com",
spend=0.0,
max_budget=100.0,
)
token: Final = UserAPIKeyAuth(token="hashed-key-1", user_id="user-1")
slack_alerting: Final = SlackAlerting(alerting=["webhook"])
send_alert: Final = AsyncMock()
alert_finished: Final = asyncio.Event()
async def _get_spend(
counter_key: str,
fallback_spend: float,
max_budget: float | None = None,
**kwargs: object,
) -> float:
assert counter_key == "spend:user:user-1"
assert fallback_spend == 0.0
assert max_budget == 100.0
return spend
async def _budget_alerts(
*,
type: Literal["user_budget"],
user_info: CallInfo,
) -> None:
assert type == "user_budget"
try:
await slack_alerting.budget_alerts(type=type, user_info=user_info)
finally:
alert_finished.set()
proxy_logging_obj: Final = MagicMock(budget_alerts=_budget_alerts)
async def _check() -> bool:
return await common_checks(
request_body={"messages": [{"role": "user", "content": "hi"}]},
team_object=None,
user_object=user,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=proxy_logging_obj,
valid_token=token,
request=MagicMock(spec=Request),
)
async def _check_for_error() -> litellm.BudgetExceededError | None:
if spend < 100.0:
assert await _check() is True
return None
with pytest.raises(litellm.BudgetExceededError) as raised:
await _check()
return raised.value
with (
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam
patch("litellm.proxy.proxy_server.get_current_spend", _get_spend), # test-quality-ok: common_checks imports it locally
patch.object(slack_alerting, "send_alert", send_alert),
):
error: Final = await _check_for_error()
await asyncio.wait_for(alert_finished.wait(), timeout=1.0)
return send_alert, error
@pytest.mark.asyncio
async def test_common_checks_internal_user_budget_below_threshold_does_not_emit_alert():
send_alert, error = await _run_internal_user_budget_alert(spend=84.0)
assert error is None
send_alert.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_checks_internal_user_budget_emits_user_threshold_event():
send_alert, error = await _run_internal_user_budget_alert(spend=85.0)
assert error is None
send_alert.assert_awaited_once()
event: Final = send_alert.await_args.kwargs["user_info"]
assert isinstance(event, WebhookEvent)
assert event.event == "threshold_crossed"
assert event.event_group == Litellm_EntityType.USER
assert event.user_id == "user-1"
assert event.user_email == "person@example.com"
assert event.spend == 85.0
assert event.max_budget == 100.0
assert event.token is None
assert event.key_alias is None
assert event.team_id is None
assert event.organization_id is None
@pytest.mark.asyncio
async def test_common_checks_internal_user_budget_emits_crossed_event_and_rejects():
send_alert, error = await _run_internal_user_budget_alert(spend=100.0)
assert error is not None
assert error.current_cost == 100.0
assert error.max_budget == 100.0
send_alert.assert_awaited_once()
event: Final = send_alert.await_args.kwargs["user_info"]
assert isinstance(event, WebhookEvent)
assert event.event == "budget_crossed"
assert event.event_group == Litellm_EntityType.USER
assert event.user_id == "user-1"
assert event.user_email == "person@example.com"
@pytest.mark.asyncio
async def test_common_checks_personal_user_budget_skipped_for_team_key():
"""A user's personal max_budget does not apply to a team-scoped key.
@ -5475,6 +5603,9 @@ async def test_common_checks_personal_user_budget_skipped_for_team_key():
async def _no_membership(*args, **kwargs):
return None
proxy_logging_obj: Final = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_by_counter),
@ -5489,11 +5620,12 @@ async def test_common_checks_personal_user_budget_skipped_for_team_key():
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
proxy_logging_obj=proxy_logging_obj,
valid_token=token,
request=MagicMock(spec=Request),
)
assert result is True
proxy_logging_obj.budget_alerts.assert_not_awaited()
@pytest.mark.asyncio
@ -5518,6 +5650,9 @@ async def test_common_checks_personal_user_budget_enforced_on_team_key_when_flag
async def _no_membership(*args, **kwargs):
return None
proxy_logging_obj: Final = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_by_counter),
@ -5533,10 +5668,11 @@ async def test_common_checks_personal_user_budget_enforced_on_team_key_when_flag
general_settings={"apply_user_budget_to_team_keys": True},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
proxy_logging_obj=proxy_logging_obj,
valid_token=token,
request=MagicMock(spec=Request),
)
await asyncio.sleep(0)
assert "ExceededBudget: User=u1" in str(exc_info.value)
@ -5553,6 +5689,9 @@ async def test_common_checks_personal_user_budget_still_enforced_on_personal_key
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
return 999.0 if counter_key == "spend:user:u1" else 0.0
proxy_logging_obj: Final = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_by_counter),
@ -5567,10 +5706,11 @@ async def test_common_checks_personal_user_budget_still_enforced_on_personal_key
general_settings={"apply_user_budget_to_team_keys": True},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
proxy_logging_obj=proxy_logging_obj,
valid_token=token,
request=MagicMock(spec=Request),
)
await asyncio.sleep(0)
@pytest.mark.parametrize(