mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Fix tag spend reset for budget tiers
Co-authored-by: ishaan-berri <ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
f2e97380d2
commit
93b005d0fb
2 changed files with 176 additions and 0 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue