This commit is contained in:
yuneng-jiang 2025-12-15 18:04:27 -08:00
parent 14feda694b
commit 9907a0d93c
3 changed files with 506 additions and 3 deletions

View file

@ -19,7 +19,7 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import (
)
from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER
from litellm.proxy._types import Litellm_EntityType, WebhookEvent
from litellm.proxy._types import CallInfo, Litellm_EntityType, WebhookEvent
@pytest.fixture(autouse=True)
@ -605,4 +605,238 @@ async def test_get_email_params_default_templates(monkeypatch):
)
assert key_params.subject == "LiteLLM: API Key Created"
assert key_params.signature == EMAIL_FOOTER
assert key_params.signature == EMAIL_FOOTER
@pytest.mark.asyncio
async def test_send_soft_budget_alert_email(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Test that send_soft_budget_alert_email sends an email with the correct parameters and content"""
event = WebhookEvent(
user_id="test_user",
user_email="test@example.com",
event_group=Litellm_EntityType.USER,
event="soft_budget_crossed",
event_message="Soft Budget Crossed - Total Soft Budget: $100.0",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
)
with mock.patch.dict(
os.environ,
{
"EMAIL_LOGO_URL": "https://litellm-listing.s3.amazonaws.com/litellm_logo.png",
"EMAIL_SUPPORT_CONTACT": "support@berri.ai",
"PROXY_BASE_URL": "http://test.com",
},
):
await base_email_logger.send_soft_budget_alert_email(event)
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert call_args["from_email"] == BaseEmailLogger.DEFAULT_LITELLM_EMAIL
assert call_args["to_email"] == ["test@example.com"]
assert call_args["subject"] == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
assert "$100.0" in call_args["html_body"] # soft_budget
assert "$105.0" in call_args["html_body"] # spend
assert "$200.0" in call_args["html_body"] # max_budget
@pytest.mark.asyncio
async def test_send_soft_budget_alert_email_no_max_budget(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Test that send_soft_budget_alert_email handles missing max_budget correctly"""
event = WebhookEvent(
user_id="test_user",
user_email="test@example.com",
event_group=Litellm_EntityType.USER,
event="soft_budget_crossed",
event_message="Soft Budget Crossed - Total Soft Budget: $100.0",
spend=105.0,
max_budget=None,
soft_budget=100.0,
)
with mock.patch.dict(
os.environ,
{
"PROXY_BASE_URL": "http://test.com",
},
):
await base_email_logger.send_soft_budget_alert_email(event)
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert "$100.0" in call_args["html_body"] # soft_budget
assert "$105.0" in call_args["html_body"] # spend
assert "Maximum Budget" not in call_args["html_body"] # max_budget should not be shown
@pytest.mark.asyncio
async def test_budget_alerts_soft_budget_crossed(
base_email_logger, mock_send_email
):
"""Test that budget_alerts sends email when soft budget is crossed"""
user_info = CallInfo(
user_id="test_user",
user_email="test@example.com",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
event_group=Litellm_EntityType.USER,
)
# Mock the cache to return None (no previous alert sent)
mock_cache = mock.AsyncMock()
mock_cache.async_get_cache = mock.AsyncMock(return_value=None)
mock_cache.async_set_cache = mock.AsyncMock()
base_email_logger.internal_usage_cache = mock_cache
with mock.patch.dict(
os.environ,
{
"PROXY_BASE_URL": "http://test.com",
},
):
await base_email_logger.budget_alerts(type="soft_budget", user_info=user_info)
# Verify email was sent
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert call_args["to_email"] == ["test@example.com"]
# Verify cache was set to prevent duplicate alerts
mock_cache.async_set_cache.assert_called_once()
cache_call_args = mock_cache.async_set_cache.call_args[1]
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:test_user"
assert cache_call_args["value"] == "SENT"
assert cache_call_args["ttl"] == BaseEmailLogger.DEFAULT_BUDGET_ALERT_TTL
@pytest.mark.asyncio
async def test_budget_alerts_soft_budget_not_crossed(
base_email_logger, mock_send_email
):
"""Test that budget_alerts does not send email when soft budget is not crossed"""
user_info = CallInfo(
user_id="test_user",
user_email="test@example.com",
spend=50.0,
max_budget=200.0,
soft_budget=100.0,
event_group=Litellm_EntityType.USER,
)
mock_cache = mock.AsyncMock()
base_email_logger.internal_usage_cache = mock_cache
await base_email_logger.budget_alerts(type="soft_budget", user_info=user_info)
# Verify email was NOT sent
mock_send_email.assert_not_called()
@pytest.mark.asyncio
async def test_budget_alerts_soft_budget_duplicate_prevention(
base_email_logger, mock_send_email
):
"""Test that budget_alerts does not send duplicate alerts within TTL period"""
user_info = CallInfo(
user_id="test_user",
user_email="test@example.com",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
event_group=Litellm_EntityType.USER,
)
# Mock the cache to return "SENT" (previous alert already sent)
mock_cache = mock.AsyncMock()
mock_cache.async_get_cache = mock.AsyncMock(return_value="SENT")
base_email_logger.internal_usage_cache = mock_cache
await base_email_logger.budget_alerts(type="soft_budget", user_info=user_info)
# Verify email was NOT sent (duplicate prevention)
mock_send_email.assert_not_called()
@pytest.mark.asyncio
async def test_budget_alerts_no_budgets(
base_email_logger, mock_send_email
):
"""Test that budget_alerts returns early when no budgets are set"""
user_info = CallInfo(
user_id="test_user",
user_email="test@example.com",
spend=50.0,
max_budget=None,
soft_budget=None,
event_group=Litellm_EntityType.USER,
)
await base_email_logger.budget_alerts(type="soft_budget", user_info=user_info)
# Verify email was NOT sent
mock_send_email.assert_not_called()
@pytest.mark.asyncio
async def test_budget_alerts_uses_token_for_cache_key(
base_email_logger, mock_send_email
):
"""Test that budget_alerts uses token for cache key when available"""
user_info = CallInfo(
user_id="test_user",
user_email="test@example.com",
token="hashed_token_123",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
event_group=Litellm_EntityType.KEY,
)
# Mock the cache to return None (no previous alert sent)
mock_cache = mock.AsyncMock()
mock_cache.async_get_cache = mock.AsyncMock(return_value=None)
mock_cache.async_set_cache = mock.AsyncMock()
base_email_logger.internal_usage_cache = mock_cache
with mock.patch.dict(
os.environ,
{
"PROXY_BASE_URL": "http://test.com",
},
):
await base_email_logger.budget_alerts(type="soft_budget", user_info=user_info)
# Verify cache key uses token instead of user_id
mock_cache.async_set_cache.assert_called_once()
cache_call_args = mock_cache.async_set_cache.call_args[1]
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:hashed_token_123"
@pytest.mark.asyncio
async def test_get_email_params_soft_budget_crossed(
base_email_logger, mock_lookup_user_email
):
"""Test that _get_email_params handles soft_budget_crossed event correctly"""
with mock.patch.dict(
os.environ,
{
"PROXY_BASE_URL": "http://test.com",
},
):
result = await base_email_logger._get_email_params(
email_event=EmailEvent.soft_budget_crossed,
user_email="test@example.com",
event_message="Soft Budget Crossed - Total Soft Budget: $100.0",
)
# Should use default subject template for soft_budget_crossed
assert result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
assert result.recipient_email == "test@example.com"
assert result.base_url == "http://test.com"

View file

@ -14,9 +14,11 @@ import pytest
import litellm
from litellm.proxy._types import (
CallInfo,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Litellm_EntityType,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
@ -27,6 +29,7 @@ from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_can_object_call_vector_stores,
_get_team_db_check,
_virtual_key_soft_budget_check,
get_user_object,
vector_store_access_check,
)
@ -988,3 +991,144 @@ async def test_reject_clientside_metadata_tags_non_llm_route():
)
assert result is True
@pytest.mark.asyncio
async def test_virtual_key_soft_budget_check_with_user_obj():
"""Test _virtual_key_soft_budget_check includes user_email when user_obj is provided"""
alert_triggered = False
captured_call_info = None
class MockProxyLogging:
async def budget_alerts(self, type, user_info):
nonlocal alert_triggered, captured_call_info
alert_triggered = True
captured_call_info = user_info
assert type == "soft_budget"
assert isinstance(user_info, CallInfo)
valid_token = UserAPIKeyAuth(
token="test-token",
spend=100.0,
soft_budget=50.0,
user_id="test-user",
team_id="test-team",
team_alias="test-team-alias",
org_id="test-org",
key_alias="test-key",
max_budget=200.0,
)
user_obj = LiteLLM_UserTable(
user_id="test-user",
user_email="test@example.com",
max_budget=None,
)
proxy_logging_obj = MockProxyLogging()
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
await asyncio.sleep(0.1)
assert alert_triggered is True
assert captured_call_info is not None
assert captured_call_info.user_email == "test@example.com"
assert captured_call_info.token == "test-token"
assert captured_call_info.spend == 100.0
assert captured_call_info.soft_budget == 50.0
assert captured_call_info.max_budget == 200.0
assert captured_call_info.user_id == "test-user"
assert captured_call_info.team_id == "test-team"
assert captured_call_info.team_alias == "test-team-alias"
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
@pytest.mark.asyncio
async def test_virtual_key_soft_budget_check_without_user_obj():
"""Test _virtual_key_soft_budget_check sets user_email to None when user_obj is not provided"""
alert_triggered = False
captured_call_info = None
class MockProxyLogging:
async def budget_alerts(self, type, user_info):
nonlocal alert_triggered, captured_call_info
alert_triggered = True
captured_call_info = user_info
assert type == "soft_budget"
assert isinstance(user_info, CallInfo)
valid_token = UserAPIKeyAuth(
token="test-token",
spend=100.0,
soft_budget=50.0,
user_id="test-user",
team_id="test-team",
key_alias="test-key",
)
proxy_logging_obj = MockProxyLogging()
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=None,
)
await asyncio.sleep(0.1)
assert alert_triggered is True
assert captured_call_info is not None
assert captured_call_info.user_email is None
@pytest.mark.parametrize(
"spend, soft_budget, expect_alert",
[
(100.0, 50.0, True), # Over soft budget
(50.0, 50.0, True), # At soft budget
(25.0, 50.0, False), # Under soft budget
(100.0, None, False), # No soft budget set
],
)
@pytest.mark.asyncio
async def test_virtual_key_soft_budget_check_scenarios(
spend, soft_budget, expect_alert
):
"""Test _virtual_key_soft_budget_check with various spend and soft_budget scenarios"""
alert_triggered = False
class MockProxyLogging:
async def budget_alerts(self, type, user_info):
nonlocal alert_triggered
alert_triggered = True
assert type == "soft_budget"
assert isinstance(user_info, CallInfo)
valid_token = UserAPIKeyAuth(
token="test-token",
spend=spend,
soft_budget=soft_budget,
user_id="test-user",
key_alias="test-key",
)
proxy_logging_obj = MockProxyLogging()
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=None,
)
await asyncio.sleep(0.1)
assert (
alert_triggered == expect_alert
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"

View file

@ -1,7 +1,7 @@
import json
import os
import sys
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from jsonschema import validate
@ -2602,3 +2602,128 @@ class TestIsCachedMessage:
"""Empty list content should return False."""
message = {"role": "user", "content": []}
assert is_cached_message(message) is False
@pytest.mark.asyncio
class TestProxyLoggingBudgetAlerts:
"""Test budget_alerts method in ProxyLogging class."""
async def test_budget_alerts_when_alerting_is_none(self):
"""Test that budget_alerts returns early when alerting is None."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
# Should return without calling any alerting instances
await proxy_logging.budget_alerts(type="user_budget", user_info=user_info)
# Verify no calls were made
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
proxy_logging.email_logging_instance.budget_alerts.assert_not_called()
async def test_budget_alerts_with_slack_only(self):
"""Test that budget_alerts calls slack_alerting_instance when slack is in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack"]
proxy_logging.slack_alerting_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="token_budget", user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type="token_budget", user_info=user_info
)
async def test_budget_alerts_with_email_only(self):
"""Test that budget_alerts calls email_logging_instance when email is in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["email"]
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="team_budget", user_info=user_info)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type="team_budget", user_info=user_info
)
async def test_budget_alerts_with_email_when_instance_is_none(self):
"""Test that budget_alerts does not call email_logging_instance when it is None."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["email"]
proxy_logging.email_logging_instance = None
user_info = MagicMock()
# Should not raise an error
await proxy_logging.budget_alerts(type="organization_budget", user_info=user_info)
async def test_budget_alerts_with_both_slack_and_email(self):
"""Test that budget_alerts calls both slack and email instances when both are in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack", "email"]
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="proxy_budget", user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type="proxy_budget", user_info=user_info
)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type="proxy_budget", user_info=user_info
)
@pytest.mark.parametrize(
"alert_type",
[
"token_budget",
"user_budget",
"soft_budget",
"team_budget",
"organization_budget",
"proxy_budget",
"projected_limit_exceeded",
],
)
async def test_budget_alerts_with_all_alert_types(self, alert_type):
"""Test that budget_alerts works with all supported alert types."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack", "email"]
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type=alert_type, user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type=alert_type, user_info=user_info
)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type=alert_type, user_info=user_info
)