From 47401f29f8dc8fca4802906b149bc1478c6b369b Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 4 Aug 2026 13:27:31 -0700 Subject: [PATCH 1/3] fix(proxy): dedupe budget alert emails across replicas without redis The dedup claim added in #32011 lives in DualCache, which is memory-only when Redis is not configured, so each replica wins its own claim and sends its own copy of the same alert. It is also lost on restart and expires with EMAIL_BUDGET_ALERT_TTL, so a long-crossed threshold keeps re-alerting. Record the claim in a new LiteLLM_BudgetAlertSent table instead. The unique constraint makes the insert an atomic cross-replica claim, and the stored budget window re-arms the alert when the budget period resets or max_budget changes rather than on a timer. Falls open when the claim cannot be recorded. Resolves LIT-4172 --- .../send_emails/base_email.py | 145 +++++- .../migration.sql | 15 + .../litellm_proxy_extras/schema.prisma | 12 + litellm/constants.py | 2 + .../SlackAlerting/slack_alerting.py | 1 + litellm/proxy/_types.py | 4 + litellm/proxy/auth/auth_checks.py | 2 + litellm/proxy/db/budget_alert_claim.py | 171 +++++++ litellm/proxy/db/exception_handler.py | 14 + .../key_management_endpoints.py | 6 + litellm/proxy/schema.prisma | 12 + schema.prisma | 12 + .../send_emails/test_base_email.py | 454 +++++++++++++++++- .../proxy/auth/test_auth_checks.py | 4 + .../proxy/db/test_exception_handler.py | 23 + .../test_key_management_endpoints.py | 8 + 16 files changed, 864 insertions(+), 21 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql create mode 100644 litellm/proxy/db/budget_alert_claim.py 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 e7898cac565..dbcd1a7042f 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -20,6 +20,8 @@ from litellm.caching.caching import DualCache from litellm.constants import ( EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE, EMAIL_BUDGET_ALERT_TTL, + MAX_BUDGET_ALERT_TYPE, + MAX_BUDGET_DEFAULT_ALERT_TYPE, ) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER @@ -44,6 +46,11 @@ from litellm.proxy._types import ( UserAPIKeyAuth, WebhookEvent, ) +from litellm.proxy.db.budget_alert_claim import ( + claim_budget_alert_slot, + get_budget_window, + release_budget_alert_slot, +) from litellm.secret_managers.main import get_secret_bool from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL @@ -542,16 +549,18 @@ class BaseEmailLogger(CustomLogger): _id = user_info.token or user_info.user_id or "default_id" _cache_key = f"email_budget_alerts:max_budget_alert:{_id}" - send_count = await _cache.async_increment_cache( - key=_cache_key, - value=1, - ttl=EMAIL_BUDGET_ALERT_TTL, + # Calculate percentage + percentage = int( + EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100 ) - if send_count is None or send_count <= 1: - # Calculate percentage - percentage = int( - EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100 - ) + + if await self._claim_max_budget_alert_send( + cache=_cache, + cache_key=_cache_key, + user_info=user_info, + threshold_pct=percentage, + alert_type=MAX_BUDGET_DEFAULT_ALERT_TYPE, + ): # Create WebhookEvent for max budget alert event_message = f"Max Budget Alert - {percentage}% of Maximum Budget Reached" @@ -572,6 +581,7 @@ class BaseEmailLogger(CustomLogger): projected_exceeded_date=user_info.projected_exceeded_date, projected_spend=user_info.projected_spend, event_group=user_info.event_group, + budget_reset_at=user_info.budget_reset_at, ) try: @@ -581,7 +591,13 @@ class BaseEmailLogger(CustomLogger): f"Error sending max budget alert email: {e}", exc_info=True, ) - await self._release_budget_alert_claim(_cache, _cache_key) + await self._release_max_budget_alert_claim( + cache=_cache, + cache_key=_cache_key, + user_info=user_info, + threshold_pct=percentage, + alert_type=MAX_BUDGET_DEFAULT_ALERT_TYPE, + ) return async def _handle_multi_threshold_max_budget_alert( @@ -624,12 +640,12 @@ class BaseEmailLogger(CustomLogger): continue recipient_emails = list(set(emails)) - send_count = await _cache.async_increment_cache( - key=_cache_key, - value=1, - ttl=EMAIL_BUDGET_ALERT_TTL, - ) - if send_count is not None and send_count > 1: + if not await self._claim_max_budget_alert_send( + cache=_cache, + cache_key=_cache_key, + user_info=user_info, + threshold_pct=threshold_pct, + ): continue event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached" @@ -650,6 +666,7 @@ class BaseEmailLogger(CustomLogger): projected_exceeded_date=user_info.projected_exceeded_date, projected_spend=user_info.projected_spend, event_group=user_info.event_group, + budget_reset_at=user_info.budget_reset_at, ) try: @@ -663,7 +680,12 @@ class BaseEmailLogger(CustomLogger): f"Error sending multi-threshold max budget alert email for {threshold_pct}%: {e}", exc_info=True, ) - await self._release_budget_alert_claim(_cache, _cache_key) + await self._release_max_budget_alert_claim( + cache=_cache, + cache_key=_cache_key, + user_info=user_info, + threshold_pct=threshold_pct, + ) async def _release_budget_alert_claim(self, cache: DualCache, cache_key: str) -> None: try: @@ -674,6 +696,95 @@ class BaseEmailLogger(CustomLogger): cache_key, ) + @staticmethod + def _budget_window(user_info: CallInfo) -> str: + return get_budget_window(user_info.budget_reset_at, user_info.max_budget) + + @staticmethod + def _windowed_cache_key(cache_key: str, budget_window: str) -> str: + """ + Scope the in-memory claim to the budget window it belongs to. Without this + the claim would still suppress the alert for up to EMAIL_BUDGET_ALERT_TTL + after the budget resets, even though the durable claim has re-armed. + """ + return f"{cache_key}:{budget_window}" + + @staticmethod + def _durable_claim_entity(user_info: CallInfo) -> tuple[str, str] | None: + """ + (entity_type, entity_id) to key the durable claim on, or None when the alert + carries no stable identifier and can only be deduped in memory. + """ + entity_id = user_info.token or user_info.user_id + if not entity_id: + return None + return user_info.event_group.value, entity_id + + async def _claim_max_budget_alert_send( + self, + cache: DualCache, + cache_key: str, + user_info: CallInfo, + threshold_pct: int, + alert_type: str = MAX_BUDGET_ALERT_TYPE, + ) -> bool: + """ + Decide whether this process should send one max budget alert. + + The in-memory claim stops a replica re-sending on every request it serves. + The durable claim is what stops a SECOND replica sending the same alert, and + what keeps the alert from firing again after a restart or a cache TTL expiry: + the in-memory claim is process-local whenever Redis is not configured, so on + its own it cannot dedupe across replicas or across restarts. + """ + budget_window = self._budget_window(user_info) + send_count = await cache.async_increment_cache( + key=self._windowed_cache_key(cache_key, budget_window), + value=1, + ttl=EMAIL_BUDGET_ALERT_TTL, + ) + if send_count is not None and send_count > 1: + return False + + entity = self._durable_claim_entity(user_info) + if entity is None: + return True + + entity_type, entity_id = entity + return await claim_budget_alert_slot( + entity_type=entity_type, + entity_id=entity_id, + alert_type=alert_type, + threshold_pct=threshold_pct, + budget_window=budget_window, + ) + + async def _release_max_budget_alert_claim( + self, + cache: DualCache, + cache_key: str, + user_info: CallInfo, + threshold_pct: int, + alert_type: str = MAX_BUDGET_ALERT_TYPE, + ) -> None: + """Hand a won claim back after a failed send, so the alert can be retried.""" + budget_window = self._budget_window(user_info) + await self._release_budget_alert_claim( + cache, self._windowed_cache_key(cache_key, budget_window) + ) + + entity = self._durable_claim_entity(user_info) + if entity is None: + return + entity_type, entity_id = entity + await release_budget_alert_slot( + entity_type=entity_type, + entity_id=entity_id, + alert_type=alert_type, + threshold_pct=threshold_pct, + budget_window=budget_window, + ) + async def _get_email_params( self, email_event: EmailEvent, diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql new file mode 100644 index 00000000000..958799cebf0 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql @@ -0,0 +1,15 @@ +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_BudgetAlertSent" ( + "id" TEXT NOT NULL, + "entity_type" TEXT NOT NULL, + "entity_id" TEXT NOT NULL, + "alert_type" TEXT NOT NULL, + "threshold_pct" INTEGER NOT NULL, + "budget_window" TEXT NOT NULL, + "sent_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_BudgetAlertSent_pkey" PRIMARY KEY ("id") +); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_BudgetAlertSent_entity_type_entity_id_alert_type_th_key" ON "LiteLLM_BudgetAlertSent"("entity_type", "entity_id", "alert_type", "threshold_pct"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 0d7fa8692c8..c5f20e3b633 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1467,3 +1467,15 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +model LiteLLM_BudgetAlertSent { + id String @id @default(uuid()) + entity_type String + entity_id String + alert_type String + threshold_pct Int + budget_window String + sent_at DateTime @default(now()) + + @@unique([entity_type, entity_id, alert_type, threshold_pct]) +} diff --git a/litellm/constants.py b/litellm/constants.py index 164f5a77a76..5b4d9af925a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -464,6 +464,8 @@ EMAIL_BUDGET_ALERT_TTL = int(os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)) EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE = float( os.getenv("EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE", 0.8) ) # 80% of max budget +MAX_BUDGET_ALERT_TYPE = "max_budget_alert" +MAX_BUDGET_DEFAULT_ALERT_TYPE = "max_budget_alert_default_threshold" ############### LLM Provider Constants ############### ### ANTHROPIC CONSTANTS ### ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01") diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 0d842a5889a..39ec7c4069a 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -631,6 +631,7 @@ class SlackAlerting(CustomBatchLogger): """ _all_fields_as_dict = user_info.model_dump(exclude_none=True) _all_fields_as_dict.pop("token") + _all_fields_as_dict.pop("budget_reset_at", None) msg = "" for k, v in _all_fields_as_dict.items(): if isinstance(v, Litellm_EntityType): diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index fd21c5b8334..9f4e158b75a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3092,6 +3092,10 @@ class CallInfo(LiteLLMPydanticObjectBase): default=None, description="Map of threshold percentage to email recipients (e.g., {'50': ['a@co.com'], '75': ['a@co.com', 'b@co.com']})", ) + budget_reset_at: datetime | None = Field( + default=None, + description="End of the budget period this alert belongs to. Used to re-arm the alert once the budget rolls over.", + ) class WebhookEvent(CallInfo): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c186bbc15ee..0bf94fa0348 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3846,6 +3846,7 @@ async def _virtual_key_max_budget_alert_check( key_alias=valid_token.key_alias, event_group=Litellm_EntityType.KEY, max_budget_alert_emails=alert_email_config, + budget_reset_at=valid_token.budget_reset_at, ) asyncio.create_task( proxy_logging_obj.budget_alerts( @@ -3877,6 +3878,7 @@ async def _virtual_key_max_budget_alert_check( user_email=owner_email, key_alias=valid_token.key_alias, event_group=Litellm_EntityType.KEY, + budget_reset_at=valid_token.budget_reset_at, ) asyncio.create_task( diff --git a/litellm/proxy/db/budget_alert_claim.py b/litellm/proxy/db/budget_alert_claim.py new file mode 100644 index 00000000000..66f209af08a --- /dev/null +++ b/litellm/proxy/db/budget_alert_claim.py @@ -0,0 +1,171 @@ +""" +Durable, cross-replica claim for "this budget alert has already been sent". + +The in-process claim used by the email alerting path lives in ``DualCache``, which +is memory-only when Redis is not configured. Every replica therefore wins its own +claim and sends its own copy of the same alert, and the claim is lost on restart. +This module records the claim in the database instead, so it is shared by all +replicas and survives restarts. + +The claim is scoped to the entity's budget window: the budget period it belongs to +plus the budget it was measured against. The alert therefore re-arms when the budget +period rolls over or the budget itself is changed, rather than on a fixed TTL. +""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from datetime import datetime, timezone + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler + + +def get_budget_window(budget_reset_at: datetime | None, max_budget: float | None) -> str: + """ + Stable identity for the budget an alert was raised against. + + ``budget_reset_at`` alone is not enough: a key with no ``budget_duration`` never + rolls over, so a claim keyed on it alone would silence that threshold for the + lifetime of the key. Including ``max_budget`` means raising or lowering the budget + moves every threshold and re-arms the alerts, which is what an operator expects. + """ + period = budget_reset_at.isoformat() if budget_reset_at is not None else "" + return f"{period}|{max_budget}" + + +def _claim_identity( + entity_type: str, + entity_id: str, + alert_type: str, + threshold_pct: int, +) -> Mapping[str, str | int]: + """The columns the unique constraint is built on, as a prisma query payload.""" + return { # mutable-ok: prisma query payloads must be plain dicts + "entity_type": entity_type, + "entity_id": entity_id, + "alert_type": alert_type, + "threshold_pct": threshold_pct, + } + + +async def claim_budget_alert_slot( + entity_type: str, + entity_id: str, + alert_type: str, + threshold_pct: int, + budget_window: str, +) -> bool: + """ + Try to become the single sender of one budget alert. + + Returns True iff this caller won the claim and must send the alert, and False + iff another replica, or an earlier request in this budget window, already sent + it. + + Inserting the claim row is the atomic part: the unique constraint lets exactly + one replica succeed. A row that already exists is taken over only when it + belongs to an older budget window, which is the conditional update below. + + Fails open. If there is no database, or the database rejects the claim for any + reason other than the alert already being claimed, this returns True, because a + duplicate alert is a better failure than a missed budget alert. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return True + + identity = _claim_identity(entity_type, entity_id, alert_type, threshold_pct) + + try: + await prisma_client.db.litellm_budgetalertsent.create( + data={**identity, "budget_window": budget_window} # mutable-ok: prisma query payload + ) + return True + except Exception as e: # noqa: BLE001 # best-effort dedup, any failure must fall through to sending + if not PrismaDBExceptionHandler.is_unique_constraint_violation(e): + verbose_proxy_logger.warning( + "Could not record budget alert claim for %s %s at %d%%, sending anyway: %s", + entity_type, + entity_id, + threshold_pct, + e, + ) + return True + + try: + rows_updated = await prisma_client.db.litellm_budgetalertsent.update_many( + where={**identity, "budget_window": {"not": budget_window}}, # mutable-ok: prisma query payload + data={ # mutable-ok: prisma query payload + "budget_window": budget_window, + "sent_at": datetime.now(timezone.utc), + }, + ) + return int(rows_updated) > 0 + except Exception as e: # noqa: BLE001 # best-effort dedup, any failure must fall through to sending + verbose_proxy_logger.warning( + "Could not roll over budget alert claim for %s %s at %d%%, sending anyway: %s", + entity_type, + entity_id, + threshold_pct, + e, + ) + return True + + +async def release_budget_alert_slot( + entity_type: str, + entity_id: str, + alert_type: str, + threshold_pct: int, + budget_window: str, +) -> None: + """ + Give a won claim back after the alert failed to send, so the next request can + retry it. Without this, one failed send would silence the whole fleet for the + rest of the budget window. + + Scoped to ``budget_window`` so a slow failing send cannot delete a claim that + another replica has already taken over for a later window. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return + + identity = _claim_identity(entity_type, entity_id, alert_type, threshold_pct) + + try: + await prisma_client.db.litellm_budgetalertsent.delete_many( + where={**identity, "budget_window": budget_window} # mutable-ok: prisma query payload + ) + except Exception as e: # noqa: BLE001 # releasing the claim is best-effort, it also expires with the window + verbose_proxy_logger.debug( + "Failed to release budget alert claim for %s %s at %d%%: %s", + entity_type, + entity_id, + threshold_pct, + e, + ) + + +async def delete_budget_alert_claims(entity_type: str, entity_ids: Sequence[str]) -> None: + """ + Drop the claim rows for entities that no longer exist, so deleting a key does + not leave its alert claims behind. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None or not entity_ids: + return + + try: + await prisma_client.db.litellm_budgetalertsent.delete_many( + where={ # mutable-ok: prisma query payload + "entity_type": entity_type, + "entity_id": {"in": tuple(entity_ids)}, # mutable-ok: prisma query payload + } + ) + except Exception as e: # noqa: BLE001 # cleanup is best-effort, it must never fail a deletion + verbose_proxy_logger.warning("Failed to delete budget alert claims for %s %s: %s", entity_type, entity_ids, e) diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 28551f49031..33b42e4acb6 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -92,6 +92,20 @@ class PrismaDBExceptionHandler: return type(e) is prisma.errors.DataError + @staticmethod + def is_unique_constraint_violation(e: Exception) -> bool: + """True iff ``e`` is a prisma unique-constraint violation, i.e. the row the + caller tried to insert already exists. + + Lets callers use an insert as an atomic cross-replica claim without + importing prisma themselves: prisma is a proxy-only dependency, and a + module reachable from a base ``import litellm`` cannot import it at the top + level. + """ + import prisma + + return isinstance(e, prisma.errors.UniqueViolationError) + @staticmethod def is_database_transport_error(e: Exception) -> bool: """ diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index f37bc6da23e..3cc56eaf480 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -68,6 +68,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.db.budget_alert_claim import delete_budget_alert_claims from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import ( VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, @@ -4120,6 +4121,11 @@ async def delete_verification_tokens( verbose_proxy_logger.debug(traceback.format_exc()) raise e + await delete_budget_alert_claims( + entity_type=Litellm_EntityType.KEY.value, + entity_ids=tuple(key.token for key in _keys_being_deleted), + ) + for key in tokens: user_api_key_cache.delete_cache(key) # remove hash token from cache diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 0d7fa8692c8..c5f20e3b633 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1467,3 +1467,15 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +model LiteLLM_BudgetAlertSent { + id String @id @default(uuid()) + entity_type String + entity_id String + alert_type String + threshold_pct Int + budget_window String + sent_at DateTime @default(now()) + + @@unique([entity_type, entity_id, alert_type, threshold_pct]) +} diff --git a/schema.prisma b/schema.prisma index 0d7fa8692c8..c5f20e3b633 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1467,3 +1467,15 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +model LiteLLM_BudgetAlertSent { + id String @id @default(uuid()) + entity_type String + entity_id String + alert_type String + threshold_pct Int + budget_window String + sent_at DateTime @default(now()) + + @@unique([entity_type, entity_id, alert_type, threshold_pct]) +} diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 61303340570..5548b4cecdf 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -3,6 +3,9 @@ import json import os import sys import unittest.mock as mock +from collections.abc import Mapping +from datetime import datetime, timezone +from types import SimpleNamespace from unittest.mock import patch import pytest @@ -23,6 +26,12 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import ( 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 +from litellm.constants import MAX_BUDGET_ALERT_TYPE +from litellm.proxy.db.budget_alert_claim import ( + claim_budget_alert_slot, + delete_budget_alert_claims, + release_budget_alert_slot, +) @pytest.fixture(autouse=True) @@ -914,7 +923,8 @@ async def test_budget_alerts_max_budget_alert_crossed( mock_cache.async_increment_cache.assert_called_once() cache_call_args = mock_cache.async_increment_cache.call_args[1] assert ( - cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user" + cache_call_args["key"] + == "email_budget_alerts:max_budget_alert:test_user:|200.0" ) assert cache_call_args["value"] == 1 assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL @@ -955,8 +965,8 @@ async def test_multi_threshold_sends_crossed_thresholds( cache_keys = [ c[1]["key"] for c in mock_cache.async_increment_cache.call_args_list ] - assert "email_budget_alerts:max_budget_alert:50:hashed_key_1" in cache_keys - assert "email_budget_alerts:max_budget_alert:75:hashed_key_1" in cache_keys + assert "email_budget_alerts:max_budget_alert:50:hashed_key_1:|100.0" in cache_keys + assert "email_budget_alerts:max_budget_alert:75:hashed_key_1:|100.0" in cache_keys @pytest.mark.asyncio @@ -1119,7 +1129,7 @@ async def test_no_map_preserves_old_single_threshold( assert call_args["to_email"] == ["test@example.com"] # Old path cache key has no threshold percentage cache_key = mock_cache.async_increment_cache.call_args[1]["key"] - assert cache_key == "email_budget_alerts:max_budget_alert:test_user" + assert cache_key == "email_budget_alerts:max_budget_alert:test_user:|200.0" CUSTOM_SIGNATURE = "
Best,
The Acme Platform Team
" @@ -1378,3 +1388,439 @@ async def test_budget_alert_release_failure_does_not_propagate(base_email_logger ) mock_cache.async_delete_cache.assert_awaited_once() + + +class _FakeBudgetAlertTable: + """In-memory stand-in for the LiteLLM_BudgetAlertSent prisma table. + + Enforces the same unique constraint the real table does, so a claim taken by + one replica is visible to every other replica sharing this instance. + """ + + _UNIQUE_FIELDS = ("entity_type", "entity_id", "alert_type", "threshold_pct") + + def __init__(self): + self.rows: list[dict] = [] + + @staticmethod + def _identity(row: Mapping[str, object]) -> tuple: + return tuple(row[f] for f in _FakeBudgetAlertTable._UNIQUE_FIELDS) + + def _matching(self, where: Mapping[str, object]) -> list[dict]: + matched = [] + for row in self.rows: + for field, expected in where.items(): + actual = row.get(field) + if isinstance(expected, dict) and "not" in expected: + if actual == expected["not"]: + break + elif isinstance(expected, dict) and "in" in expected: + if actual not in expected["in"]: + break + elif actual != expected: + break + else: + matched.append(row) + return matched + + async def create(self, data: Mapping[str, object]) -> dict: + import prisma + + if any(self._identity(r) == self._identity(data) for r in self.rows): + raise prisma.errors.UniqueViolationError({"user_facing_error": {"meta": {}}}) + row = dict(data) + self.rows.append(row) + return row + + async def update_many( + self, where: Mapping[str, object], data: Mapping[str, object] + ) -> int: + matched = self._matching(where) + for row in matched: + row.update(data) + return len(matched) + + async def delete_many(self, where: Mapping[str, object]) -> int: + matched = self._matching(where) + self.rows = [r for r in self.rows if r not in matched] + return len(matched) + + +@pytest.fixture +def shared_alert_table(monkeypatch): + """A single durable claim store shared by every simulated replica.""" + import litellm.proxy.proxy_server as proxy_server + + table = _FakeBudgetAlertTable() + fake_client = SimpleNamespace(db=SimpleNamespace(litellm_budgetalertsent=table)) + monkeypatch.setattr(proxy_server, "prisma_client", fake_client, raising=False) + return table + + +def _replica_logger() -> BaseEmailLogger: + """A logger with its OWN DualCache, i.e. a separate proxy process with no Redis.""" + logger = BaseEmailLogger() + logger.internal_usage_cache = DualCache() + return logger + + +@pytest.mark.parametrize( + "branch, alert_type, send_method, ci_kwargs", + [b for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], + ids=[b[0] for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], +) +@pytest.mark.asyncio +async def test_max_budget_alert_deduped_across_replicas_without_redis( + shared_alert_table, branch, alert_type, send_method, ci_kwargs +): + """Regression for LIT-4172 follow-up: two replicas with no Redis must send + exactly one email for one threshold crossing. + + The in-memory claim added by the first fix is process-local, so each replica + won its own claim and sent its own copy. The durable claim is what makes the + two replicas agree. + """ + sends = [] + + async def record_send(*args, **kwargs): + sends.append(kwargs.get("threshold_pct")) + + for _ in range(2): + replica = _replica_logger() + with mock.patch.object(replica, send_method, side_effect=record_send): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type=alert_type, user_info=_budget_alert_user_info(ci_kwargs) + ) + + assert len(sends) == 1 + assert len(shared_alert_table.rows) == 1 + + +@pytest.mark.asyncio +async def test_max_budget_alert_not_resent_after_replica_restart(shared_alert_table): + """A restart wipes the in-memory claim. The alert must not fire again, which + is what made a long-crossed threshold re-alert every restart and every TTL.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + user_info = _budget_alert_user_info( + dict(max_budget=100.0, spend=80.0, max_budget_alert_emails={"50": []}) + ) + for _ in range(3): + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type="max_budget_alert", user_info=user_info + ) + + assert len(sends) == 1 + + +@pytest.mark.parametrize( + "branch, alert_type, send_method, ci_kwargs", + [b for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], + ids=[b[0] for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], +) +@pytest.mark.asyncio +async def test_max_budget_alert_rearms_when_budget_window_rolls_over( + shared_alert_table, branch, alert_type, send_method, ci_kwargs +): + """A durable claim must not silence the alert forever: once the key's budget + period resets, the same threshold has to alert again.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + def user_info_for(reset_at): + return _budget_alert_user_info({**ci_kwargs, "budget_reset_at": reset_at}) + + same_window = datetime(2026, 9, 1, tzinfo=timezone.utc) + next_window = datetime(2026, 10, 1, tzinfo=timezone.utc) + for window in (same_window, same_window, next_window): + replica = _replica_logger() + with mock.patch.object(replica, send_method, side_effect=record_send): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type=alert_type, user_info=user_info_for(window) + ) + + assert len(sends) == 2 + assert len(shared_alert_table.rows) == 1 + + +@pytest.mark.parametrize( + "branch, alert_type, send_method, ci_kwargs", + [b for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], + ids=[b[0] for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], +) +@pytest.mark.asyncio +async def test_max_budget_alert_failed_send_releases_durable_claim( + shared_alert_table, branch, alert_type, send_method, ci_kwargs +): + """A send that fails must give the durable claim back, or the whole fleet is + permanently silenced for that threshold.""" + attempts = [] + + async def flaky_send(*args, **kwargs): + attempts.append(1) + if len(attempts) == 1: + raise ValueError("transient email backend failure") + + user_info = _budget_alert_user_info(ci_kwargs) + for _ in range(2): + replica = _replica_logger() + with mock.patch.object(replica, send_method, side_effect=flaky_send): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts(type=alert_type, user_info=user_info) + + assert len(attempts) == 2 + assert len(shared_alert_table.rows) == 1 + + +@pytest.mark.asyncio +async def test_max_budget_alert_sends_when_no_database_is_configured(monkeypatch): + """With no prisma client the durable claim cannot be taken. The alert must + still be sent rather than silently dropped.""" + import litellm.proxy.proxy_server as proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None, raising=False) + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type="max_budget_alert", + user_info=_budget_alert_user_info( + dict(max_budget=100.0, spend=80.0, max_budget_alert_emails={"50": []}) + ), + ) + + assert len(sends) == 1 + + +@pytest.mark.parametrize( + "branch, alert_type, send_method, ci_kwargs", + [b for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], + ids=[b[0] for b in _BUDGET_ALERT_BRANCHES if b[1] == "max_budget_alert"], +) +@pytest.mark.asyncio +async def test_max_budget_alert_rearms_when_max_budget_changes( + shared_alert_table, branch, alert_type, send_method, ci_kwargs +): + """A key with no budget_duration never rolls over, so a claim keyed only on the + reset date would silence the threshold for the life of the key. Changing the + budget moves every threshold amount and has to re-arm the alerts.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + def user_info_for(max_budget: float): + return _budget_alert_user_info( + {**ci_kwargs, "max_budget": max_budget, "budget_reset_at": None} + ) + + for max_budget in (100.0, 100.0, 90.0): + replica = _replica_logger() + with mock.patch.object(replica, send_method, side_effect=record_send): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type=alert_type, user_info=user_info_for(max_budget) + ) + + assert len(sends) == 2 + assert len(shared_alert_table.rows) == 1 + + +@pytest.mark.asyncio +async def test_release_does_not_steal_a_claim_from_a_later_window(shared_alert_table): + """A slow failing send must only release the claim it actually took, or it + deletes the row another replica has already taken over for a later window.""" + await claim_budget_alert_slot( + entity_type="key", + entity_id="hashed_key_1", + alert_type=MAX_BUDGET_ALERT_TYPE, + threshold_pct=50, + budget_window="later-window", + ) + + await release_budget_alert_slot( + entity_type="key", + entity_id="hashed_key_1", + alert_type=MAX_BUDGET_ALERT_TYPE, + threshold_pct=50, + budget_window="earlier-window", + ) + + assert len(shared_alert_table.rows) == 1 + + +@pytest.mark.asyncio +async def test_each_configured_threshold_claims_its_own_slot(shared_alert_table): + """The claim is per threshold. A key crossing two thresholds at once must send + both emails; collapsing them onto one claim silently drops the higher one.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(kwargs.get("threshold_pct")) + + user_info = _budget_alert_user_info( + dict( + max_budget=100.0, + spend=80.0, + max_budget_alert_emails={"50": ["finance@co.com"], "75": ["finance@co.com"]}, + ) + ) + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts(type="max_budget_alert", user_info=user_info) + + assert sorted(sends) == [50, 75] + assert len(shared_alert_table.rows) == 2 + + +@pytest.mark.asyncio +async def test_max_budget_alert_sends_when_claim_table_is_unavailable(monkeypatch): + """Day one of an upgrade the migration may not have run yet. A database error + that is not "already claimed" must fall through to sending, never silence.""" + import prisma + + import litellm.proxy.proxy_server as proxy_server + + class _BrokenTable: + async def create(self, data): + raise prisma.errors.TableNotFoundError({"user_facing_error": {"meta": {}}}) + + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace(db=SimpleNamespace(litellm_budgetalertsent=_BrokenTable())), + raising=False, + ) + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + for _ in range(2): + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type="max_budget_alert", + user_info=_budget_alert_user_info( + dict( + max_budget=100.0, + spend=80.0, + max_budget_alert_emails={"50": []}, + ) + ), + ) + + assert len(sends) == 2 + + +@pytest.mark.asyncio +async def test_deleting_a_key_drops_its_claims_only(shared_alert_table): + """Claim rows outlive the key unless deletion sweeps them, and a stale row for a + recycled id would suppress a real alert.""" + for entity_id in ("hashed_key_1", "hashed_key_2"): + await claim_budget_alert_slot( + entity_type="key", + entity_id=entity_id, + alert_type=MAX_BUDGET_ALERT_TYPE, + threshold_pct=50, + budget_window="|100.0", + ) + + await delete_budget_alert_claims(entity_type="key", entity_ids=("hashed_key_1",)) + + assert [r["entity_id"] for r in shared_alert_table.rows] == ["hashed_key_2"] + + +@pytest.mark.asyncio +async def test_in_memory_claim_does_not_outlive_its_budget_window(shared_alert_table): + """One long-lived replica must re-alert after the window moves. The in-memory + claim outlives the window by up to EMAIL_BUDGET_ALERT_TTL unless it is scoped + to the window too.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(1) + + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + for reset_at in ( + datetime(2026, 9, 1, tzinfo=timezone.utc), + datetime(2026, 9, 1, tzinfo=timezone.utc), + datetime(2026, 10, 1, tzinfo=timezone.utc), + ): + await replica.budget_alerts( + type="max_budget_alert", + user_info=_budget_alert_user_info( + dict( + max_budget=100.0, + spend=80.0, + max_budget_alert_emails={"50": []}, + budget_reset_at=reset_at, + ) + ), + ) + + assert len(sends) == 2 + + +@pytest.mark.asyncio +async def test_default_threshold_and_configured_threshold_do_not_share_a_claim( + shared_alert_table, +): + """The 80% default alert and a configured "80" threshold are different alerts + with different recipients; one must not consume the other's claim.""" + sends = [] + + async def record_send(*args, **kwargs): + sends.append(kwargs.get("threshold_pct")) + + replica = _replica_logger() + with mock.patch.object( + replica, "send_max_budget_alert_email", side_effect=record_send + ): + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await replica.budget_alerts( + type="max_budget_alert", + user_info=_budget_alert_user_info(dict(max_budget=100.0, spend=85.0)), + ) + await replica.budget_alerts( + type="max_budget_alert", + user_info=_budget_alert_user_info( + dict( + max_budget=100.0, + spend=85.0, + max_budget_alert_emails={"80": ["finance@co.com"]}, + ) + ), + ) + + assert len(sends) == 2 + assert len(shared_alert_table.rows) == 2 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index d1b5395c73d..34beba45098 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2624,6 +2624,7 @@ async def test_virtual_key_max_budget_alert_check_with_user_obj(): org_id="test-org", key_alias="test-key", soft_budget=50.0, + budget_reset_at=datetime(2026, 9, 1, tzinfo=timezone.utc), ) user_obj = LiteLLM_UserTable( @@ -2655,6 +2656,7 @@ async def test_virtual_key_max_budget_alert_check_with_user_obj(): assert captured_call_info.organization_id == "test-org" assert captured_call_info.key_alias == "test-key" assert captured_call_info.event_group == Litellm_EntityType.KEY + assert captured_call_info.budget_reset_at == valid_token.budget_reset_at @pytest.mark.asyncio @@ -2767,6 +2769,7 @@ async def test_virtual_key_max_budget_alert_check_with_multi_threshold_map(): user_id="test-user", key_alias="test-key", metadata={"max_budget_alert_emails": alert_config}, + budget_reset_at=datetime(2026, 9, 1, tzinfo=timezone.utc), ) user_obj = LiteLLM_UserTable( user_id="test-user", @@ -2786,6 +2789,7 @@ async def test_virtual_key_max_budget_alert_check_with_multi_threshold_map(): assert captured_call_info.max_budget_alert_emails == alert_config assert captured_call_info.user_email == "owner@co.com" assert captured_call_info.event_group == Litellm_EntityType.KEY + assert captured_call_info.budget_reset_at == valid_token.budget_reset_at @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 23099177812..de9bf809808 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -426,3 +426,26 @@ def test_handle_db_exception_with_non_db_error(): ) with pytest.raises(litellm.BudgetExceededError): PrismaDBExceptionHandler.handle_db_exception(regular_error) + + +def test_is_unique_constraint_violation(): + """ + Callers use an insert as an atomic cross-replica claim, so this predicate is what + separates "someone else already claimed it" from a database problem that has to + fail open. Table-not-found in particular must not read as a claim. + """ + import prisma + + assert ( + PrismaDBExceptionHandler.is_unique_constraint_violation( + prisma.errors.UniqueViolationError({"user_facing_error": {"meta": {}}}) + ) + is True + ) + for other in ( + prisma.errors.TableNotFoundError({"user_facing_error": {"meta": {}}}), + prisma.errors.DataError({"user_facing_error": {"meta": {}}}), + ClientNotConnectedError(), + ValueError("not a prisma error"), + ): + assert PrismaDBExceptionHandler.is_unique_constraint_violation(other) is False diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 867ef759fb3..15af6be3b59 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -4923,6 +4923,14 @@ async def test_delete_verification_tokens_persists_deleted_keys(monkeypatch): # delete_data returns the list directly, which gets wrapped in {"deleted_keys": ...} assert isinstance(result["deleted_keys"], list) assert set(result["deleted_keys"]) == {"hashed-token-1", "hashed-token-2"} + + # budget alert claim rows are keyed on the token and have no FK, so deleting a key + # has to sweep them or a stale claim silences the alert for a recycled token + claim_delete = mock_prisma_client.db.litellm_budgetalertsent.delete_many + claim_delete.assert_awaited_once() + claim_where = claim_delete.await_args.kwargs["where"] + assert claim_where["entity_type"] == "key" + assert set(claim_where["entity_id"]["in"]) == {"hashed-token-1", "hashed-token-2"} assert len(deleted_keys) == 2 From 8bcc1a3bd4fb352b89c114f8e7d2c8be7318d10f Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 4 Aug 2026 14:10:53 -0700 Subject: [PATCH 2/3] fix(proxy): cascade budget alert claims from the key instead of sweeping call sites Key deletion is not the only path that removes verification tokens; team deletion and the expired-session cleanup delete them directly, so an application-level sweep on /key/delete left claim rows orphaned. Key the claim table on the token with an ON DELETE CASCADE foreign key so every deletion path clears the rows with no application code. Verified live: deleting a key, and deleting a team that owns the key, both drop the claim rows. --- .../send_emails/base_email.py | 26 ++----- .../migration.sql | 9 ++- .../litellm_proxy_extras/schema.prisma | 8 ++- litellm/proxy/db/budget_alert_claim.py | 67 +++++-------------- .../key_management_endpoints.py | 6 -- litellm/proxy/schema.prisma | 8 ++- schema.prisma | 8 ++- .../send_emails/test_base_email.py | 30 +-------- .../test_key_management_endpoints.py | 8 --- 9 files changed, 47 insertions(+), 123 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 dbcd1a7042f..cd2e30d65fe 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -709,17 +709,6 @@ class BaseEmailLogger(CustomLogger): """ return f"{cache_key}:{budget_window}" - @staticmethod - def _durable_claim_entity(user_info: CallInfo) -> tuple[str, str] | None: - """ - (entity_type, entity_id) to key the durable claim on, or None when the alert - carries no stable identifier and can only be deduped in memory. - """ - entity_id = user_info.token or user_info.user_id - if not entity_id: - return None - return user_info.event_group.value, entity_id - async def _claim_max_budget_alert_send( self, cache: DualCache, @@ -746,14 +735,11 @@ class BaseEmailLogger(CustomLogger): if send_count is not None and send_count > 1: return False - entity = self._durable_claim_entity(user_info) - if entity is None: + if not user_info.token: return True - entity_type, entity_id = entity return await claim_budget_alert_slot( - entity_type=entity_type, - entity_id=entity_id, + token=user_info.token, alert_type=alert_type, threshold_pct=threshold_pct, budget_window=budget_window, @@ -773,13 +759,11 @@ class BaseEmailLogger(CustomLogger): cache, self._windowed_cache_key(cache_key, budget_window) ) - entity = self._durable_claim_entity(user_info) - if entity is None: + if not user_info.token: return - entity_type, entity_id = entity + await release_budget_alert_slot( - entity_type=entity_type, - entity_id=entity_id, + token=user_info.token, alert_type=alert_type, threshold_pct=threshold_pct, budget_window=budget_window, diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql index 958799cebf0..9bb0ed640f3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260804000000_add_budget_alert_sent/migration.sql @@ -1,8 +1,7 @@ -- CreateTable CREATE TABLE IF NOT EXISTS "LiteLLM_BudgetAlertSent" ( "id" TEXT NOT NULL, - "entity_type" TEXT NOT NULL, - "entity_id" TEXT NOT NULL, + "token" TEXT NOT NULL, "alert_type" TEXT NOT NULL, "threshold_pct" INTEGER NOT NULL, "budget_window" TEXT NOT NULL, @@ -12,4 +11,8 @@ CREATE TABLE IF NOT EXISTS "LiteLLM_BudgetAlertSent" ( ); -- CreateIndex -CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_BudgetAlertSent_entity_type_entity_id_alert_type_th_key" ON "LiteLLM_BudgetAlertSent"("entity_type", "entity_id", "alert_type", "threshold_pct"); +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_BudgetAlertSent_token_alert_type_threshold_pct_key" ON "LiteLLM_BudgetAlertSent"("token", "alert_type", "threshold_pct"); + +-- AddForeignKey +ALTER TABLE "LiteLLM_BudgetAlertSent" DROP CONSTRAINT IF EXISTS "LiteLLM_BudgetAlertSent_token_fkey"; +ALTER TABLE "LiteLLM_BudgetAlertSent" ADD CONSTRAINT "LiteLLM_BudgetAlertSent_token_fkey" FOREIGN KEY ("token") REFERENCES "LiteLLM_VerificationToken"("token") ON DELETE CASCADE ON UPDATE CASCADE; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index c5f20e3b633..5077b203940 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -464,6 +464,7 @@ model LiteLLM_VerificationToken { litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id]) object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) jwt_key_mappings LiteLLM_JWTKeyMapping[] + budget_alerts_sent LiteLLM_BudgetAlertSent[] // SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub" // SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2 @@ -1470,12 +1471,13 @@ model LiteLLM_WorkflowMessage { model LiteLLM_BudgetAlertSent { id String @id @default(uuid()) - entity_type String - entity_id String + token String alert_type String threshold_pct Int budget_window String sent_at DateTime @default(now()) - @@unique([entity_type, entity_id, alert_type, threshold_pct]) + verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade) + + @@unique([token, alert_type, threshold_pct]) } diff --git a/litellm/proxy/db/budget_alert_claim.py b/litellm/proxy/db/budget_alert_claim.py index 66f209af08a..9f808121669 100644 --- a/litellm/proxy/db/budget_alert_claim.py +++ b/litellm/proxy/db/budget_alert_claim.py @@ -7,14 +7,15 @@ claim and sends its own copy of the same alert, and the claim is lost on restart This module records the claim in the database instead, so it is shared by all replicas and survives restarts. -The claim is scoped to the entity's budget window: the budget period it belongs to -plus the budget it was measured against. The alert therefore re-arms when the budget -period rolls over or the budget itself is changed, rather than on a fixed TTL. +The claim is scoped to the key's budget window: the budget period it belongs to +plus the budget it was measured against. The alert therefore re-arms when the +budget period rolls over or the budget itself is changed, rather than on a fixed +TTL. Rows are removed by the schema's cascade when the key is deleted. """ from __future__ import annotations -from collections.abc import Mapping, Sequence +from collections.abc import Mapping from datetime import datetime, timezone from litellm._logging import verbose_proxy_logger @@ -34,24 +35,17 @@ def get_budget_window(budget_reset_at: datetime | None, max_budget: float | None return f"{period}|{max_budget}" -def _claim_identity( - entity_type: str, - entity_id: str, - alert_type: str, - threshold_pct: int, -) -> Mapping[str, str | int]: +def _claim_identity(token: str, alert_type: str, threshold_pct: int) -> Mapping[str, str | int]: """The columns the unique constraint is built on, as a prisma query payload.""" return { # mutable-ok: prisma query payloads must be plain dicts - "entity_type": entity_type, - "entity_id": entity_id, + "token": token, "alert_type": alert_type, "threshold_pct": threshold_pct, } async def claim_budget_alert_slot( - entity_type: str, - entity_id: str, + token: str, alert_type: str, threshold_pct: int, budget_window: str, @@ -76,7 +70,7 @@ async def claim_budget_alert_slot( if prisma_client is None: return True - identity = _claim_identity(entity_type, entity_id, alert_type, threshold_pct) + identity = _claim_identity(token, alert_type, threshold_pct) try: await prisma_client.db.litellm_budgetalertsent.create( @@ -86,9 +80,8 @@ async def claim_budget_alert_slot( except Exception as e: # noqa: BLE001 # best-effort dedup, any failure must fall through to sending if not PrismaDBExceptionHandler.is_unique_constraint_violation(e): verbose_proxy_logger.warning( - "Could not record budget alert claim for %s %s at %d%%, sending anyway: %s", - entity_type, - entity_id, + "Could not record budget alert claim for key %s at %d%%, sending anyway: %s", + token, threshold_pct, e, ) @@ -105,9 +98,8 @@ async def claim_budget_alert_slot( return int(rows_updated) > 0 except Exception as e: # noqa: BLE001 # best-effort dedup, any failure must fall through to sending verbose_proxy_logger.warning( - "Could not roll over budget alert claim for %s %s at %d%%, sending anyway: %s", - entity_type, - entity_id, + "Could not roll over budget alert claim for key %s at %d%%, sending anyway: %s", + token, threshold_pct, e, ) @@ -115,8 +107,7 @@ async def claim_budget_alert_slot( async def release_budget_alert_slot( - entity_type: str, - entity_id: str, + token: str, alert_type: str, threshold_pct: int, budget_window: str, @@ -134,38 +125,16 @@ async def release_budget_alert_slot( if prisma_client is None: return - identity = _claim_identity(entity_type, entity_id, alert_type, threshold_pct) + identity = _claim_identity(token, alert_type, threshold_pct) try: await prisma_client.db.litellm_budgetalertsent.delete_many( where={**identity, "budget_window": budget_window} # mutable-ok: prisma query payload ) except Exception as e: # noqa: BLE001 # releasing the claim is best-effort, it also expires with the window - verbose_proxy_logger.debug( - "Failed to release budget alert claim for %s %s at %d%%: %s", - entity_type, - entity_id, + verbose_proxy_logger.warning( + "Failed to release budget alert claim for key %s at %d%%: %s", + token, threshold_pct, e, ) - - -async def delete_budget_alert_claims(entity_type: str, entity_ids: Sequence[str]) -> None: - """ - Drop the claim rows for entities that no longer exist, so deleting a key does - not leave its alert claims behind. - """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None or not entity_ids: - return - - try: - await prisma_client.db.litellm_budgetalertsent.delete_many( - where={ # mutable-ok: prisma query payload - "entity_type": entity_type, - "entity_id": {"in": tuple(entity_ids)}, # mutable-ok: prisma query payload - } - ) - except Exception as e: # noqa: BLE001 # cleanup is best-effort, it must never fail a deletion - verbose_proxy_logger.warning("Failed to delete budget alert claims for %s %s: %s", entity_type, entity_ids, e) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3cc56eaf480..f37bc6da23e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -68,7 +68,6 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.db.budget_alert_claim import delete_budget_alert_claims from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import ( VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, @@ -4121,11 +4120,6 @@ async def delete_verification_tokens( verbose_proxy_logger.debug(traceback.format_exc()) raise e - await delete_budget_alert_claims( - entity_type=Litellm_EntityType.KEY.value, - entity_ids=tuple(key.token for key in _keys_being_deleted), - ) - for key in tokens: user_api_key_cache.delete_cache(key) # remove hash token from cache diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index c5f20e3b633..5077b203940 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -464,6 +464,7 @@ model LiteLLM_VerificationToken { litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id]) object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) jwt_key_mappings LiteLLM_JWTKeyMapping[] + budget_alerts_sent LiteLLM_BudgetAlertSent[] // SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub" // SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2 @@ -1470,12 +1471,13 @@ model LiteLLM_WorkflowMessage { model LiteLLM_BudgetAlertSent { id String @id @default(uuid()) - entity_type String - entity_id String + token String alert_type String threshold_pct Int budget_window String sent_at DateTime @default(now()) - @@unique([entity_type, entity_id, alert_type, threshold_pct]) + verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade) + + @@unique([token, alert_type, threshold_pct]) } diff --git a/schema.prisma b/schema.prisma index c5f20e3b633..5077b203940 100644 --- a/schema.prisma +++ b/schema.prisma @@ -464,6 +464,7 @@ model LiteLLM_VerificationToken { litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id]) object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) jwt_key_mappings LiteLLM_JWTKeyMapping[] + budget_alerts_sent LiteLLM_BudgetAlertSent[] // SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub" // SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2 @@ -1470,12 +1471,13 @@ model LiteLLM_WorkflowMessage { model LiteLLM_BudgetAlertSent { id String @id @default(uuid()) - entity_type String - entity_id String + token String alert_type String threshold_pct Int budget_window String sent_at DateTime @default(now()) - @@unique([entity_type, entity_id, alert_type, threshold_pct]) + verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade) + + @@unique([token, alert_type, threshold_pct]) } diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 5548b4cecdf..387d3c4100c 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -29,7 +29,6 @@ from litellm.constants import EMAIL_BUDGET_ALERT_TTL from litellm.constants import MAX_BUDGET_ALERT_TYPE from litellm.proxy.db.budget_alert_claim import ( claim_budget_alert_slot, - delete_budget_alert_claims, release_budget_alert_slot, ) @@ -1397,7 +1396,7 @@ class _FakeBudgetAlertTable: one replica is visible to every other replica sharing this instance. """ - _UNIQUE_FIELDS = ("entity_type", "entity_id", "alert_type", "threshold_pct") + _UNIQUE_FIELDS = ("token", "alert_type", "threshold_pct") def __init__(self): self.rows: list[dict] = [] @@ -1414,9 +1413,6 @@ class _FakeBudgetAlertTable: if isinstance(expected, dict) and "not" in expected: if actual == expected["not"]: break - elif isinstance(expected, dict) and "in" in expected: - if actual not in expected["in"]: - break elif actual != expected: break else: @@ -1650,16 +1646,14 @@ async def test_release_does_not_steal_a_claim_from_a_later_window(shared_alert_t """A slow failing send must only release the claim it actually took, or it deletes the row another replica has already taken over for a later window.""" await claim_budget_alert_slot( - entity_type="key", - entity_id="hashed_key_1", + token="hashed_key_1", alert_type=MAX_BUDGET_ALERT_TYPE, threshold_pct=50, budget_window="later-window", ) await release_budget_alert_slot( - entity_type="key", - entity_id="hashed_key_1", + token="hashed_key_1", alert_type=MAX_BUDGET_ALERT_TYPE, threshold_pct=50, budget_window="earlier-window", @@ -1738,24 +1732,6 @@ async def test_max_budget_alert_sends_when_claim_table_is_unavailable(monkeypatc assert len(sends) == 2 -@pytest.mark.asyncio -async def test_deleting_a_key_drops_its_claims_only(shared_alert_table): - """Claim rows outlive the key unless deletion sweeps them, and a stale row for a - recycled id would suppress a real alert.""" - for entity_id in ("hashed_key_1", "hashed_key_2"): - await claim_budget_alert_slot( - entity_type="key", - entity_id=entity_id, - alert_type=MAX_BUDGET_ALERT_TYPE, - threshold_pct=50, - budget_window="|100.0", - ) - - await delete_budget_alert_claims(entity_type="key", entity_ids=("hashed_key_1",)) - - assert [r["entity_id"] for r in shared_alert_table.rows] == ["hashed_key_2"] - - @pytest.mark.asyncio async def test_in_memory_claim_does_not_outlive_its_budget_window(shared_alert_table): """One long-lived replica must re-alert after the window moves. The in-memory diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 15af6be3b59..867ef759fb3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -4923,14 +4923,6 @@ async def test_delete_verification_tokens_persists_deleted_keys(monkeypatch): # delete_data returns the list directly, which gets wrapped in {"deleted_keys": ...} assert isinstance(result["deleted_keys"], list) assert set(result["deleted_keys"]) == {"hashed-token-1", "hashed-token-2"} - - # budget alert claim rows are keyed on the token and have no FK, so deleting a key - # has to sweep them or a stale claim silences the alert for a recycled token - claim_delete = mock_prisma_client.db.litellm_budgetalertsent.delete_many - claim_delete.assert_awaited_once() - claim_where = claim_delete.await_args.kwargs["where"] - assert claim_where["entity_type"] == "key" - assert set(claim_where["entity_id"]["in"]) == {"hashed-token-1", "hashed-token-2"} assert len(deleted_keys) == 2 From e01a05472ef0bc52a95ea3c5267f57d02c651775 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 10 Aug 2026 19:54:15 +0000 Subject: [PATCH 3/3] fix: annotate budget alert claim locals and constants as Final Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 4 ++-- litellm/proxy/db/budget_alert_claim.py | 9 +++++---- 2 files changed, 7 insertions(+), 6 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 487384abb68..9031d9e6d79 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -468,8 +468,8 @@ EMAIL_BUDGET_ALERT_TTL: Final = int(os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE: Final = float( os.getenv("EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE", 0.8) ) # 80% of max budget -MAX_BUDGET_ALERT_TYPE = "max_budget_alert" -MAX_BUDGET_DEFAULT_ALERT_TYPE = "max_budget_alert_default_threshold" +MAX_BUDGET_ALERT_TYPE: Final = "max_budget_alert" +MAX_BUDGET_DEFAULT_ALERT_TYPE: Final = "max_budget_alert_default_threshold" ############### LLM Provider Constants ############### ### ANTHROPIC CONSTANTS ### ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01") diff --git a/litellm/proxy/db/budget_alert_claim.py b/litellm/proxy/db/budget_alert_claim.py index 9f808121669..d74c3c3b946 100644 --- a/litellm/proxy/db/budget_alert_claim.py +++ b/litellm/proxy/db/budget_alert_claim.py @@ -17,6 +17,7 @@ from __future__ import annotations from collections.abc import Mapping from datetime import datetime, timezone +from typing import Final from litellm._logging import verbose_proxy_logger from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler @@ -31,7 +32,7 @@ def get_budget_window(budget_reset_at: datetime | None, max_budget: float | None lifetime of the key. Including ``max_budget`` means raising or lowering the budget moves every threshold and re-arms the alerts, which is what an operator expects. """ - period = budget_reset_at.isoformat() if budget_reset_at is not None else "" + period: Final = budget_reset_at.isoformat() if budget_reset_at is not None else "" return f"{period}|{max_budget}" @@ -70,7 +71,7 @@ async def claim_budget_alert_slot( if prisma_client is None: return True - identity = _claim_identity(token, alert_type, threshold_pct) + identity: Final = _claim_identity(token, alert_type, threshold_pct) try: await prisma_client.db.litellm_budgetalertsent.create( @@ -88,7 +89,7 @@ async def claim_budget_alert_slot( return True try: - rows_updated = await prisma_client.db.litellm_budgetalertsent.update_many( + rows_updated: Final = await prisma_client.db.litellm_budgetalertsent.update_many( where={**identity, "budget_window": {"not": budget_window}}, # mutable-ok: prisma query payload data={ # mutable-ok: prisma query payload "budget_window": budget_window, @@ -125,7 +126,7 @@ async def release_budget_alert_slot( if prisma_client is None: return - identity = _claim_identity(token, alert_type, threshold_pct) + identity: Final = _claim_identity(token, alert_type, threshold_pct) try: await prisma_client.db.litellm_budgetalertsent.delete_many(