diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index d4c5d76ac13..3ca66d3c26c 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -83,6 +83,29 @@ class ResetBudgetJob: "Failed to reset spend counter %s: %s", counter_key, e ) + async def _invalidate_source_cache(self, source_cache_key: str) -> None: + """Drop a cached entity so the next budget check re-reads spend from DB. + + Without this, a row whose spend was just zeroed in DB can still be + served from user_api_key_cache (e.g. ``tag:``) with the + pre-reset spend value, which is consulted as the cold-start / + DB-unavailable fallback in get_current_spend(). Run this AFTER the + DB write commits, mirroring _invalidate_spend_counter. + """ + if self.proxy_logging_obj is None: + return + user_api_key_cache = self.proxy_logging_obj.call_details.get( + "user_api_key_cache" + ) + if user_api_key_cache is None: + return + try: + await user_api_key_cache.async_delete_cache(key=source_cache_key) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to invalidate source cache key %s: %s", source_cache_key, e + ) + async def _cascade_reset_spend_for_budget_link( self, budgets_to_reset: List[LiteLLM_BudgetTableFull], @@ -90,9 +113,14 @@ class ResetBudgetJob: counter_key_fn: Callable[[Any], str], log_subject: str, extra_where: Optional[dict] = None, + source_cache_key_fn: Optional[Callable[[Any], str]] = None, ): """ Generic cascade: zero spend on rows whose budget_id is in the reset set. + + When ``source_cache_key_fn`` is supplied, the corresponding entry in + user_api_key_cache is also evicted so the next request re-reads the + zeroed spend from DB rather than the stale cached object. """ budget_ids = [b.budget_id for b in budgets_to_reset if b.budget_id is not None] if not budget_ids: @@ -114,6 +142,8 @@ class ResetBudgetJob: for row in rows: await self._invalidate_spend_counter(counter_key_fn(row)) + if source_cache_key_fn is not None: + await self._invalidate_source_cache(source_cache_key_fn(row)) return update_result @@ -166,6 +196,14 @@ class ResetBudgetJob: ): """ Resets the spend for tags linked to budget tiers that are being reset. + + Also evicts the cached LiteLLM_TagTable object at ``tag:`` in + user_api_key_cache. The auth-time tag budget check (see + ``_tag_max_budget_check`` in litellm/proxy/auth/auth_checks.py) reads + ``tag_object.spend`` as the DB-unavailable fallback in + ``get_current_spend``; if that cached object survives a reset it can + keep blocking otherwise-unblocked tenants under cold-start / + Redis-down scenarios. """ return await self._cascade_reset_spend_for_budget_link( budgets_to_reset=budgets_to_reset, @@ -173,6 +211,7 @@ class ResetBudgetJob: counter_key_fn=lambda t: f"spend:tag:{t.tag_name}", log_subject="tags", extra_where={"spend": {"gt": 0}}, + source_cache_key_fn=lambda t: f"tag:{t.tag_name}", ) async def reset_budget_for_litellm_budget_table(self): 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 82511fbc55e..2dc472f82a5 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 @@ -1458,3 +1458,58 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monke counter_cache.redis_cache.async_set_cache.assert_any_await( key="spend:tag:tenant-42", value=0.0, ttl=60 ) + + +def test_reset_budget_for_tags_linked_to_budgets_invalidates_source_cache(monkeypatch): + """Resetting tags must also evict the cached LiteLLM_TagTable object so + the auth-time fallback (``tag_object.spend``) does not keep blocking a + tenant after spend has been zeroed in the DB. + """ + _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) + + user_api_key_cache = MagicMock() + user_api_key_cache.async_delete_cache = AsyncMock() + + proxy_logging_obj = MagicMock() + proxy_logging_obj.call_details = {"user_api_key_cache": user_api_key_cache} + + job = ResetBudgetJob( + proxy_logging_obj=proxy_logging_obj, prisma_client=prisma_client + ) + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + user_api_key_cache.async_delete_cache.assert_any_await(key="tag:tenant-42") + + +def test_reset_budget_for_tags_linked_to_budgets_no_user_api_key_cache(monkeypatch): + """When user_api_key_cache is not wired up (e.g. early-boot or tests), + the cascade must still complete without raising — the spend counter + invalidation is the load-bearing path. + """ + _make_counter_invalidation_job(monkeypatch) + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) + + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag]) + prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1}) + + proxy_logging_obj = MagicMock() + proxy_logging_obj.call_details = {} + + job = ResetBudgetJob( + proxy_logging_obj=proxy_logging_obj, prisma_client=prisma_client + ) + # Should not raise even though user_api_key_cache is missing. + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + # The DB write must still have happened. + prisma_client.db.litellm_tagtable.update_many.assert_awaited_once()