mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
commit
50cd26cd9c
4 changed files with 174 additions and 6 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue