diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py index fcdb3c9246c..8b5e6c5497b 100644 --- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py +++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py @@ -1210,18 +1210,19 @@ async def test_a_pre_upgrade_counter_keyed_on_the_request_model_still_enforces(e model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key=f"{prefix}:entity-1:openai/gpt-4:1d", value=25.0, ttl=86400) + if entity_type == Litellm_EntityType.KEY: + budget_check = limiter.is_key_within_model_budget( + user_api_key_dict=UserAPIKeyAuth(token="entity-1", model_max_budget=model_max_budget), + model="openai/gpt-4", + ) + else: + budget_check = limiter.is_end_user_within_model_budget( + end_user_id="entity-1", + end_user_model_max_budget=model_max_budget, + model="openai/gpt-4", + ) with pytest.raises(litellm.BudgetExceededError) as exc_info: - if entity_type == Litellm_EntityType.KEY: - await limiter.is_key_within_model_budget( - user_api_key_dict=UserAPIKeyAuth(token="entity-1", model_max_budget=model_max_budget), - model="openai/gpt-4", - ) - else: - await limiter.is_end_user_within_model_budget( - end_user_id="entity-1", - end_user_model_max_budget=model_max_budget, - model="openai/gpt-4", - ) + await budget_check assert exc_info.value.current_cost == 25.0 diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index e9566254dbc..ea7e298c380 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -1357,9 +1357,8 @@ async def fake_get_user_object(**kwargs): new=fake_get_user_object, ): if expect_refusal: - with pytest.raises(Exception) as exc: + with pytest.raises(Exception, match=r"(?i)budget") as exc: await user_api_key_auth(request=request, api_key="Bearer " + key) - assert "budget" in str(exc.value).lower() assert user_id in str(exc.value) else: result = await user_api_key_auth(request=request, api_key="Bearer " + key)