mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(alerting): configure common budget notification percentages
This commit is contained in:
parent
0fdce6ef63
commit
be50d86c74
11 changed files with 429 additions and 39 deletions
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue