fix: end-user budget enforcement was bypassed on master key auth

Budget checks were performed in get_end_user_object() which ran before
model resolution. This had two issues:
1. Master key auth returned early, skipping the budget check entirely
2. Over-budget users were blocked from zero-cost models

Moved budget check after model resolution in user_api_key_auth.py to:
- Enforce budgets for master key auth paths
- Allow over-budget users to access zero-cost models
This commit is contained in:
Fabian Hoeldin 2026-04-08 15:49:56 +02:00
parent 62757ff48f
commit 1ab2814b85
5 changed files with 419 additions and 229 deletions

View file

@ -8,6 +8,7 @@ Run checks for:
2. If user is in budget
3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
"""
import asyncio
import re
import time
@ -196,9 +197,7 @@ def _is_model_cost_zero(
return True
def _is_cost_explicitly_configured(
model: str, llm_router: "Router"
) -> bool:
def _is_cost_explicitly_configured(model: str, llm_router: "Router") -> bool:
"""
Check if any deployment in the model group has cost fields explicitly
set in its litellm.model_cost entry.
@ -215,10 +214,7 @@ def _is_cost_explicitly_configured(
if model_id is None:
continue
raw_entry = litellm.model_cost.get(model_id, {})
if (
"input_cost_per_token" in raw_entry
or "output_cost_per_token" in raw_entry
):
if "input_cost_per_token" in raw_entry or "output_cost_per_token" in raw_entry:
return True
return False
@ -456,9 +452,9 @@ async def common_checks( # noqa: PLR0915
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases
if valid_token
else None,
team_model_aliases=(
valid_token.team_model_aliases if valid_token else None
),
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
@ -925,35 +921,6 @@ async def _apply_default_budget_to_end_user(
return end_user_obj
def _check_end_user_budget(
end_user_obj: LiteLLM_EndUserTable,
route: str,
) -> None:
"""
Check if end user is within their budget limit.
Args:
end_user_obj: The end user object to check
route: The request route
Raises:
litellm.BudgetExceededError: If end user has exceeded their budget
"""
if RouteChecks.is_info_route(route):
return
if end_user_obj.litellm_budget_table is None:
return
end_user_budget = end_user_obj.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_obj.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_obj.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}",
)
@log_db_metrics
async def get_end_user_object(
end_user_id: Optional[str],
@ -969,11 +936,14 @@ async def get_end_user_object(
If end user exists but has no budget_id, applies the default budget
(if configured via litellm.max_end_user_budget_id).
Note: Budget checks are handled in common_checks(), not here.
This allows zero-cost models to skip budget checks properly.
Args:
end_user_id: The ID of the end user
prisma_client: Database client instance
user_api_key_cache: Cache for storing/retrieving data
route: The request route
route: The request route (unused, kept for API compatibility)
parent_otel_span: Optional OpenTelemetry span for tracing
proxy_logging_obj: Optional proxy logging object
@ -1001,9 +971,6 @@ async def get_end_user_object(
parent_otel_span=parent_otel_span,
)
# Check budget limits
_check_end_user_budget(end_user_obj=return_obj, route=route)
return return_obj
# Fetch from database
@ -1032,9 +999,6 @@ async def get_end_user_object(
key="end_user_id:{}".format(end_user_id), value=_response.dict()
)
# Check budget limits
_check_end_user_budget(end_user_obj=_response, route=route)
return _response
except Exception as e:

View file

@ -247,6 +247,50 @@ def _apply_budget_limits_to_end_user_params(
verbose_proxy_logger.debug(f"Applied budget limits to end user {end_user_id}")
async def _check_end_user_budget(
end_user_object: Optional[Any],
request_data: dict,
route: str,
llm_router: Optional[Any],
) -> None:
"""
Check if end user has exceeded their budget. Raises BudgetExceededError if so.
Skips check for zero-cost models.
Args:
end_user_object: The end user object from DB
request_data: Request body data
route: Current route
llm_router: LLM router instance
Raises:
BudgetExceededError: If end user has exceeded their max_budget
"""
if end_user_object is None:
return
# Check if model has zero cost - if so, skip budget check
model = get_model_from_request(request_data, route)
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
if _is_model_cost_zero(model=model, llm_router=llm_router):
verbose_proxy_logger.info(
f"Skipping all budget checks for zero-cost model: {model}"
)
return
# Check end-user budget
if end_user_object.litellm_budget_table is not None:
end_user_budget = end_user_object.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_object.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_object.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
async def user_api_key_auth_websocket(websocket: WebSocket):
# Accept the WebSocket connection
@ -314,6 +358,8 @@ def update_valid_token_with_end_user_params(
valid_token.end_user_rpm_limit = end_user_params["end_user_rpm_limit"]
if end_user_params.get("allowed_model_region") is not None:
valid_token.allowed_model_region = end_user_params["allowed_model_region"]
if end_user_params.get("end_user_max_budget") is not None:
valid_token.end_user_max_budget = end_user_params["end_user_max_budget"]
if end_user_params.get("end_user_model_max_budget") is not None:
valid_token.end_user_model_max_budget = end_user_params[
"end_user_model_max_budget"
@ -696,11 +742,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Routing uses unverified JWT claims only to choose auth path.
# Final authentication is enforced by the selected validator.
route_jwt_to_oauth2 = (
is_jwt
and _should_route_jwt_to_oauth2_override(
token=api_key, jwt_handler=jwt_handler
)
route_jwt_to_oauth2 = is_jwt and _should_route_jwt_to_oauth2_override(
token=api_key, jwt_handler=jwt_handler
)
# OAuth2 applies for:
@ -716,6 +759,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if (should_apply_global_oauth2 and not is_jwt) or should_apply_override_oauth2:
from litellm.proxy.proxy_server import premium_user
if premium_user is not True:
raise ValueError(
"Oauth2 token validation is only available for premium users"
@ -743,10 +787,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None:
# Decode JWT to get claims without running full auth_builder
jwt_claims: Optional[dict]
if (
jwt_handler.litellm_jwtauth.oidc_userinfo_enabled
and not is_jwt
):
if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt:
jwt_claims = await jwt_handler.get_oidc_userinfo(token=api_key)
else:
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
@ -981,9 +1022,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
route=route,
)
if _end_user_object is not None:
end_user_params[
"allowed_model_region"
] = _end_user_object.allowed_model_region
end_user_params["allowed_model_region"] = (
_end_user_object.allowed_model_region
)
if _end_user_object.litellm_budget_table is not None:
_apply_budget_limits_to_end_user_params(
end_user_params=end_user_params,
@ -1078,6 +1119,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
_end_user_object.object_permission
)
# Check end-user budget before returning
await _check_end_user_budget(
end_user_object=_end_user_object,
request_data=request_data,
route=route,
llm_router=llm_router,
)
return valid_token
if (
@ -1156,6 +1205,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
valid_token=_user_api_key_obj, end_user_params=end_user_params
)
# Check end-user budget before returning
await _check_end_user_budget(
end_user_object=_end_user_object,
request_data=request_data,
route=route,
llm_router=llm_router,
)
return _user_api_key_obj
## IF it's not a master key
@ -1537,9 +1594,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if _end_user_object is not None:
valid_token_dict.update(end_user_params)
valid_token_dict[
"end_user_object_permission"
] = _end_user_object.object_permission
valid_token_dict["end_user_object_permission"] = (
_end_user_object.object_permission
)
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
# sso/login, ui/login, /key functions and /user functions
@ -1802,12 +1859,9 @@ async def _enforce_key_and_fallback_model_access(
if config != {}:
model_list = config.get("model_list", [])
new_model_list = model_list
verbose_proxy_logger.debug(
f"\n new llm router model list {new_model_list}"
)
verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}")
elif (
isinstance(valid_token.models, list)
and "all-team-models" in valid_token.models
isinstance(valid_token.models, list) and "all-team-models" in valid_token.models
):
pass
else:

View file

@ -8,9 +8,7 @@ from dotenv import load_dotenv
load_dotenv()
import os
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
import pytest, litellm
import httpx
from litellm.proxy._types import UserAPIKeyAuth
@ -38,7 +36,10 @@ from litellm.proxy.utils import CallInfo
async def test_get_end_user_object(customer_spend, customer_budget):
"""
Scenario 1: normal
Scenario 2: user over budget
Scenario 2: user over budget - budget check moved to user_api_key_auth.py after model resolution
get_end_user_object() should return the user object regardless of budget status.
Budget enforcement happens later in user_api_key_auth.py to allow zero-cost model access.
"""
end_user_id = "my-test-customer"
_budget = LiteLLM_BudgetTable(max_budget=customer_budget)
@ -51,31 +52,15 @@ async def test_get_end_user_object(customer_spend, customer_budget):
_cache = DualCache()
_key = "end_user_id:{}".format(end_user_id)
_cache.set_cache(key=_key, value=end_user_obj.model_dump())
try:
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",
)
if customer_spend > customer_budget:
pytest.fail(
"Expected call to fail. 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:
pytest.fail(
"Expected call to work. Customer Spend={}, Customer Budget={}, Error={}".format(
customer_spend, customer_budget, str(e)
)
)
result = await get_end_user_object(
end_user_id=end_user_id,
prisma_client="RANDOM VALUE",
user_api_key_cache=_cache,
route="/v1/chat/completions",
)
assert result is not None
assert result.user_id == end_user_id
assert result.spend == customer_spend
@pytest.mark.parametrize(
@ -402,16 +387,12 @@ async def test_is_valid_fallback_model():
)
try:
await is_valid_fallback_model(
model="gpt-3.5-turbo", llm_router=router, user_model=None
)
await is_valid_fallback_model(model="gpt-3.5-turbo", llm_router=router, user_model=None)
except Exception as e:
pytest.fail(f"Expected is_valid_fallback_model to work, got exception: {e}")
try:
await is_valid_fallback_model(
model="gpt-4o", llm_router=router, user_model=None
)
await is_valid_fallback_model(model="gpt-4o", llm_router=router, user_model=None)
pytest.fail("Expected is_valid_fallback_model to fail")
except Exception as e:
assert "Invalid" in str(e)
@ -426,9 +407,7 @@ async def test_is_valid_fallback_model():
],
)
@pytest.mark.asyncio
async def test_virtual_key_max_budget_check(
token_spend, max_budget, expect_budget_error
):
async def test_virtual_key_max_budget_check(token_spend, max_budget, expect_budget_error):
"""
Test if virtual key budget checks work as expected:
1. Triggers budget alert for all cases
@ -472,14 +451,10 @@ async def test_virtual_key_max_budget_check(
user_obj=user_obj,
)
if expect_budget_error:
pytest.fail(
f"Expected BudgetExceededError for spend={token_spend}, max_budget={max_budget}"
)
pytest.fail(f"Expected BudgetExceededError for spend={token_spend}, max_budget={max_budget}")
except litellm.BudgetExceededError as e:
if not expect_budget_error:
pytest.fail(
f"Unexpected BudgetExceededError for spend={token_spend}, max_budget={max_budget}"
)
pytest.fail(f"Unexpected BudgetExceededError for spend={token_spend}, max_budget={max_budget}")
assert e.current_cost == token_spend
assert e.max_budget == max_budget
@ -546,9 +521,7 @@ async def test_can_team_access_model(model, team_models, expect_to_work):
team_model_aliases=None,
)
if not expect_to_work:
pytest.fail(
f"Expected model access check to fail for model={model}, team_models={team_models}"
)
pytest.fail(f"Expected model access check to fail for model={model}, team_models={team_models}")
except Exception as e:
if expect_to_work:
pytest.fail(
@ -601,9 +574,9 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert):
await asyncio.sleep(0.1) # Allow time for the alert task to complete
assert (
alert_triggered == expect_alert
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
assert alert_triggered == expect_alert, (
f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
)
@pytest.mark.parametrize(
@ -613,9 +586,27 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert):
(50, 50, False, None, None), # At soft budget, no metadata - no alert_emails configured, so no alert
(25, 50, False, None, None), # Under soft budget
(100, None, False, None, None), # No soft budget set
(100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with list of emails
(100, 50, True, {"soft_budget_alerting_emails": "team1@example.com,team2@example.com"}, ["team1@example.com", "team2@example.com"]), # Over soft budget with comma-separated emails
(100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "", " ", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with empty strings filtered
(
100,
50,
True,
{"soft_budget_alerting_emails": ["team1@example.com", "team2@example.com"]},
["team1@example.com", "team2@example.com"],
), # Over soft budget with list of emails
(
100,
50,
True,
{"soft_budget_alerting_emails": "team1@example.com,team2@example.com"},
["team1@example.com", "team2@example.com"],
), # Over soft budget with comma-separated emails
(
100,
50,
True,
{"soft_budget_alerting_emails": ["team1@example.com", "", " ", "team2@example.com"]},
["team1@example.com", "team2@example.com"],
), # Over soft budget with empty strings filtered
],
)
@pytest.mark.asyncio
@ -667,9 +658,9 @@ async def test_team_soft_budget_check(spend, soft_budget, expect_alert, metadata
await asyncio.sleep(0.1) # Allow time for the alert task to complete
assert (
alert_triggered == expect_alert
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
assert alert_triggered == expect_alert, (
f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
)
if expect_alert:
assert captured_call_info is not None
@ -813,9 +804,7 @@ async def test_get_fuzzy_user_object():
# Test 5: Only email provided (no SSO ID)
mock_prisma.db.litellm_usertable.find_first = AsyncMock(return_value=test_user)
result = await _get_fuzzy_user_object(
prisma_client=mock_prisma, user_email="test@example.com"
)
result = await _get_fuzzy_user_object(prisma_client=mock_prisma, user_email="test@example.com")
assert result == test_user
mock_prisma.db.litellm_usertable.find_first.assert_called_with(
where={"user_email": {"equals": "test@example.com", "mode": "insensitive"}},
@ -824,9 +813,7 @@ async def test_get_fuzzy_user_object():
# Test 6: Only SSO ID provided (no email)
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=test_user)
result = await _get_fuzzy_user_object(
prisma_client=mock_prisma, sso_user_id="sso_123"
)
result = await _get_fuzzy_user_object(prisma_client=mock_prisma, sso_user_id="sso_123")
assert result == test_user
mock_prisma.db.litellm_usertable.find_unique.assert_called_with(
where={"sso_user_id": "sso_123"}, include={"organization_memberships": True}

View file

@ -8,14 +8,14 @@ a default budget to end users without explicit budgets.
import sys
import os
import uuid
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_EndUserTable
from litellm.proxy._types import LiteLLM_BudgetTable
from litellm.proxy.auth.auth_checks import get_end_user_object
from litellm.caching import DualCache
@ -29,14 +29,14 @@ async def test_default_budget_applied_to_end_user_without_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
default_budget = LiteLLM_BudgetTable(
budget_id=default_budget_id,
max_budget=10.0,
rpm_limit=2,
tpm_limit=10,
)
# Mock end user in DB without budget
mock_end_user_data = {
"user_id": end_user_id,
@ -47,7 +47,7 @@ async def test_default_budget_applied_to_end_user_without_budget():
"default_model": None,
"blocked": False,
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: mock_end_user_data)
@ -55,18 +55,18 @@ async def test_default_budget_applied_to_end_user_without_budget():
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: default_budget.dict())
)
mock_cache = AsyncMock(spec=DualCache)
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
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 default budget was applied
assert result is not None
assert result.litellm_budget_table is not None
@ -74,7 +74,7 @@ async def test_default_budget_applied_to_end_user_without_budget():
assert result.litellm_budget_table.max_budget == 10.0
assert result.litellm_budget_table.rpm_limit == 2
assert result.litellm_budget_table.tpm_limit == 10
litellm.max_end_user_budget_id = None
@ -88,13 +88,13 @@ async def test_explicit_budget_not_overridden_by_default():
explicit_budget_id = str(uuid.uuid4())
default_budget_id = str(uuid.uuid4())
litellm.max_end_user_budget_id = default_budget_id
explicit_budget = LiteLLM_BudgetTable(
budget_id=explicit_budget_id,
max_budget=100.0,
rpm_limit=50,
)
# Mock end user with explicit budget
mock_end_user_data = {
"user_id": end_user_id,
@ -105,48 +105,49 @@ async def test_explicit_budget_not_overridden_by_default():
"default_model": None,
"blocked": False,
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: mock_end_user_data)
)
mock_cache = AsyncMock(spec=DualCache)
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
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 explicit budget is kept (not replaced with default)
assert result is not None
assert result.litellm_budget_table.budget_id == explicit_budget_id
assert result.litellm_budget_table.max_budget == 100.0
assert result.litellm_budget_table.rpm_limit == 50
litellm.max_end_user_budget_id = None
@pytest.mark.asyncio
async def test_budget_enforcement_blocks_over_budget_users():
async def test_over_budget_user_gets_budget_applied():
"""
Core scenario: Budget limits are actually enforced.
Users who exceed their budget should be blocked.
Budget enforcement moved to user_api_key_auth.py after model resolution.
get_end_user_object() should still apply the budget, but NOT raise error.
This allows zero-cost models to bypass budget checks.
"""
end_user_id = f"test_user_{uuid.uuid4().hex}"
default_budget_id = str(uuid.uuid4())
litellm.max_end_user_budget_id = default_budget_id
default_budget = LiteLLM_BudgetTable(
budget_id=default_budget_id,
max_budget=10.0,
rpm_limit=2,
)
# Mock end user who has already spent more than budget
mock_end_user_data = {
"user_id": end_user_id,
@ -157,7 +158,7 @@ async def test_budget_enforcement_blocks_over_budget_users():
"default_model": None,
"blocked": False,
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: mock_end_user_data)
@ -165,23 +166,25 @@ async def test_budget_enforcement_blocks_over_budget_users():
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: default_budget.dict())
)
mock_cache = AsyncMock(spec=DualCache)
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
# Should raise BudgetExceededError
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,
route="/chat/completions",
)
assert "ExceededBudget" in str(exc_info.value)
assert end_user_id in str(exc_info.value)
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 default budget was applied (even though user is over budget)
# Budget check happens later in user_api_key_auth.py after model resolution
assert result is not None
assert result.litellm_budget_table is not None
assert result.litellm_budget_table.max_budget == 10.0
assert result.spend == 15.0 # Over budget, but no error raised here
litellm.max_end_user_budget_id = None
@ -193,7 +196,7 @@ async def test_system_works_without_default_budget_configured():
"""
end_user_id = f"test_user_{uuid.uuid4().hex}"
litellm.max_end_user_budget_id = None # Not configured
# Mock end user without budget
mock_end_user_data = {
"user_id": end_user_id,
@ -204,25 +207,24 @@ async def test_system_works_without_default_budget_configured():
"default_model": None,
"blocked": False,
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(
return_value=MagicMock(dict=lambda: mock_end_user_data)
)
mock_cache = AsyncMock(spec=DualCache)
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
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",
)
# Should work fine, just without budget limits
assert result is not None
assert result.user_id == end_user_id
assert result.litellm_budget_table is None # No budget applied

View file

@ -7,9 +7,7 @@ import sys
import litellm.proxy
import litellm.proxy.proxy_server
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
from typing import Dict, List, Optional
from unittest.mock import MagicMock, patch, AsyncMock
@ -50,9 +48,7 @@ class Request:
), # Request with no client IP should not be allowed
],
)
def test_check_valid_ip(
allowed_ips: Optional[List[str]], client_ip: Optional[str], expected_result: bool
):
def test_check_valid_ip(allowed_ips: Optional[List[str]], client_ip: Optional[str], expected_result: bool):
from litellm.proxy.auth.auth_utils import _check_valid_ip
request = Request(client_ip)
@ -121,9 +117,7 @@ async def test_check_blocked_team():
last_refreshed_at=time.time(),
)
await asyncio.sleep(1)
team_obj = LiteLLM_TeamTableCachedObj(
team_id=_team_id, blocked=False, last_refreshed_at=time.time()
)
team_obj = LiteLLM_TeamTableCachedObj(team_id=_team_id, blocked=False, last_refreshed_at=time.time())
hashed_token = hash_token(user_key)
print(f"STORING TOKEN UNDER KEY={hashed_token}")
user_api_key_cache.set_cache(key=hashed_token, value=valid_token)
@ -173,9 +167,7 @@ async def test_team_object_has_object_permission_id():
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
) as mock_common_checks:
with patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock) as mock_common_checks:
mock_common_checks.return_value = True
await user_api_key_auth(request=request, api_key="Bearer " + user_key)
@ -200,9 +192,7 @@ async def test_returned_user_api_key_auth(user_role, expected_role):
from datetime import datetime
new_obj = await _return_user_api_key_auth_obj(
user_obj=LiteLLM_UserTable(
user_role=user_role, user_id="", max_budget=None, user_email=""
),
user_obj=LiteLLM_UserTable(user_role=user_role, user_id="", max_budget=None, user_email=""),
api_key="hello-world",
parent_otel_span=None,
valid_token_dict={},
@ -253,9 +243,7 @@ async def test_aaauser_personal_budgets(key_ownership):
spend=20,
)
user_obj = LiteLLM_UserTable(
user_id=_user_id, spend=11, max_budget=10, user_email=""
)
user_obj = LiteLLM_UserTable(user_id=_user_id, spend=11, max_budget=10, user_email="")
user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token)
user_api_key_cache.set_cache(key="{}".format(_user_id), value=user_obj)
@ -305,9 +293,7 @@ async def test_user_api_key_auth_fails_with_prohibited_params(prohibited_param):
request.body = return_body
try:
response = await user_api_key_auth(
request=request, api_key="Bearer " + user_key
)
response = await user_api_key_auth(request=request, api_key="Bearer " + user_key)
except Exception as e:
print("error str=", str(e))
error_message = str(e.message)
@ -502,9 +488,7 @@ def test_get_api_key_from_custom_header(headers, custom_header_name, expected_ap
# Call the function and verify it doesn't raise an exception
api_key = get_api_key_from_custom_header(
request=request, custom_litellm_key_header_name=custom_header_name
)
api_key = get_api_key_from_custom_header(request=request, custom_litellm_key_header_name=custom_header_name)
assert api_key == expected_api_key
@ -521,9 +505,7 @@ from litellm.proxy._types import LitellmUserRoles
(LitellmUserRoles.TEAM, "1234", "1234", True),
],
)
def test_allowed_route_inside_route(
user_role, auth_user_id, requested_user_id, expected_result
):
def test_allowed_route_inside_route(user_role, auth_user_id, requested_user_id, expected_result):
from litellm.proxy.auth.auth_checks import allowed_route_check_inside_route
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles
@ -664,9 +646,7 @@ async def test_soft_budget_alert():
try:
# Call user_api_key_auth
response = await user_api_key_auth(
request=request, api_key="Bearer " + user_key
)
response = await user_api_key_auth(request=request, api_key="Bearer " + user_key)
# Assert the request was allowed (no exception raised)
assert response is not None
@ -833,10 +813,7 @@ async def test_user_api_key_auth_websocket():
mock_websocket.url = URL(url="/ws")
# Mock the return value of `user_api_key_auth` when it's called within the `user_api_key_auth_websocket` function
with patch(
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True
) as mock_user_api_key_auth:
with patch("litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True) as mock_user_api_key_auth:
# Make the call to the WebSocket function
await user_api_key_auth_websocket(mock_websocket)
@ -845,15 +822,13 @@ async def test_user_api_key_auth_websocket():
# Get the request object that was passed to user_api_key_auth
request_arg = mock_user_api_key_auth.call_args.kwargs["request"]
# Verify that the request has headers set
assert hasattr(request_arg, "headers"), "Request object should have headers attribute"
assert "authorization" in request_arg.headers, "Request headers should contain authorization"
assert request_arg.headers["authorization"] == "Bearer some_api_key"
assert (
mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key"
)
assert mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key"
@pytest.mark.parametrize("enforce_rbac", [True, False])
@ -1035,14 +1010,10 @@ async def test_jwt_non_admin_team_route_access(monkeypatch):
}
# Create request
request = Request(
scope={"type": "http", "headers": [(b"authorization", b"Bearer fake.jwt.token")]}
)
request = Request(scope={"type": "http", "headers": [(b"authorization", b"Bearer fake.jwt.token")]})
request._url = URL(url="/team/new")
monkeypatch.setattr(
litellm.proxy.proxy_server, "general_settings", {"enable_jwt_auth": True}
)
monkeypatch.setattr(litellm.proxy.proxy_server, "general_settings", {"enable_jwt_auth": True})
# Initialize jwt_handler with a default LiteLLM_JWTAuth so that the
# virtual_key_claim_field check in user_api_key_auth doesn't fail with
@ -1059,18 +1030,19 @@ async def test_jwt_non_admin_team_route_access(monkeypatch):
# Mock enterprise license check and JWTAuthManager.auth_builder
# License check must be mocked to avoid environment variable pollution
# in parallel test execution
with patch(
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
"litellm.proxy.auth.handle_jwt.JWTAuthManager.auth_builder",
return_value=mock_jwt_response,
with (
patch(
"litellm.proxy.proxy_server.premium_user",
True,
),
patch(
"litellm.proxy.auth.handle_jwt.JWTAuthManager.auth_builder",
return_value=mock_jwt_response,
),
):
try:
await user_api_key_auth(request=request, api_key="Bearer fake.jwt.token")
pytest.fail(
"Expected this call to fail. Non-admin user should not access team routes."
)
pytest.fail("Expected this call to fail. Non-admin user should not access team routes.")
except ProxyException as e:
print("e", e)
assert "Only proxy admin can be used to generate" in str(e.message)
@ -1101,14 +1073,12 @@ async def test_x_litellm_api_key():
ignored_key = "aj12445"
# Create request with headers as bytes
request = Request(
scope={
"type": "http"
}
)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
valid_token = await user_api_key_auth(request=request, api_key="Bearer " + ignored_key, custom_litellm_key_header=master_key)
valid_token = await user_api_key_auth(
request=request, api_key="Bearer " + ignored_key, custom_litellm_key_header=master_key
)
assert valid_token.token == hash_token(master_key)
@ -1146,3 +1116,216 @@ async def test_user_api_key_from_query_param():
valid_token = await user_api_key_auth(request=request, api_key="")
assert valid_token.token == hash_token(user_key)
class TestEndUserBudgetWithVirtualKey:
"""Tests for end user budget enforcement with virtual key authentication."""
@pytest.mark.asyncio
async def test_over_budget_end_user_can_access_zero_cost_model_with_virtual_key(self):
"""
Test that an over-budget end user can still access zero-cost models
when authenticating with a virtual key.
This verifies that budget checks properly skip for zero-cost models
regardless of authentication method.
"""
import time
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
from litellm.router import Router
master_key = "sk-master-1234"
user_key = "sk-virtual-key-1234"
hashed_key = hash_token(user_key)
end_user_id = "over-budget-end-user"
# Create a valid virtual key token
valid_token = UserAPIKeyAuth(
token=hashed_key,
user_id="test-user",
last_refreshed_at=time.time(),
)
user_api_key_cache.set_cache(key=hashed_key, value=valid_token)
# Create an over-budget end user
end_user_budget = LiteLLM_BudgetTable(
budget_id="test-budget",
max_budget=10.0,
)
over_budget_end_user = LiteLLM_EndUserTable(
user_id=end_user_id,
spend=50.0, # Over budget of 10.0
litellm_budget_table=end_user_budget,
blocked=False,
)
# Create a router with zero-cost model
mock_router = Router(
model_list=[
{
"model_name": "free-model",
"litellm_params": {
"model": "ollama/llama2",
"api_base": "http://localhost:11434",
},
"model_info": {
"id": "free-model-id",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
},
},
]
)
# Set up the proxy server state
setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "prisma_client", "mock-prisma")
setattr(litellm.proxy.proxy_server, "llm_router", mock_router)
# Mock get_end_user_object to return our over-budget user
with patch(
"litellm.proxy.auth.user_api_key_auth.get_end_user_object",
new_callable=AsyncMock,
return_value=over_budget_end_user,
):
request = Request(
scope={
"type": "http",
"headers": [(b"x-litellm-end-user-id", end_user_id.encode())],
}
)
request._url = URL(url="/v1/chat/completions")
async def return_body():
return b'{"model": "free-model", "user": "' + end_user_id.encode() + b'"}'
request.body = return_body
# Should NOT raise BudgetExceededError for zero-cost model
result = await user_api_key_auth(
request=request,
api_key="Bearer " + user_key,
)
assert result is not None
# Cleanup
setattr(litellm.proxy.proxy_server, "llm_router", None)
@pytest.mark.asyncio
async def test_over_budget_end_user_blocked_from_paid_model_with_virtual_key(self):
"""
Test that an over-budget end user is blocked from paid models
when authenticating with a virtual key.
This verifies the budget enforcement works correctly with virtual keys.
"""
import time
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
from litellm.router import Router
master_key = "sk-master-1234"
user_key = "sk-virtual-key-5678"
hashed_key = hash_token(user_key)
end_user_id = "over-budget-end-user-2"
# Create a valid virtual key token
valid_token = UserAPIKeyAuth(
token=hashed_key,
user_id="test-user-2",
last_refreshed_at=time.time(),
)
user_api_key_cache.set_cache(key=hashed_key, value=valid_token)
# Create an over-budget end user
end_user_budget = LiteLLM_BudgetTable(
budget_id="test-budget-2",
max_budget=10.0,
)
over_budget_end_user = LiteLLM_EndUserTable(
user_id=end_user_id,
spend=50.0, # Over budget of 10.0
litellm_budget_table=end_user_budget,
blocked=False,
)
# Create a router with a paid model
mock_router = Router(
model_list=[
{
"model_name": "paid-model",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "sk-test",
},
"model_info": {
"id": "paid-model-id",
},
},
]
)
# Set up the proxy server state
setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "prisma_client", "mock-prisma")
setattr(litellm.proxy.proxy_server, "llm_router", mock_router)
# Mock get_end_user_object and get_model_info
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_end_user_object",
new_callable=AsyncMock,
return_value=over_budget_end_user,
),
patch("litellm.get_model_info") as mock_get_model_info,
):
# Mock paid model cost
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
request = Request(
scope={
"type": "http",
"headers": [(b"x-litellm-end-user-id", end_user_id.encode())],
}
)
request._url = URL(url="/v1/chat/completions")
async def return_body():
return b'{"model": "paid-model", "user": "' + end_user_id.encode() + b'"}'
request.body = return_body
# Should raise ProxyException for paid model (budget exceeded)
from litellm.proxy._types import ProxyException
with pytest.raises(ProxyException) as exc_info:
await user_api_key_auth(
request=request,
api_key="Bearer " + user_key,
)
assert "ExceededBudget" in exc_info.value.message
assert end_user_id in exc_info.value.message
# Cleanup
setattr(litellm.proxy.proxy_server, "llm_router", None)