diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e17e359e183..275d278c14e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -659,7 +659,6 @@ async def common_checks( # noqa: PLR0915 valid_token=valid_token, model=_model, llm_router=llm_router, - prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) @@ -1297,32 +1296,17 @@ async def get_team_membership( return None -@log_db_metrics async def get_team_member_model_spend( user_id: str, team_id: str, model: str, - prisma_client: Optional[PrismaClient], user_api_key_cache: DualCache, ) -> float: - if prisma_client is None: - return 0.0 _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) - row = await prisma_client.db.litellm_teammembermodelspend.find_unique( - where={ - "user_id_team_id_model": { - "user_id": user_id, - "team_id": team_id, - "model": model, - } - } - ) - spend = float(row.spend if row is not None else 0.0) - await user_api_key_cache.async_set_cache(key=_key, value=spend, ttl=5) - return spend + return 0.0 def model_in_access_group( @@ -3381,7 +3365,6 @@ async def _check_team_member_model_budget( valid_token: Optional[UserAPIKeyAuth], model: Optional[Union[str, List[str]]], llm_router: Optional[Router], - prisma_client: Optional[PrismaClient] = None, user_api_key_cache: Optional[DualCache] = None, ): """ @@ -3458,12 +3441,11 @@ async def _check_team_member_model_budget( model=counter_model, ) fallback_spend = 0.0 - if prisma_client is not None and user_api_key_cache is not None: + 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, - prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) model_spend = await get_current_spend( diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 1b24f86964e..87d4b2aae14 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -73,6 +73,7 @@ class ResetBudgetJob: _get_team_member_counter_key, _get_team_member_model_counter_key, spend_counter_cache, + user_api_key_cache, ) memberships = await self.prisma_client.db.litellm_teammembership.find_many( @@ -116,9 +117,13 @@ class ResetBudgetJob: team_id=row.team_id, model=row.model, ) + model_spend_cache_key = f"team_member_model_spend:{row.user_id}:{row.team_id}:{row.model}" spend_counter_cache.in_memory_cache.set_cache( key=model_counter_key, value=0.0 ) + await user_api_key_cache.async_set_cache( + key=model_spend_cache_key, value=0.0, ttl=5 + ) if spend_counter_cache.redis_cache is not None: try: await spend_counter_cache.redis_cache.async_set_cache( diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index e8f992dacc8..5ded031efef 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -859,7 +859,10 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format(): spend_counter_cache = MagicMock() spend_counter_cache.in_memory_cache.set_cache = MagicMock() spend_counter_cache.redis_cache = None + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock(return_value=None) fake_module.spend_counter_cache = spend_counter_cache + fake_module.user_api_key_cache = user_api_key_cache fake_module._get_team_member_counter_key = ( lambda user_id, team_id: f"spend:team_member:team_id::{team_id}::user_id::{user_id}" ) @@ -885,6 +888,11 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format(): where={"OR": [{"user_id": "user-1", "team_id": "team-1"}]}, data={"spend": 0}, ) + user_api_key_cache.async_set_cache.assert_awaited_once_with( + key="team_member_model_spend:user-1:team-1:gpt-4o", + value=0.0, + ttl=5, + ) # ---------------------------------------------------------------------------