mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
parent
62757ff48f
commit
1ab2814b85
5 changed files with 419 additions and 229 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue