fix: invalidate cached tag object on tag budget reset (#27481) (#27572)

Squash-merged by litellm-agent from oss-agent-shin's PR.
This commit is contained in:
oss-agent-shin 2026-05-09 18:31:43 -07:00 • committed by GitHub
parent 02edaef50c
commit 0a6bec161b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 94 additions and 0 deletions

View file

@ -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:<name>``) 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:<name>`` 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):

View file

@ -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()