mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
refactor(tests): update budget enforcement tests to reflect changes in get_end_user_object
- test_get_end_user_object() verifies data fetching - test_check_end_user_budget() verifies enforcement - test_budget_enforcement_blocks_over_budget_users() integrates _check_end_user_budget() - test_resolve_end_user_reraises_budget_exceeded() is now test_resolve_end_user since no budget exceeded is thrown in get_end_user_object()
This commit is contained in:
parent
49c7604ed8
commit
b95cf4e6fe
3 changed files with 100 additions and 35 deletions
|
|
@ -38,8 +38,12 @@ from litellm.proxy.utils import CallInfo
|
|||
@pytest.mark.asyncio
|
||||
async def test_get_end_user_object(customer_spend, customer_budget):
|
||||
"""
|
||||
Scenario 1: normal
|
||||
Scenario 2: user over budget
|
||||
Scenario 1: normal - get_end_user_object returns the cached user
|
||||
Scenario 2: user over budget - NOTE: budget enforcement now happens in
|
||||
common_checks() via _check_end_user_budget(), not in get_end_user_object()
|
||||
|
||||
This test verifies that get_end_user_object correctly retrieves the end user
|
||||
from cache. Budget enforcement is tested separately in test_check_end_user_budget().
|
||||
"""
|
||||
end_user_id = "my-test-customer"
|
||||
_budget = LiteLLM_BudgetTable(max_budget=customer_budget)
|
||||
|
|
@ -58,31 +62,62 @@ async def test_get_end_user_object(customer_spend, customer_budget):
|
|||
value=end_user_obj,
|
||||
model_type=LiteLLM_EndUserTable,
|
||||
)
|
||||
# get_end_user_object only fetches data - it no longer enforces budget
|
||||
# Budget enforcement happens in common_checks() via _check_end_user_budget()
|
||||
result = await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client="RANDOM VALUE", # type: ignore
|
||||
user_api_key_cache=_cache,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.user_id == end_user_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize("customer_spend, customer_budget", [(0, 10), (10, 0)])
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_end_user_budget(customer_spend, customer_budget):
|
||||
"""
|
||||
Test _check_end_user_budget enforcement:
|
||||
- Scenario 1: customer_spend=0, customer_budget=10 - should pass (under budget)
|
||||
- Scenario 2: customer_spend=10, customer_budget=0 - should fail (over budget)
|
||||
|
||||
Note: Budget enforcement for end users happens in common_checks() via
|
||||
_check_end_user_budget(), not in get_end_user_object().
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _check_end_user_budget
|
||||
|
||||
_budget = LiteLLM_BudgetTable(max_budget=customer_budget)
|
||||
end_user_obj = LiteLLM_EndUserTable(
|
||||
user_id="my-test-customer",
|
||||
spend=customer_spend,
|
||||
litellm_budget_table=_budget,
|
||||
blocked=False,
|
||||
)
|
||||
|
||||
should_exceed = customer_spend > customer_budget
|
||||
|
||||
try:
|
||||
await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client="RANDOM VALUE", # type: ignore
|
||||
user_api_key_cache=_cache,
|
||||
await _check_end_user_budget(
|
||||
end_user_obj=end_user_obj,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
if customer_spend > customer_budget:
|
||||
if should_exceed:
|
||||
pytest.fail(
|
||||
"Expected call to fail. Customer Spend={}, Customer Budget={}".format(
|
||||
"Expected BudgetExceededError. Customer Spend={}, Customer Budget={}".format(
|
||||
customer_spend, customer_budget
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
if (
|
||||
isinstance(e, litellm.BudgetExceededError)
|
||||
and customer_spend > customer_budget
|
||||
):
|
||||
pass
|
||||
else:
|
||||
except litellm.BudgetExceededError as e:
|
||||
if not should_exceed:
|
||||
pytest.fail(
|
||||
"Expected call to work. Customer Spend={}, Customer Budget={}, Error={}".format(
|
||||
"Unexpected BudgetExceededError. Customer Spend={}, Customer Budget={}, Error={}".format(
|
||||
customer_spend, customer_budget, str(e)
|
||||
)
|
||||
)
|
||||
# Verify the error has correct info
|
||||
assert e.current_cost == customer_spend
|
||||
assert e.max_budget == customer_budget
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -134,9 +134,14 @@ async def test_explicit_budget_not_overridden_by_default():
|
|||
@pytest.mark.asyncio
|
||||
async def test_budget_enforcement_blocks_over_budget_users():
|
||||
"""
|
||||
Core scenario: Budget limits are actually enforced.
|
||||
Core scenario: Budget limits are actually enforced via _check_end_user_budget.
|
||||
Users who exceed their budget should be blocked.
|
||||
|
||||
Note: Budget enforcement happens in common_checks() via _check_end_user_budget(),
|
||||
not in get_end_user_object(). get_end_user_object only fetches the user data.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _check_end_user_budget
|
||||
|
||||
end_user_id = f"test_user_{uuid.uuid4().hex}"
|
||||
default_budget_id = str(uuid.uuid4())
|
||||
litellm.max_end_user_budget_id = default_budget_id
|
||||
|
|
@ -170,12 +175,23 @@ async def test_budget_enforcement_blocks_over_budget_users():
|
|||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
# Should raise BudgetExceededError
|
||||
# First, get the end user object (this just fetches data, doesn't enforce budget)
|
||||
result = await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
# Verify user was fetched with default budget applied
|
||||
assert result is not None
|
||||
assert result.litellm_budget_table is not None
|
||||
assert result.litellm_budget_table.max_budget == 10.0
|
||||
|
||||
# Now test budget enforcement separately via _check_end_user_budget
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
await _check_end_user_budget(
|
||||
end_user_obj=result,
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -3390,30 +3390,44 @@ async def test_resolve_end_user_swallows_db_errors_and_returns_none(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_end_user_reraises_budget_exceeded(
|
||||
async def test_resolve_end_user(
|
||||
_validate_flag_on, monkeypatch
|
||||
):
|
||||
"""BudgetExceededError from get_end_user_object must bubble up so the
|
||||
auth path enforces spend limits instead of silently dropping the id."""
|
||||
import litellm
|
||||
"""Verify that resolve_and_validate_end_user_id does NOT raise BudgetExceededError.
|
||||
|
||||
Note: As of the refactor that moved _check_end_user_budget out of
|
||||
get_end_user_object, budget enforcement now happens in common_checks().
|
||||
|
||||
The end-user validation path should return the user ID regardless of budget status.
|
||||
Budget enforcement for end users happens later in common_checks() via
|
||||
_check_end_user_budget(), which respects skip_budget_checks for zero-cost models.
|
||||
|
||||
This test verifies that even when get_end_user_object returns a user with a budget,
|
||||
resolve_and_validate_end_user_id does not block the request - budget enforcement
|
||||
is deferred to common_checks() where skip_budget_checks logic can be applied.
|
||||
"""
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
||||
|
||||
# Mock get_end_user_object to return a user with budget info
|
||||
# (simulating a user who may have exceeded their budget)
|
||||
mock_end_user = MagicMock()
|
||||
mock_end_user.user_id = "customer-over-budget"
|
||||
monkeypatch.setattr(
|
||||
auth_checks,
|
||||
"get_end_user_object",
|
||||
AsyncMock(
|
||||
side_effect=litellm.BudgetExceededError(current_cost=10.0, max_budget=5.0)
|
||||
),
|
||||
AsyncMock(return_value=mock_end_user),
|
||||
)
|
||||
cache = _validation_cache()
|
||||
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await resolve_and_validate_end_user_id(
|
||||
raw_end_user_id="customer-over-budget",
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
# resolve_and_validate_end_user_id should return the user ID without raising
|
||||
# BudgetExceededError - budget enforcement happens in common_checks()
|
||||
result = await resolve_and_validate_end_user_id(
|
||||
raw_end_user_id="customer-over-budget",
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
assert result == "customer-over-budget"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue