mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Tests
This commit is contained in:
parent
14feda694b
commit
9907a0d93c
3 changed files with 506 additions and 3 deletions
|
|
@ -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"
|
||||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue