feat(alerting): configure common budget notification percentages

This commit is contained in:
XAVIER ALMENDROS 2026-10-04 19:17:41 +02:00
parent 0fdce6ef63
commit be50d86c74
11 changed files with 429 additions and 39 deletions

View file

@ -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}",

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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.",

View file

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

View file

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

View file

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

View file

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