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 6b140d489cf..fbf57afeeb4 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 @@ -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 \ No newline at end of file + 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" \ No newline at end of file diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 3d4b68ce441..52f1c433433 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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}" diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 2bd94488ba2..b9485f6a317 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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 + )