From be50d86c749bc6c8c91f0749270199744f302f70 Mon Sep 17 00:00:00 2001 From: XAVIER ALMENDROS Date: Sun, 4 Oct 2026 19:17:41 +0200 Subject: [PATCH] feat(alerting): configure common budget notification percentages --- .../send_emails/base_email.py | 35 ++++- .../SlackAlerting/budget_alert_types.py | 8 ++ .../SlackAlerting/slack_alerting.py | 28 +++- litellm/proxy/auth/auth_checks.py | 10 +- litellm/proxy/auth/user_api_key_auth.py | 6 + litellm/proxy/utils.py | 20 ++- litellm/types/integrations/slack_alerting.py | 18 ++- .../send_emails/test_base_email.py | 102 ++++++++++++++- .../SlackAlerting/test_slack_alerting.py | 122 ++++++++++++++++++ ...st_auth_checks_object_access_and_lookup.py | 54 +++++++- .../utils/proxy_logging/test_alerting.py | 65 ++++++++-- 11 files changed, 429 insertions(+), 39 deletions(-) diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index 6e33d9f1bf3..0a4a21c25f1 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -6,7 +6,7 @@ Base class for sending emails to user after creating keys or invite links import html import json import os -from typing import List, Literal, Optional +from typing import Final, List, Literal, Optional from litellm_enterprise.types.enterprise_callbacks.send_emails import ( EmailEvent, @@ -38,6 +38,7 @@ from litellm.integrations.email_templates.templates import ( from litellm.integrations.email_templates.user_invitation_email import ( USER_INVITATION_EMAIL_TEMPLATE, ) +from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_threshold from litellm.proxy._types import ( CallInfo, InvitationNew, @@ -450,6 +451,7 @@ class BaseEmailLogger(CustomLogger): "projected_limit_exceeded", ], user_info: CallInfo, + budget_alert_thresholds: tuple[int, ...] | None = None, ): """ Send a budget alert via email @@ -551,8 +553,17 @@ class BaseEmailLogger(CustomLogger): ) return - alert_threshold = ( - user_info.max_budget * EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE + configured_threshold: Final = ( + get_budget_alert_threshold(user_info.spend, user_info.max_budget, budget_alert_thresholds) + if budget_alert_thresholds is not None + else None + ) + if budget_alert_thresholds is not None and configured_threshold is None: + return + alert_threshold = user_info.max_budget * ( + configured_threshold / 100 + if configured_threshold is not None + else EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE ) # Only alert if we've crossed the threshold but haven't exceeded max_budget yet @@ -562,7 +573,12 @@ class BaseEmailLogger(CustomLogger): ): # Generate cache key based on event type and identifier _id = user_info.token or user_info.user_id or "default_id" - _cache_key = f"email_budget_alerts:max_budget_alert:{_id}" + cache_id: Final = f"{_id}:{configured_threshold}" if configured_threshold is not None else _id + _cache_key = f"email_budget_alerts:max_budget_alert:{cache_id}" + if configured_threshold is not None and await _cache.async_get_cache( + key=f"email_budget_alerts:max_budget_alert:{_id}" + ) is not None: + return send_count = await _cache.async_increment_cache( key=_cache_key, @@ -571,8 +587,10 @@ class BaseEmailLogger(CustomLogger): ) if send_count is None or send_count <= 1: # Calculate percentage - percentage = int( - EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100 + percentage = ( + configured_threshold + if configured_threshold is not None + else int(EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100) ) # Create WebhookEvent for max budget alert @@ -597,7 +615,10 @@ class BaseEmailLogger(CustomLogger): ) try: - await self.send_max_budget_alert_email(webhook_event) + if configured_threshold is None: + await self.send_max_budget_alert_email(webhook_event) + else: + await self.send_max_budget_alert_email(webhook_event, threshold_pct=configured_threshold) except Exception as e: verbose_proxy_logger.error( f"Error sending max budget alert email: {e}", diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index 4fe833acecc..32645eadd68 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -1,9 +1,17 @@ +import math from abc import ABC, abstractmethod +from collections.abc import Sequence from typing import Final, Literal from litellm.proxy._types import CallInfo, Litellm_EntityType +def get_budget_alert_threshold(spend: float, max_budget: float | None, thresholds: Sequence[int]) -> int | None: + if max_budget is None or not math.isfinite(max_budget) or max_budget <= 0 or not math.isfinite(spend): + return None + return max((threshold for threshold in thresholds if spend >= max_budget * (threshold / 100)), default=None) + + class BaseBudgetAlertType(ABC): """Base class for different budget alert types""" diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 44e22f3c0fc..f29cf705acd 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -25,7 +25,7 @@ from litellm.constants import ( SLACK_MODEL_DEPRECATION_LOCK_ID, ) from litellm.integrations.custom_batch_logger import CustomBatchLogger -from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type +from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_threshold, get_budget_alert_type from litellm.integrations.SlackAlerting.hanging_request_check import ( AlertingHangingRequestCheck, ) @@ -536,6 +536,7 @@ class SlackAlerting(CustomBatchLogger): "project_budget", ], user_info: CallInfo, + send_threshold_email: bool = True, ): """ Send a budget alert on slack or webhook @@ -586,9 +587,19 @@ class SlackAlerting(CustomBatchLogger): # send alert if event is not None and user_info.event_group is not None: - _cache_key: Final = f"budget_alerts:{event}:{_id}" + threshold: Final = ( + get_budget_alert_threshold( + user_info.spend, user_info.max_budget, self.alerting_args.budget_alert_thresholds + ) + if event == "threshold_crossed" and self.alerting_args.budget_alert_thresholds is not None + else None + ) + cache_id: Final = f"{_id}:{threshold}" if threshold is not None else _id + _cache_key: Final = f"budget_alerts:{event}:{cache_id}" + if threshold is not None and await _cache.async_get_cache(key=f"budget_alerts:{event}:{_id}") == "SENT": + return result: Final = await _cache.async_get_cache(key=_cache_key) - slack_cache_key: Final = f"budget_alerts:slack:{event}:{_id}" + slack_cache_key: Final = f"budget_alerts:slack:{event}:{cache_id}" slack_due: Final[bool] = ( "slack" in self.alerting and self._slack_budget_alert_allowed(user_info) @@ -623,6 +634,7 @@ class SlackAlerting(CustomBatchLogger): user_info=webhook_event, alerting_metadata={}, budget_alert_destination=("slack" if result is not None else "all") if slack_due else "non_slack", + budget_alert_email=send_threshold_email or event != "threshold_crossed", ) if slack_accepted: await _cache.async_set_cache( @@ -678,6 +690,14 @@ class SlackAlerting(CustomBatchLogger): if user_info.spend >= user_info.max_budget: event = "budget_crossed" event_message += f"Budget Crossed\n Total Budget:`{user_info.max_budget}`" + elif self.alerting_args.budget_alert_thresholds is not None: + if event in ("soft_budget_crossed", "projected_limit_exceeded"): + return event, event_message + threshold: Final = get_budget_alert_threshold( + user_info.spend, user_info.max_budget, self.alerting_args.budget_alert_thresholds + ) + if threshold is not None: + return "threshold_crossed", event_message + f"{threshold}% of budget consumed" elif percent_left <= SLACK_ALERTING_THRESHOLD_5_PERCENT: event = "threshold_crossed" event_message += "5% or less of budget remaining" @@ -1464,6 +1484,7 @@ Model Info: request_model: str | None = None, api_base: str | None = None, budget_alert_destination: Literal["all", "slack", "non_slack"] = "all", + budget_alert_email: bool = True, **kwargs: object, ) -> bool: """ @@ -1499,6 +1520,7 @@ Model Info: if ( budget_alert_destination != "slack" + and budget_alert_email and "email" in self.alerting and alert_type == "budget_alerts" and user_info is not None diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index dbd6f28a183..3e3bb4d5817 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -5684,6 +5684,7 @@ async def _virtual_key_max_budget_alert_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, + budget_alert_thresholds: tuple[int, ...] | None = None, ): """ Triggers a budget alert if the token has reached EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE @@ -5729,8 +5730,13 @@ async def _virtual_key_max_budget_alert_check( ) ) else: - # Old path: existing single 80% threshold — completely unchanged - alert_threshold: Final = valid_token.max_budget * EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE + if budget_alert_thresholds == (): + return + alert_threshold: Final = valid_token.max_budget * ( + min(budget_alert_thresholds) / 100 + if budget_alert_thresholds is not None + else EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE + ) if valid_token.spend >= alert_threshold and valid_token.spend < valid_token.max_budget: verbose_proxy_logger.debug( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e82f3eed7cc..742a3ac7760 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2452,10 +2452,16 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e # Check 4. Max Budget Alert Check (runs before budget enforcement # so multi-threshold 100% alerts fire on the request that crosses # max_budget, before BudgetExceededError is raised below) + configured_thresholds: Final = ( + proxy_logging_obj.slack_alerting_instance.alerting_args.budget_alert_thresholds + ) await _virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, + budget_alert_thresholds=( + tuple(configured_thresholds) if configured_thresholds is not None else None + ), ) # Check 5. Token Spend is under budget diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5828ed94984..9c6f4141d95 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -156,6 +156,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ( AlertType, CallInfo, + Litellm_EntityType, LiteLLM_VerificationTokenView, Member, UserAPIKeyAuth, @@ -1314,7 +1315,7 @@ class ProxyLogging: if alert_to_webhook_url is not None: self.alert_to_webhook_url = alert_to_webhook_url updated_slack_alerting = True - if alert_type_config is not None: + if alert_type_config is not None or alerting_args is not None: updated_slack_alerting = True if updated_slack_alerting is True: @@ -3070,14 +3071,22 @@ class ProxyLogging: # do nothing if alerting is not switched on (unless it's a soft_budget alert with team-specific emails) return + configured_thresholds: Final = self.slack_alerting_instance.alerting_args.budget_alert_thresholds if self.alerting is not None and ( "slack" in self.alerting or "ms_teams" in self.alerting or "webhook" in self.alerting ): if self.slack_alerting_instance is not None: - await self.slack_alerting_instance.budget_alerts( - type=type, - user_info=user_info, - ) + if configured_thresholds is None: + await self.slack_alerting_instance.budget_alerts(type=type, user_info=user_info) + else: + await self.slack_alerting_instance.budget_alerts( + type=type, + user_info=user_info, + send_threshold_email=( + self.email_logging_instance is None + or user_info.event_group not in (Litellm_EntityType.KEY, Litellm_EntityType.TEAM_MEMBER) + ), + ) # Call email_logging_instance if: # 1. "email" is in alerting config, OR @@ -3088,6 +3097,7 @@ class ProxyLogging: await self.email_logging_instance.budget_alerts( type=type, user_info=user_info, + budget_alert_thresholds=(tuple(configured_thresholds) if configured_thresholds is not None else None), ) async def alerting_handler( diff --git a/litellm/types/integrations/slack_alerting.py b/litellm/types/integrations/slack_alerting.py index 66d0bb7f799..ead017ad8e9 100644 --- a/litellm/types/integrations/slack_alerting.py +++ b/litellm/types/integrations/slack_alerting.py @@ -5,7 +5,7 @@ from datetime import datetime as dt from enum import Enum from typing import Annotated, Final -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.types.utils import LiteLLMPydanticObjectBase @@ -74,6 +74,22 @@ class SlackAlertingArgs(LiteLLMPydanticObjectBase): default=None, description="Case-sensitive key alias glob patterns for Slack budget alerts. Null allows all budget alerts; an empty list disables them.", ) + budget_alert_thresholds: ( # mutable-ok: public configuration accepts and serializes a list + Annotated[list[Annotated[int, Field(strict=True, ge=1, le=99)]], Field(strict=True)] | None + ) = Field( + default=None, + description="Consumed budget percentages for enabled alert destinations. Null preserves legacy thresholds; an empty list disables percentage warnings.", + ) + + @field_validator("budget_alert_thresholds") + @classmethod + def validate_budget_alert_thresholds( + cls, thresholds: list[int] | None + ) -> list[int] | None: # mutable-ok: preserve Pydantic's validated list and JSON schema + if thresholds is not None and len(thresholds) != len(set(thresholds)): + raise ValueError("Budget alert thresholds must be unique") + return thresholds + outage_alert_ttl: int = Field( default=SlackAlertingArgsEnum.outage_alert_ttl.value, description="Cache ttl for model outage alerts. Sets time-window for errors. Default is 1 minute. Value is in seconds.", diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 52e44ca5448..73682ec5991 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -2,25 +2,117 @@ import asyncio import json import os import unittest.mock as mock -from unittest.mock import patch +from typing import Final +from unittest.mock import AsyncMock, patch +import httpx import pytest from fastapi.testclient import TestClient - -from litellm.caching.caching import DualCache from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( BaseEmailLogger, ) - +from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import ( + SendGridEmailLogger, +) from litellm_enterprise.types.enterprise_callbacks.send_emails import ( EmailEvent, SendKeyCreatedEmailEvent, SendKeyRotatedEmailEvent, ) +from litellm.caching.caching import DualCache +from litellm.constants import EMAIL_BUDGET_ALERT_TTL from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER from litellm.proxy._types import CallInfo, Litellm_EntityType, WebhookEvent -from litellm.constants import EMAIL_BUDGET_ALERT_TTL + + +@pytest.mark.asyncio +async def test_common_budget_thresholds_use_real_email_rendering_and_per_level_dedup( + monkeypatch, +) -> None: + monkeypatch.setenv("SENDGRID_API_KEY", "synthetic-not-a-secret") + transport: Final = AsyncMock() + transport.post.return_value = httpx.Response(202) + logger: Final = SendGridEmailLogger() + logger.async_httpx_client = transport + info: Final = CallInfo( + spend=70, + max_budget=100, + token="synthetic", + user_email="owner@example.test", + event_group=Litellm_EntityType.KEY, + ) + for spend, count in ( + (69, 0), + (70, 1), + (70, 1), + (85, 2), + (96, 3), + (96, 3), + (100, 3), + ): + await logger.budget_alerts( + type="max_budget_alert", + user_info=info.model_copy(update={"spend": spend}), + budget_alert_thresholds=(95, 70, 85), + ) + assert transport.post.await_count == count + for call, pct in zip(transport.post.await_args_list, (70, 85, 95)): + payload: Final = call.kwargs["json"] + assert payload["personalizations"][0]["to"] == [{"email": "owner@example.test"}] + assert f"{pct}%" in payload["personalizations"][0]["subject"] + assert f"{pct}%" in payload["content"][0]["value"] + await logger.budget_alerts( + type="max_budget_alert", user_info=info, budget_alert_thresholds=() + ) + assert transport.post.await_count == 3 + await logger.budget_alerts( + type="max_budget_alert", + user_info=info.model_copy( + update={"max_budget_alert_emails": {"50": ["finance@example.test"]}} + ), + budget_alert_thresholds=(), + ) + assert transport.post.await_count == 4 + assert { + email["email"] + for email in transport.post.await_args.kwargs["json"]["personalizations"][0][ + "to" + ] + } == { + "owner@example.test", + "finance@example.test", + } + assert "50%" in transport.post.await_args.kwargs["json"]["content"][0]["value"] + + +@pytest.mark.asyncio +async def test_common_email_threshold_failed_send_releases_claim(monkeypatch) -> None: + monkeypatch.setenv("SENDGRID_API_KEY", "synthetic-not-a-secret") + transport: Final = AsyncMock() + transport.post.side_effect = ( + httpx.ConnectError("synthetic failure"), + httpx.Response(202), + ) + logger: Final = SendGridEmailLogger() + logger.async_httpx_client = transport + info: Final = CallInfo( + spend=70, + max_budget=100, + token="synthetic", + user_email="owner@example.test", + event_group=Litellm_EntityType.KEY, + ) + await logger.budget_alerts( + type="max_budget_alert", user_info=info, budget_alert_thresholds=(70,) + ) + await logger.budget_alerts( + type="max_budget_alert", user_info=info, budget_alert_thresholds=(70,) + ) + await logger.budget_alerts( + type="max_budget_alert", user_info=info, budget_alert_thresholds=(70,) + ) + assert transport.post.await_count == 2 @pytest.fixture(autouse=True) diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py index 29ac5c21439..7ddc732fda1 100644 --- a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py @@ -26,6 +26,128 @@ from litellm.types.integrations.slack_alerting import ( ) +@pytest.mark.parametrize("thresholds", ([0], [100], [70, 70], [True], [70.0], ["70"], "70", (70,), {})) +def test_budget_alert_thresholds_reject_invalid_configuration(thresholds: object) -> None: + with pytest.raises(ValidationError): + SlackAlertingArgs.model_validate({"budget_alert_thresholds": thresholds}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("digest", (False, True)) +async def test_configured_budget_thresholds_notify_each_level_and_destination(digest: bool, monkeypatch) -> None: + http_handler: Final = _webhook_accepting_posts() + monkeypatch.setenv("WEBHOOK_URL", "https://example.test/budget") + monkeypatch.setenv("MS_TEAMS_WEBHOOK_URL", "https://example.test/teams") + alerts: Final = SlackAlerting( + alerting=["slack", "webhook", "ms_teams"], + default_webhook_url=SLACK_WEBHOOK_URL, + alerting_args={"budget_alert_thresholds": [95, 70, 85]}, + alert_type_config={"budget_alerts": {"digest": digest, "digest_interval": 0}}, + async_http_handler=http_handler, + ) + for spend, expected_posts in ((69, 0), (70, 3), (70, 3), (85, 6), (96, 9), (96, 9), (100, 12)): + await alerts.budget_alerts( + type="token_budget", + user_info=CallInfo( + spend=spend, + max_budget=100, + token="synthetic", + key_alias="example-api", + event_group=Litellm_EntityType.KEY, + ), + ) + await alerts._flush_digest_buckets() + await alerts.flush_queue() + assert http_handler.post.await_count == expected_posts + webhook_events: Final = tuple( + json.loads(call.kwargs["data"]) + for call in http_handler.post.await_args_list + if call.kwargs["url"] == "https://example.test/budget" + ) + assert tuple(event["event"] for event in webhook_events) == ( + "threshold_crossed", + "threshold_crossed", + "threshold_crossed", + "budget_crossed", + ) + assert tuple(event["event_message"] for event in webhook_events[:3]) == tuple( + f"Key Budget: {pct}% of budget consumed" for pct in (70, 85, 95) + ) + + +@pytest.mark.asyncio +async def test_budget_threshold_jump_empty_reload_and_soft_budget() -> None: + http_handler: Final = _webhook_accepting_posts() + alerts: Final = SlackAlerting( + alerting=["slack"], + default_webhook_url=SLACK_WEBHOOK_URL, + alerting_args={"budget_alert_thresholds": [70, 85, 95]}, + async_http_handler=http_handler, + ) + info: Final = CallInfo(spend=96, max_budget=100, token="synthetic", event_group=Litellm_EntityType.KEY) + await alerts.budget_alerts(type="token_budget", user_info=info) + await alerts.flush_queue() + assert len(_posted_slack_bodies(http_handler)) == 1 + assert "95% of budget consumed" in _posted_slack_bodies(http_handler)[0]["text"] + alerts.update_values(alerting_args={"budget_alert_thresholds": []}) + await alerts.budget_alerts(type="token_budget", user_info=info.model_copy(update={"spend": 99})) + await alerts.flush_queue() + assert http_handler.post.await_count == 1 + await alerts.budget_alerts(type="soft_budget", user_info=info.model_copy(update={"spend": 60, "soft_budget": 50})) + await alerts.budget_alerts(type="token_budget", user_info=info.model_copy(update={"spend": 100})) + await alerts.flush_queue() + assert http_handler.post.await_count == 3 + alerts.update_values(alerting_args={"budget_alert_thresholds": [50]}) + await alerts.budget_alerts(type="token_budget", user_info=info.model_copy(update={"spend": 50})) + await alerts.flush_queue() + assert "50% of budget consumed" in _posted_slack_bodies(http_handler)[-1]["text"] + + +@pytest.mark.asyncio +async def test_configured_thresholds_preserve_soft_and_projected_events() -> None: + http_handler: Final = _webhook_accepting_posts() + alerts: Final = SlackAlerting( + alerting=["slack"], + default_webhook_url=SLACK_WEBHOOK_URL, + alerting_args={"budget_alert_thresholds": [70, 95]}, + async_http_handler=http_handler, + ) + info: Final = CallInfo( + spend=70, max_budget=100, soft_budget=50, token="synthetic", event_group=Litellm_EntityType.KEY + ) + await alerts.budget_alerts(type="soft_budget", user_info=info, send_threshold_email=False) + await alerts.budget_alerts(type="token_budget", user_info=info, send_threshold_email=False) + await alerts.budget_alerts( + type="projected_limit_exceeded", + user_info=info.model_copy(update={"soft_budget": 90}), + send_threshold_email=False, + ) + await alerts.flush_queue() + assert http_handler.post.await_count == 2 + assert "Soft Budget Crossed" in _posted_slack_bodies(http_handler)[0]["text"] + assert "Projected Limit Exceeded" in _posted_slack_bodies(http_handler)[1]["text"] + + +@pytest.mark.asyncio +async def test_configured_threshold_respects_pre_upgrade_sent_marker() -> None: + cache: Final = DualCache() + await cache.async_set_cache("budget_alerts:threshold_crossed:synthetic", "SENT", ttl=86400) + http_handler: Final = _webhook_accepting_posts() + alerts: Final = SlackAlerting( + alerting=["slack"], + internal_usage_cache=cache, + default_webhook_url=SLACK_WEBHOOK_URL, + alerting_args={"budget_alert_thresholds": [70]}, + async_http_handler=http_handler, + ) + await alerts.budget_alerts( + type="token_budget", + user_info=CallInfo(spend=70, max_budget=100, token="synthetic", event_group=Litellm_EntityType.KEY), + ) + await alerts.flush_queue() + http_handler.post.assert_not_awaited() + + class TestSlackAlerting(unittest.TestCase): def setUp(self): self.slack_alerting = SlackAlerting() diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 6c8b6571991..327c084373b 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -3820,6 +3820,43 @@ async def test_virtual_key_soft_budget_check_scenarios(spend, soft_budget, expec ) +@pytest.mark.asyncio +@pytest.mark.parametrize("thresholds,spend,expected", (([70, 95], 69, 0), ([70, 95], 70, 1), ([], 85, 0))) +async def test_common_email_threshold_auth_gate(thresholds, spend: int, expected: int, monkeypatch) -> None: + from unittest.mock import AsyncMock + import httpx + from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting + from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import SendGridEmailLogger + + monkeypatch.setenv("SENDGRID_API_KEY", "synthetic-not-a-secret") + transport: Final = AsyncMock() + transport.post.return_value = httpx.Response(202) + from litellm.proxy.utils import ProxyLogging + + proxy_logging: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) + proxy_logging.alerting = ["email"] + proxy_logging.slack_alerting_instance = SlackAlerting( + alerting=["email"], alerting_args={"budget_alert_thresholds": thresholds} + ) + email_logger: Final = SendGridEmailLogger() + email_logger.async_httpx_client = transport + proxy_logging.email_logging_instance = email_logger + await _virtual_key_max_budget_alert_check( + valid_token=UserAPIKeyAuth(token="synthetic", spend=spend, max_budget=100), + proxy_logging_obj=proxy_logging, + user_obj=LiteLLM_UserTable(user_id="owner", user_email="owner@example.test"), + budget_alert_thresholds=tuple(thresholds), + ) + pending: Final = tuple( + task + for task in asyncio.all_tasks() + if task is not asyncio.current_task() and task.get_coro().__qualname__ == "ProxyLogging.budget_alerts" + ) + if pending: + await asyncio.gather(*pending) + assert transport.post.await_count == expected + + @pytest.mark.asyncio async def test_virtual_key_max_budget_alert_check_with_user_obj(): """Test _virtual_key_max_budget_alert_check includes user_email when user_obj is provided""" @@ -10151,7 +10188,9 @@ async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache from litellm.proxy._types import LiteLLM_AccessGroupTable from litellm.proxy.auth.auth_checks import get_access_object - stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"]) + stale: Final = LiteLLM_AccessGroupTable( + access_group_id="group", access_group_name="Policy", access_model_names=["old"] + ) current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []}) client: Final = MagicMock() client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current) @@ -10303,9 +10342,12 @@ async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailab from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache database: Final = MagicMock() - database.get_data = AsyncMock(return_value=UserAPIKeyAuth( - object_permission_id="grant", object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"]) - )) + database.get_data = AsyncMock( + return_value=UserAPIKeyAuth( + object_permission_id="grant", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"]), + ) + ) database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock( return_value=None, side_effect=None if missing else RuntimeError("writer unavailable") ) @@ -10328,7 +10370,9 @@ async def test_authoritative_group_grants_propagate_policy_outages( database: Final = MagicMock() database.db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable")) - database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable")) + database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock( + side_effect=RuntimeError("database unavailable") + ) monkeypatch.setattr(proxy_server, "prisma_client", database) monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache()) if strict: diff --git a/tests/unit/proxy/utils/proxy_logging/test_alerting.py b/tests/unit/proxy/utils/proxy_logging/test_alerting.py index 77c0f71dbf9..a19eca2e7c0 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_alerting.py +++ b/tests/unit/proxy/utils/proxy_logging/test_alerting.py @@ -7,14 +7,63 @@ Covers ``failed_tracking_alert``, ``budget_alerts``, ``alerting_handler``, from __future__ import annotations from datetime import datetime -from typing import Any, Dict +from typing import Any, Dict, Final from unittest.mock import AsyncMock, MagicMock +import httpx import pytest from fastapi import HTTPException +from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import SendGridEmailLogger import litellm -from litellm.proxy._types import AlertType, CallInfo +from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting +from litellm.proxy._types import AlertType, CallInfo, Litellm_EntityType + + +@pytest.mark.asyncio +@pytest.mark.parametrize("destinations", (["email"], ["email", "slack", "webhook", "ms_teams"])) +async def test_common_thresholds_propagate_to_email_without_duplicate_smtp( + proxy_logging, monkeypatch, destinations +) -> None: + monkeypatch.setenv("SENDGRID_API_KEY", "synthetic-not-a-secret") + monkeypatch.setenv("WEBHOOK_URL", "https://example.test/budget") + monkeypatch.setenv("MS_TEAMS_WEBHOOK_URL", "https://example.test/teams") + transport: Final = AsyncMock() + transport.post.return_value = httpx.Response(200) + proxy_logging.alerting = destinations + proxy_logging.slack_alerting_instance = SlackAlerting( + alerting=destinations, + default_webhook_url="https://example.test/slack", + alerting_args={"budget_alert_thresholds": [70, 95]}, + async_http_handler=transport, + ) + email_logger: Final = SendGridEmailLogger() + email_logger.async_httpx_client = transport + proxy_logging.email_logging_instance = email_logger + info: Final = CallInfo( + spend=70, max_budget=100, token="synthetic", user_email="owner@example.test", event_group=Litellm_EntityType.KEY + ) + for spend in (70, 70, 96, 96): + event: Final = info.model_copy(update={"spend": spend}) + await proxy_logging.budget_alerts(type="token_budget", user_info=event) + await proxy_logging.budget_alerts(type="max_budget_alert", user_info=event) + await proxy_logging.slack_alerting_instance.flush_queue() + email_posts: Final = tuple(call for call in transport.post.await_args_list if "json" in call.kwargs) + assert len(email_posts) == 2 + assert "70%" in email_posts[0].kwargs["json"]["personalizations"][0]["subject"] + assert "95%" in email_posts[1].kwargs["json"]["personalizations"][0]["subject"] + assert transport.post.await_count == (2 if destinations == ["email"] else 8) + + +@pytest.mark.asyncio +async def test_alerting_args_only_reload_validates_and_preserves_last_valid(proxy_logging) -> None: + proxy_logging.update_values(alerting_args={"budget_alert_thresholds": [70, 95]}) + assert proxy_logging.slack_alerting_instance.alerting_args.budget_alert_thresholds == [70, 95] + from pydantic import ValidationError + + with pytest.raises(ValidationError): + proxy_logging.update_values(alerting_args={"budget_alert_thresholds": [True]}) + assert proxy_logging.slack_alerting_instance.alerting_args.budget_alert_thresholds == [70, 95] # --------------------------------------------------------------------------- @@ -158,9 +207,7 @@ async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_global(proxy @pytest.mark.asyncio async def test_budget_alerts_slack_failure_raises(proxy_logging): proxy_logging.alerting = ["slack"] - proxy_logging.slack_alerting_instance = MagicMock( - budget_alerts=AsyncMock(side_effect=ConnectionError("slack")) - ) + proxy_logging.slack_alerting_instance = MagicMock(budget_alerts=AsyncMock(side_effect=ConnectionError("slack"))) proxy_logging.email_logging_instance = None with pytest.raises(ConnectionError): await proxy_logging.budget_alerts(type="user_budget", user_info=_user_info()) @@ -281,11 +328,7 @@ async def test_failure_handler_with_capture_exception_invoked(proxy_logging, mon async def test_failure_handler_propagates_service_logging_error_raises(proxy_logging, monkeypatch): proxy_logging.alert_types = [AlertType.db_exceptions] proxy_logging.alerting_handler = AsyncMock() - proxy_logging.service_logging_obj = MagicMock( - async_service_failure_hook=AsyncMock(side_effect=RuntimeError("svc")) - ) + proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock(side_effect=RuntimeError("svc"))) monkeypatch.setattr(litellm.utils, "capture_exception", None) with pytest.raises(RuntimeError): - await proxy_logging.failure_handler( - original_exception=Exception("x"), duration=0.0, call_type="db_read" - ) + await proxy_logging.failure_handler(original_exception=Exception("x"), duration=0.0, call_type="db_read")