From 93b005d0fbabf9e8c692c76ed9a416d2632bab3b Mon Sep 17 00:00:00 2001 From: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Date: Sat, 9 May 2026 00:09:54 +0000 Subject: [PATCH] Fix tag spend reset for budget tiers Co-authored-by: ishaan-berri --- .../proxy/common_utils/reset_budget_job.py | 62 ++++++++++ .../common_utils/test_reset_budget_job.py | 114 ++++++++++++++++++ 2 files changed, 176 insertions(+) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 0928ce914da..3f06c277d97 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -9,6 +9,7 @@ from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( LiteLLM_BudgetTableFull, LiteLLM_EndUserTable, + LiteLLM_TagTable, LiteLLM_TeamTable, LiteLLM_UserTable, LiteLLM_VerificationToken, @@ -83,6 +84,17 @@ class ResetBudgetJob: "Failed to reset spend counter %s: %s", counter_key, e ) + @staticmethod + async def _invalidate_tag_cache(tag_name: str) -> None: + try: + from litellm.proxy.proxy_server import user_api_key_cache + + await user_api_key_cache.async_delete_cache(key=f"tag:{tag_name}") + except Exception as e: + verbose_proxy_logger.warning( + "Failed to invalidate tag cache for %s: %s", tag_name, e + ) + async def reset_budget_for_litellm_team_members( self, budgets_to_reset: List[LiteLLM_BudgetTableFull] ): @@ -171,6 +183,52 @@ class ResetBudgetJob: return update_result + async def reset_budget_for_tags_linked_to_budgets( + self, budgets_to_reset: List[LiteLLM_BudgetTableFull] + ): + """ + Resets the spend for tags linked to budget tiers that are being reset. + """ + budget_ids = [ + budget.budget_id + for budget in budgets_to_reset + if budget.budget_id is not None + ] + if not budget_ids: + return + + where_clause: dict = { + "budget_id": {"in": budget_ids}, + "spend": {"gt": 0}, + } + + try: + tags: List[LiteLLM_TagTable] = ( + await self.prisma_client.db.litellm_tagtable.find_many( + where=where_clause + ) + ) + except Exception as e: + tags = [] + verbose_proxy_logger.warning( + "Failed to fetch tags for counter invalidation: %s", e + ) + + update_result = await self.prisma_client.db.litellm_tagtable.update_many( + where=where_clause, + data={ + "spend": 0, + }, + ) + + for tag in tags: + tag_name = getattr(tag, "tag_name", None) + if tag_name: + await self._invalidate_spend_counter(f"spend:tag:{tag_name}") + await self._invalidate_tag_cache(tag_name=tag_name) + + return update_result + async def reset_budget_for_litellm_budget_table(self): """ Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired @@ -237,6 +295,10 @@ class ResetBudgetJob: budgets_to_reset=budgets_to_reset ) + await self.reset_budget_for_tags_linked_to_budgets( + budgets_to_reset=budgets_to_reset + ) + if endusers_to_reset is not None and len(endusers_to_reset) > 0: for enduser in endusers_to_reset: try: 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 5c86f9057a1..451b7acd884 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 @@ -39,6 +39,23 @@ class MockLiteLLMVerificationToken: return {"count": 1} +class MockLiteLLMTagTable: + def __init__(self): + self.find_many_calls: List[Dict[str, Any]] = [] + self.update_many_calls: List[Dict[str, Any]] = [] + self.find_many_results: List[Any] = [] + + async def find_many(self, where: Dict[str, Any]) -> List[Any]: + self.find_many_calls.append({"where": where}) + return self.find_many_results + + async def update_many( + self, where: Dict[str, Any], data: Dict[str, Any] + ) -> Dict[str, Any]: + self.update_many_calls.append({"where": where, "data": data}) + return {"count": len(self.find_many_results)} + + class MockLiteLLMEndUserTable: def __init__(self): self.find_many_calls: List[Dict[str, Any]] = [] @@ -56,6 +73,7 @@ class MockDB: def __init__(self): self.litellm_teammembership = MockLiteLLMTeamMembership() self.litellm_verificationtoken = MockLiteLLMVerificationToken() + self.litellm_tagtable = MockLiteLLMTagTable() self.litellm_endusertable = MockLiteLLMEndUserTable() @@ -618,6 +636,102 @@ def test_budget_table_reset_also_resets_linked_keys( assert calls[0]["data"]["spend"] == 0 +def test_reset_budget_for_tags_linked_to_budgets_invalidates_caches(monkeypatch): + """Resetting tags via budget tier must clear DB spend and stale caches.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + + user_api_key_cache = MagicMock() + user_api_key_cache.async_delete_cache = AsyncMock() + fake_module = sys.modules["litellm.proxy.proxy_server"] + fake_module.user_api_key_cache = user_api_key_cache + + expired_budget = type("B", (), {"budget_id": "budget-1"}) + linked_tag = type("Tag", (), {"tag_name": "tenant:example"}) + + 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}) + + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) + asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) + + prisma_client.db.litellm_tagtable.find_many.assert_awaited_once_with( + where={"budget_id": {"in": ["budget-1"]}, "spend": {"gt": 0}} + ) + prisma_client.db.litellm_tagtable.update_many.assert_awaited_once_with( + where={"budget_id": {"in": ["budget-1"]}, "spend": {"gt": 0}}, + data={"spend": 0}, + ) + counter_cache.in_memory_cache.set_cache.assert_any_call( + key="spend:tag:tenant:example", value=0.0, ttl=60 + ) + user_api_key_cache.async_delete_cache.assert_awaited_once_with( + key="tag:tenant:example" + ) + + +def test_reset_budget_for_tags_linked_to_budgets_empty( + reset_budget_job, mock_prisma_client +): + """No tag-table query should run when no expiring budget has an ID.""" + asyncio.run( + reset_budget_job.reset_budget_for_tags_linked_to_budgets( + budgets_to_reset=[type("B", (), {"budget_id": None})] + ) + ) + + assert mock_prisma_client.db.litellm_tagtable.find_many_calls == [] + assert mock_prisma_client.db.litellm_tagtable.update_many_calls == [] + + +def test_budget_table_reset_also_resets_linked_tags( + reset_budget_job, mock_prisma_client +): + """ + Integration-style test: when a budget tier resets, tags linked to that + budget must have spend reset so future tagged requests are not blocked. + """ + now = datetime.now(timezone.utc) + + test_budget = type( + "LiteLLM_BudgetTableFull", + (), + { + "max_budget": 10.0, + "budget_duration": "7d", + "budget_reset_at": now - timedelta(hours=1), + "budget_id": "7d-budget-tier", + "created_at": now - timedelta(days=7), + }, + ) + + mock_prisma_client.data["budget"] = [test_budget] + mock_prisma_client.db.litellm_tagtable.find_many_results = [ + type( + "LiteLLM_TagTable", + (), + { + "tag_name": "tenant:example", + "spend": 11.0, + "budget_id": "7d-budget-tier", + }, + ) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + calls = mock_prisma_client.db.litellm_tagtable.update_many_calls + assert len(calls) == 1, ( + "Expected reset_budget_for_litellm_budget_table to also reset tags " + f"linked to expiring budgets, but got {len(calls)} update_many calls" + ) + assert calls[0]["where"] == { + "budget_id": {"in": ["7d-budget-tier"]}, + "spend": {"gt": 0}, + } + assert calls[0]["data"]["spend"] == 0 + + def test_reset_budget_resets_endusers_with_null_budget_id( reset_budget_job, mock_prisma_client ):