diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 20ab9904f46..22692a88b62 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -604,6 +604,8 @@ def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_use 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"] return valid_token diff --git a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py index cf1f665ad21..232036ff328 100644 --- a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py +++ b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py @@ -277,3 +277,75 @@ def test_update_valid_token_db_values_override_custom_auth_when_set(): # DB values should win assert result.end_user_tpm_limit == 500 assert result.end_user_model_max_budget == db_budget + + +def test_update_valid_token_copies_end_user_max_budget(): + valid_token = UserAPIKeyAuth(token="test_token", end_user_id="customer-1") + + end_user_params = { + "end_user_id": "customer-1", + "end_user_max_budget": 0.000000001, + } + + result = update_valid_token_with_end_user_params(valid_token, end_user_params) + + assert result.end_user_max_budget == 0.000000001 + + +def test_update_valid_token_preserves_custom_auth_max_budget_when_db_has_none(): + valid_token = UserAPIKeyAuth( + token="test_token", + end_user_id="customer-1", + end_user_max_budget=50.0, + ) + + end_user_params = { + "end_user_id": "customer-1", + } + + result = update_valid_token_with_end_user_params(valid_token, end_user_params) + + assert result.end_user_max_budget == 50.0 + + +@pytest.mark.asyncio +async def test_end_user_budget_counter_created_from_token_max_budget(): + from litellm.proxy.spend_tracking.budget_reservation import ( + _get_end_user_budget_counter, + ) + + token = UserAPIKeyAuth( + token="test_token", + end_user_id="customer-1", + end_user_max_budget=0.000000001, + ) + + counter = await _get_end_user_budget_counter( + valid_token=token, + end_user_id="customer-1", + end_user_object=None, + ) + + assert counter is not None + assert counter.max_budget == 0.000000001 + assert counter.counter_key == "spend:end_user:customer-1" + + +@pytest.mark.asyncio +async def test_end_user_budget_counter_none_when_max_budget_missing(): + from litellm.proxy.spend_tracking.budget_reservation import ( + _get_end_user_budget_counter, + ) + + token = UserAPIKeyAuth( + token="test_token", + end_user_id="customer-1", + ) + + counter = await _get_end_user_budget_counter( + valid_token=token, + end_user_id="customer-1", + end_user_object=None, + ) + + assert counter is None