diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3e24f50ed97..275d278c14e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1296,6 +1296,19 @@ async def get_team_membership( return None +async def get_team_member_model_spend( + user_id: str, + team_id: str, + model: str, + user_api_key_cache: DualCache, +) -> float: + _key = f"team_member_model_spend:{user_id}:{team_id}:{model}" + cached_spend = await user_api_key_cache.async_get_cache(key=_key) + if cached_spend is not None: + return float(cached_spend) + return 0.0 + + def model_in_access_group( model: str, team_models: Optional[List[str]], llm_router: Optional[Router] ) -> bool: @@ -3385,9 +3398,7 @@ async def _check_team_member_model_budget( from litellm.proxy.proxy_server import ( _get_team_member_model_counter_key, - _reseed_spend_from_db, get_current_spend, - spend_counter_cache, ) for requested_model_str in models_to_check: @@ -3429,15 +3440,18 @@ async def _check_team_member_model_budget( team_id=team_object.team_id, model=counter_model, ) + fallback_spend = 0.0 + if user_api_key_cache is not None: + fallback_spend = await get_team_member_model_spend( + user_id=valid_token.user_id, + team_id=team_object.team_id, + model=counter_model, + user_api_key_cache=user_api_key_cache, + ) model_spend = await get_current_spend( counter_key=counter_key, - fallback_spend=-1.0, + fallback_spend=fallback_spend, ) - if model_spend < 0: - model_spend = await _reseed_spend_from_db(counter_key) - await spend_counter_cache.async_set_cache( - key=counter_key, value=model_spend, nx=True - ) if model_spend >= max_budget: raise litellm.BudgetExceededError( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0dca9832593..1a1b8692ff3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1926,6 +1926,10 @@ async def increment_spend_counters( source_cache_key=None, increment=response_cost, ) + await user_api_key_cache.async_increment_cache( + key=f"team_member_model_spend:{user_id}:{team_id}:{model}", + value=response_cost, + ) if user_id is not None: await _init_and_increment_spend_counter( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index c8af66954e4..07567834911 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2390,7 +2390,7 @@ async def test_team_member_model_budget_accepts_numeric_string_max_budget(): @pytest.mark.asyncio -async def test_team_member_model_budget_reseeds_on_counter_miss(): +async def test_team_member_model_budget_uses_cached_fallback_spend(): team_object = LiteLLM_TeamTable( team_id="test-team", metadata={"team_member_model_max_budget": {"gpt-4o": {"max_budget": 10.0}}}, @@ -2401,20 +2401,15 @@ async def test_team_member_model_budget_reseeds_on_counter_miss(): team_id="test-team", ) + user_api_key_cache = MagicMock() + user_api_key_cache.async_get_cache = AsyncMock(return_value=11.0) + async def mock_get_current_spend(counter_key, fallback_spend): + assert fallback_spend == 11.0 return fallback_spend with ( patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), - patch( - "litellm.proxy.proxy_server._reseed_spend_from_db", - new_callable=AsyncMock, - return_value=11.0, - ) as mock_reseed, - patch( - "litellm.proxy.proxy_server.spend_counter_cache.async_set_cache", - new_callable=AsyncMock, - ) as mock_set_cache, ): with pytest.raises(litellm.BudgetExceededError) as exc_info: await _check_team_member_model_budget( @@ -2422,15 +2417,7 @@ async def test_team_member_model_budget_reseeds_on_counter_miss(): valid_token=valid_token, model="gpt-4o", llm_router=None, - user_api_key_cache=None, + user_api_key_cache=user_api_key_cache, ) assert exc_info.value.current_cost == 11.0 assert exc_info.value.max_budget == 10.0 - mock_reseed.assert_awaited_once_with( - "spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o" - ) - mock_set_cache.assert_awaited_once_with( - key="spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o", - value=11.0, - nx=True, - ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d4367b91548..e4d618f5d67 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4968,6 +4968,10 @@ async def test_increment_spend_counters_team_and_member(): key="spend:team_member:team_id::team-1::user_id::user-1::model::gpt-4o" ) assert member_model_counter == 0.30 + model_fallback_spend = key_cache.in_memory_cache.get_cache( + key="team_member_model_spend:user-1:team-1:gpt-4o" + ) + assert model_fallback_spend == 0.30 finally: ps.user_api_key_cache = original_key_cache ps.spend_counter_cache = original_counter_cache