Merge pull request #39729 from amasen02/fix/end-user-budget-reset-cache-invalidation-39726

fix(proxy): invalidate end-user spend counter and cache on budget reset (#39726)
This commit is contained in:
Mateo Wang 2026-09-04 18:36:05 -07:00 committed by GitHub
commit 3d08daecfe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 283 additions and 30 deletions

View file

@ -38,6 +38,7 @@ from litellm.proxy.common_utils.timezone_utils import (
get_budget_reset_settings,
)
from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
model_access_group_cache_key,
model_access_group_spend_counter_key,
tag_cache_key,
@ -177,6 +178,21 @@ def _model_access_group_cache_keys(row: _ModelAccessGroupRow) -> tuple[str, ...]
return (model_access_group_cache_key(row.access_group_name),)
def _enduser_counter_key(row: _EndUserRow) -> str:
return f"spend:end_user:{row.user_id}"
def _enduser_cache_keys(row: _EndUserRow) -> tuple[str, ...]:
return (end_user_cache_key(row.user_id),)
def _enduser_carried_spend(row: _EndUserRow, caps: Mapping[str, float]) -> float:
if not caps:
return 0.0
effective_budget_id: Final[str | None] = row.budget_id or litellm.max_end_user_budget_id
return _carried_spend(row.spend, caps.get(effective_budget_id) if effective_budget_id is not None else None)
def _budget_link_where(
budget_ids: Sequence[str],
extra: Mapping[str, object] = MappingProxyType({}),
@ -650,6 +666,7 @@ class ResetBudgetJob:
if _rollover_enabled()
else {} # mutable-ok: empty sentinel immediately frozen by MappingProxyType
)
endusers: Final[tuple[_EndUserRow, ...]] = await self._collect_endusers_to_reset(budget_ids)
return _BudgetCascade(
budgets=tuple(budgets_to_reset),
budget_ids=budget_ids,
@ -661,7 +678,7 @@ class ResetBudgetJob:
for b in budgets_to_reset
if b.budget_id is not None and b.budget_duration is not None
),
endusers=await self._collect_endusers_to_reset(budget_ids),
endusers=endusers,
counter_resets=(
*(
(_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps))
@ -674,6 +691,7 @@ class ResetBudgetJob:
(_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps))
for row in model_access_groups
),
*((_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) for row in endusers),
),
rollover_caps=rollover_caps,
cache_keys=(
@ -682,6 +700,7 @@ class ResetBudgetJob:
*(key for row in orgs for key in _org_cache_keys(row)),
*(key for row in tags for key in _tag_cache_keys(row)),
*(key for row in model_access_groups for key in _model_access_group_cache_keys(row)),
*(key for row in endusers for key in _enduser_cache_keys(row)),
),
)

View file

@ -26,6 +26,7 @@ from litellm.proxy._types import Litellm_EntityType
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.table_repositories import (
BudgetWindowSpendRepository,
EndUserRepository,
SpendLogsRepository,
TeamMembershipRepository,
)
@ -36,6 +37,8 @@ from litellm.repositories.verification_token_repository import (
)
if TYPE_CHECKING:
from prisma.types import LiteLLM_EndUserTableWhereUniqueInput
from litellm.caching.dual_cache import DualCache
from litellm.proxy.utils import PrismaClient
@ -47,6 +50,8 @@ _WINDOW_SPEND_ENTITY_TYPES: Final[Mapping[str, str]] = MappingProxyType(
}
)
END_USER_COUNTER_PREFIX: Final = "spend:end_user:"
_WINDOW_SPEND_LOG_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{
"Key": "api_key",
@ -74,6 +79,10 @@ class SpendCounterReseed:
End-user and tag spend counters intentionally do not reseed here. Their
auth paths already load the corresponding objects via get_end_user_object()
and get_tag_objects_batch(); callers pass those values as fallback_spend.
end_user_from_db is the one end-user read, used only as the budget floor when
a counter sits below that cached spend: a worker that did not run the budget
reset still caches the pre-reset end-user object, and LiteLLM_EndUserTable
is the row the reset zeroed.
"""
_locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict()
@ -129,7 +138,7 @@ class SpendCounterReseed:
elif counter_key.startswith("spend:user:"):
user_id = counter_key[len("spend:user:") :]
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
elif counter_key.startswith("spend:end_user:") or counter_key.startswith("spend:tag:"):
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
return None
elif counter_key.startswith("spend:org:"):
org_id: Final = counter_key[len("spend:org:") :]
@ -143,6 +152,20 @@ class SpendCounterReseed:
return None
return float(getattr(row, "spend", 0.0) or 0.0)
@staticmethod
async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None:
if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX):
return None
where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]}
try:
row: Final = await EndUserRepository(prisma_client).table.find_unique(where=where)
except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db
verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key)
return None
if row is None:
return None
return float(row.spend or 0.0)
@staticmethod
def _is_key_or_team_window_counter(counter_key: str) -> bool:
for prefix in ("spend:key:", "spend:team:"):

View file

@ -423,7 +423,7 @@ from litellm.proxy.db.proxy_worker_heartbeat import (
PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS,
ProxyWorkerHeartbeat,
)
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
from litellm.proxy.db.spend_counter_reseed import END_USER_COUNTER_PREFIX, SpendCounterReseed
from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router
from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router
from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config
@ -2477,7 +2477,8 @@ async def get_current_spend(
authoritative source depends on the counter: primary key/team/user/org
counters read the DB row; per-window counters (``window_start`` supplied)
read the maintained window-spend row and only aggregate spend logs when
that row is missing or stale; end-user/tag counters have no DB row, so the caller's
that row is missing or stale; end-user counters read ``LiteLLM_EndUserTable``, the
row the budget reset zeroes; tag counters have no DB row, so the caller's
``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is
skipped for healthy primary counters (counter at or above recorded spend)
and cached in-process for a few seconds, so a persistently stale counter
@ -2511,8 +2512,8 @@ async def get_current_spend(
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=authoritative)
return authoritative
elif fallback_spend > current:
# end-user / tag counters have no DB row; fallback_spend is the
# authoritative recorded value loaded in auth.
# nothing to read (tag counters, an end user without a row or a DB client, a
# failed read); fallback_spend is the authoritative recorded value loaded in auth.
return fallback_spend
# Opt-in hard guarantee: when the spend backing this admit decision came
@ -2580,6 +2581,29 @@ async def reseed_spend_counter_from_db(counter_key: str) -> None:
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend)
async def _floor_spend_from_db(
counter_key: str,
window_entity_type: str | None,
window_entity_id: str | None,
window_duration: str | None,
window_start: datetime | None,
) -> float | None:
if counter_key.startswith(END_USER_COUNTER_PREFIX):
return await SpendCounterReseed.end_user_from_db(prisma_client=prisma_client, counter_key=counter_key)
entity_spend: Final = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key)
if entity_spend is not None:
return entity_spend
if window_entity_type is None or window_entity_id is None or window_start is None:
return None
return await SpendCounterReseed.window_from_db(
prisma_client=prisma_client,
entity_type=window_entity_type,
entity_id=window_entity_id,
window_duration=window_duration,
window_start=window_start,
)
async def _authoritative_floor_spend(
counter_key: str,
window_entity_type: str | None = None,
@ -2592,20 +2616,13 @@ async def _authoritative_floor_spend(
if cached is not None:
return float(cached)
db_spend = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key)
if (
db_spend is None
and window_entity_type is not None
and window_entity_id is not None
and window_start is not None
):
db_spend = await SpendCounterReseed.window_from_db(
prisma_client=prisma_client,
entity_type=window_entity_type,
entity_id=window_entity_id,
window_duration=window_duration,
window_start=window_start,
)
db_spend: Final = await _floor_spend_from_db(
counter_key=counter_key,
window_entity_type=window_entity_type,
window_entity_id=window_entity_id,
window_duration=window_duration,
window_start=window_start,
)
if db_spend is None:
return None

View file

@ -4,7 +4,7 @@ import sys
import types
from datetime import datetime, timedelta, timezone
from datetime import time as dt_time
from typing import Any, Dict, List
from typing import Any, Dict, Final, List
from unittest.mock import AsyncMock, MagicMock
import httpx
@ -1495,6 +1495,32 @@ def test_budget_table_reset_invalidates_every_tag_not_just_the_first(reset_budge
assert deleted == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"}
def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_job, mock_prisma_client, monkeypatch):
"""When an end user's budget resets, its Redis spend counter is zeroed and its management cache is evicted."""
counter_cache: Final = _make_counter_invalidation_job(monkeypatch)
budget: Final = _budget_row(budget_id="budget-1")
mock_prisma_client.data["budget"] = [budget]
test_enduser: Final = type(
"LiteLLM_EndUserTable",
(),
{
"spend": 20.0,
"litellm_budget_table": budget,
"budget_id": "budget-1",
"user_id": "customer-42",
},
)
mock_prisma_client.data["enduser"] = [test_enduser]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60)
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:customer-42" in deleted
def test_budget_table_reset_commits_even_when_cache_eviction_fails(reset_budget_job, mock_prisma_client, monkeypatch):
"""Eviction runs after the commit, so a broken cache cannot undo the write."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
@ -3028,6 +3054,38 @@ def test_budget_cascade_carries_enduser_overage_when_rollover_enabled(
} in enduser_writes
def test_budget_cascade_carries_default_tier_enduser_counter_when_rollover_enabled(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
"""An end user on the default budget (no budget_id on its row) 5 over the cap
keeps a counter of 5 in the next window and loses its cached object."""
import litellm
counter_cache: Final = _make_counter_invalidation_job(monkeypatch)
monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-enduser-budget")
mock_prisma_client.data["budget"] = [
_budget_row(budget_id="default-enduser-budget", budget_duration="1d", max_budget=10.0)
]
implicit_enduser: Final = type(
"EndUserRow",
(),
{
"spend": 15.0,
"user_id": "enduser-implicit",
"budget_id": None,
"model_dump": lambda self=None: {"spend": 15.0, "user_id": "enduser-implicit", "budget_id": None, "blocked": False},
},
)
mock_prisma_client.db.litellm_endusertable.set_find_many_results([implicit_enduser])
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:enduser-implicit" in deleted
def _replay_spend_writes(writes, spend):
"""Apply the queued update_many statements in order, the way the DB
transaction executes them, and return the row's final spend."""

View file

@ -9,6 +9,7 @@ from __future__ import annotations
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
import pytest
@ -18,7 +19,7 @@ from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
WINDOW_START = datetime(2026, 8, 1, tzinfo=timezone.utc)
class _FakeWindowSpendTable:
class _FakeFindUniqueTable:
def __init__(self, row: SimpleNamespace | None, error: Exception | None = None) -> None:
self._row = row
self._error = error
@ -47,10 +48,13 @@ class _FakePrismaClient:
row: SimpleNamespace | None = None,
spend_logs_total: float = 0.0,
error: Exception | None = None,
end_user_row: SimpleNamespace | None = None,
end_user_error: Exception | None = None,
) -> None:
self.db = SimpleNamespace(
litellm_budgetwindowspend=_FakeWindowSpendTable(row=row, error=error),
litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error),
litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total),
litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error),
)
@ -248,3 +252,65 @@ async def test_coalesced_window_seeds_a_cold_counter_from_the_row():
assert result == 4.5
assert cache.in_memory_cache.get_cache(key=counter_key) == 4.5
assert prisma.db.litellm_spendlogs.call_count == 0
@pytest.mark.asyncio
async def test_end_user_from_db_reads_the_end_user_row_by_user_id():
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0))
result: Final = await SpendCounterReseed.end_user_from_db(
prisma_client=prisma, counter_key="spend:end_user:customer-42"
)
assert result == 0.0
assert prisma.db.litellm_endusertable.where_clauses == [{"user_id": "customer-42"}]
@pytest.mark.asyncio
async def test_end_user_from_db_returns_the_recorded_spend():
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=12.5))
assert (
await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42")
== 12.5
)
@pytest.mark.asyncio
@pytest.mark.parametrize("counter_key", ["spend:key:hashed", "spend:team:t1", "spend:tag:t1"])
async def test_end_user_from_db_ignores_other_counter_kinds_without_touching_the_db(counter_key):
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="x", spend=5.0))
assert await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key=counter_key) is None
assert prisma.db.litellm_endusertable.where_clauses == []
@pytest.mark.asyncio
async def test_end_user_from_db_returns_none_without_a_row_a_client_or_on_db_error():
assert (
await SpendCounterReseed.end_user_from_db(prisma_client=None, counter_key="spend:end_user:customer-42")
is None
)
assert (
await SpendCounterReseed.end_user_from_db(
prisma_client=_FakePrismaClient(end_user_row=None), counter_key="spend:end_user:customer-42"
)
is None
)
assert (
await SpendCounterReseed.end_user_from_db(
prisma_client=_FakePrismaClient(end_user_error=RuntimeError("db down")),
counter_key="spend:end_user:customer-42",
)
is None
)
@pytest.mark.asyncio
async def test_from_db_still_never_reads_the_end_user_row():
"""A cold end-user counter keeps seeding from the cached end-user object the auth
path already loaded; the row is read only as the budget floor."""
prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=5.0))
assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") is None
assert prisma.db.litellm_endusertable.where_clauses == []

View file

@ -22,6 +22,7 @@ from __future__ import annotations
import asyncio
from datetime import datetime
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -222,16 +223,19 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch):
@pytest.mark.asyncio
async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch):
"""End-user and tag counters have no DB row (from_db returns None). When the
counter is stale-low, enforcement falls back to the caller's recorded spend
(loaded fresh in auth) instead of trusting the stale counter."""
fake_cache = _make_spend_counter_cache(redis_get_value=2.0)
@pytest.mark.parametrize("counter_key", ("spend:end_user:e1", "spend:tag:t1"))
async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch, counter_key):
"""Tag counters have no DB row (from_db returns None), and an end-user counter has
none to read without a DB client. When such a counter is stale-low, enforcement
falls back to the caller's recorded spend (loaded fresh in auth) instead of
trusting the stale counter."""
fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", AsyncMock(return_value=None))
result = await ps.get_current_spend(
counter_key="spend:end_user:e1",
result: Final = await ps.get_current_spend(
counter_key=counter_key,
fallback_spend=20.0,
max_budget=10.0,
)
@ -241,6 +245,72 @@ async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatc
fake_cache.redis_cache.async_set_max.assert_not_called()
def _make_prisma_with_end_user_row(spend: float | None):
prisma: Final = MagicMock()
prisma.db.litellm_endusertable.find_unique = AsyncMock(
return_value=None if spend is None else MagicMock(spend=spend)
)
return prisma
@pytest.mark.asyncio
async def test_get_current_spend_end_user_floor_admits_after_a_reset_on_a_stale_worker(monkeypatch):
"""The reset job zeroes LiteLLM_EndUserTable.spend and the shared counter, but it
evicts the cached end-user object only on the worker that ran the reset. Every
other worker still passes the pre-reset spend as fallback_spend, and that stale
copy must not out-vote the reset row."""
fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
prisma: Final = _make_prisma_with_end_user_row(spend=0.0)
monkeypatch.setattr(ps, "prisma_client", prisma)
result = await ps.get_current_spend(
counter_key="spend:end_user:customer-42",
fallback_spend=0.000032,
max_budget=0.00003,
fallback_authoritative=True,
)
assert result == 0.0
prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with(where={"user_id": "customer-42"})
fake_cache.redis_cache.async_set_max.assert_not_called()
@pytest.mark.asyncio
async def test_get_current_spend_end_user_floor_repairs_a_stale_low_counter(monkeypatch):
"""After a Redis restart the end-user counter can sit below the recorded spend;
the row wins and the shared counter is raised so other workers stop admitting on
the stale value."""
fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=12.0))
result: Final = await ps.get_current_spend(
counter_key="spend:end_user:customer-42",
fallback_spend=12.0,
max_budget=10.0,
)
assert result == 12.0
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:end_user:customer-42", value=12.0)
@pytest.mark.asyncio
async def test_get_current_spend_end_user_without_a_row_keeps_the_cached_spend(monkeypatch):
fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=None))
result: Final = await ps.get_current_spend(
counter_key="spend:end_user:customer-42",
fallback_spend=20.0,
max_budget=10.0,
)
assert result == 20.0
fake_cache.redis_cache.async_set_max.assert_not_called()
@pytest.mark.asyncio
async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
"""Per-window counters have no DB row but aggregate from spend logs. A