feat(proxy): opt-in budget rollover carrying overage into the next window (#38514)

* feat(proxy): opt-in budget rollover carrying overage into the next window

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): zero under-cap rows before decrementing over-cap rows in cascade resets

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-08-27 12:46:09 -07:00 • committed by GitHub
parent f864908cd7
commit de53283356
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 515 additions and 40 deletions

View file

@ -445,6 +445,7 @@ max_ui_session_budget: Optional[float] = (
1.0 # USD budget for each dashboard login session (playground, test connection)
)
internal_user_budget_duration: Optional[str] = None
budget_rollover: bool = False # carry spend beyond max_budget into the next window instead of zeroing it
tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None
max_end_user_budget: Optional[float] = None
max_end_user_budget_id: Optional[str] = None

View file

@ -1652,6 +1652,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
"enable_anthropic_prompt_caching",
"anthropic_prompt_caching_ttl",
"max_ui_session_budget",
"budget_rollover",
]
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))

View file

@ -1,5 +1,6 @@
import asyncio
import json
import math
import time
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
@ -45,6 +46,7 @@ from litellm.repositories.table_repositories import (
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import (
LinkedSpendResetWrites,
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
@ -59,7 +61,15 @@ _LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_dura
_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}})
class _TeamMembershipRow(Protocol):
class _BudgetLinkedRow(Protocol):
@property
def spend(self) -> float | None: ...
@property
def budget_id(self) -> str | None: ...
class _TeamMembershipRow(_BudgetLinkedRow, Protocol):
@property
def user_id(self) -> str: ...
@ -67,26 +77,48 @@ class _TeamMembershipRow(Protocol):
def team_id(self) -> str: ...
class _KeyRow(Protocol):
class _KeyRow(_BudgetLinkedRow, Protocol):
@property
def token(self) -> str: ...
class _OrgRow(Protocol):
class _OrgRow(_BudgetLinkedRow, Protocol):
@property
def organization_id(self) -> str: ...
class _TagRow(Protocol):
class _TagRow(_BudgetLinkedRow, Protocol):
@property
def tag_name(self) -> str: ...
class _EndUserRow(Protocol):
class _EndUserRow(_BudgetLinkedRow, Protocol):
@property
def user_id(self) -> str: ...
def _rollover_enabled() -> bool:
return litellm.budget_rollover is True
def _rollover_cap(max_budget: float | None) -> float | None:
if max_budget is None or not math.isfinite(max_budget):
return None
return max_budget
def _carried_spend(spend: float | None, cap: float | None) -> float:
if cap is None:
return 0.0
return max(0.0, (spend or 0.0) - cap)
def _row_carried_spend(row: _BudgetLinkedRow, caps: Mapping[str, float]) -> float:
if not caps:
return 0.0
return _carried_spend(row.spend, caps.get(row.budget_id) if row.budget_id is not None else None)
def _team_membership_counter_key(row: _TeamMembershipRow) -> str:
return f"spend:team_member:{row.user_id}:{row.team_id}"
@ -129,6 +161,59 @@ def _budget_link_where(
return {"budget_id": {"in": list(budget_ids)}, **extra}
def _queue_budget_linked_resets(
writes: LinkedSpendResetWrites,
cascade: "_BudgetCascade",
extra: Mapping[str, object] = MappingProxyType({}),
) -> None:
"""Reset one linked table's spend for every expiring tier: tiers with a
rollover cap keep spend beyond the cap (decrement preserves writes racing
the reset), everything else is zeroed as before. Zero the under-cap rows
BEFORE decrementing the over-cap ones: the statements run sequentially in
one transaction, so the reverse order lets the zero re-match a row the
decrement just moved into the (0, cap] range and erase its carried spend."""
for budget_id, cap in cascade.rollover_caps.items():
writes.queue_spend_zero(
where={"budget_id": budget_id, **extra, "spend": {"gt": 0, "lte": cap}}
) # mutable-ok: prisma where filter must be a dict
writes.queue_spend_decrement(
where={"budget_id": budget_id, **extra, "spend": {"gt": cap}}, amount=cap
) # mutable-ok: prisma where filter must be a dict
plain_ids: Final = tuple(bid for bid in cascade.budget_ids if bid not in cascade.rollover_caps)
if plain_ids:
writes.queue_spend_zero(where=_budget_link_where(plain_ids, extra))
def _queue_enduser_resets(writes: LinkedSpendResetWrites, cascade: "_BudgetCascade") -> None:
"""End users are matched by id rather than budget link: rows with no
budget_id ride the default budget tier (litellm.max_end_user_budget_id).
Zero-before-decrement ordering matters here too (see
_queue_budget_linked_resets)."""
if not cascade.rollover_caps:
if cascade.endusers:
writes.queue_spend_zero(
where={"user_id": {"in": [row.user_id for row in cascade.endusers]}}
) # mutable-ok: prisma where filter must be a dict
return
tiered: Final = tuple((row.budget_id or litellm.max_end_user_budget_id, row.user_id) for row in cascade.endusers)
for budget_id, cap in cascade.rollover_caps.items():
if not (
user_ids := [uid for bid, uid in tiered if bid == budget_id]
): # mutable-ok: prisma "in" filter takes a list
continue
writes.queue_spend_zero(
where={"user_id": {"in": user_ids}, "spend": {"lte": cap}}
) # mutable-ok: prisma where filter must be a dict
writes.queue_spend_decrement(
where={"user_id": {"in": user_ids}, "spend": {"gt": cap}}, amount=cap
) # mutable-ok: prisma where filter must be a dict
plain: Final = [
uid for bid, uid in tiered if bid is None or bid not in cascade.rollover_caps
] # mutable-ok: prisma "in" filter takes a list
if plain:
writes.queue_spend_zero(where={"user_id": {"in": plain}}) # mutable-ok: prisma where filter must be a dict
@dataclass(frozen=True, slots=True)
class _BudgetCascade:
"""Everything one budget-tier reset touches, resolved before any write."""
@ -137,8 +222,9 @@ class _BudgetCascade:
budget_ids: tuple[str, ...] = ()
budget_resets: tuple[tuple[str, datetime], ...] = ()
endusers: tuple[_EndUserRow, ...] = ()
counter_keys: tuple[str, ...] = ()
counter_resets: tuple[tuple[str, float], ...] = ()
cache_keys: tuple[str, ...] = ()
rollover_caps: Mapping[str, float] = MappingProxyType({})
@dataclass(frozen=True, slots=True)
@ -404,8 +490,10 @@ class ResetBudgetJob:
)
@staticmethod
async def _invalidate_spend_counter(counter_key: str) -> None:
"""Zero a spend counter so a DB-row reset takes effect immediately.
async def _invalidate_spend_counter(counter_key: str, new_spend: float = 0.0) -> None:
"""Overwrite a spend counter with the post-reset value (0, or the carried
overage when budget rollover is enabled) so a DB-row reset takes effect
immediately.
Call AFTER the DB write commits. Clearing Redis before the DB
commit opens a window where get_current_spend reads 0 from Redis
@ -414,10 +502,10 @@ class ResetBudgetJob:
try:
from litellm.proxy.proxy_server import spend_counter_cache
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.0, ttl=60)
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=new_spend, ttl=60)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=0.0, ttl=60)
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=new_spend, ttl=60)
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to reset spend counter %s in Redis: %s. "
@ -522,6 +610,15 @@ class ResetBudgetJob:
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="tags",
)
rollover_caps: Final[Mapping[str, float]] = MappingProxyType(
{ # mutable-ok: MappingProxyType wraps a one-shot dict comprehension
b.budget_id: cap
for b in budgets_to_reset
if b.budget_id is not None and (cap := _rollover_cap(b.max_budget)) is not None
}
if _rollover_enabled()
else {} # mutable-ok: empty sentinel immediately frozen by MappingProxyType
)
return _BudgetCascade(
budgets=tuple(budgets_to_reset),
budget_ids=budget_ids,
@ -534,12 +631,16 @@ class ResetBudgetJob:
if b.budget_id is not None and b.budget_duration is not None
),
endusers=await self._collect_endusers_to_reset(budget_ids),
counter_keys=(
*(_team_membership_counter_key(row) for row in team_memberships),
*(_key_counter_key(row) for row in keys),
*(_org_counter_key(row) for row in orgs),
*(_tag_counter_key(row) for row in tags),
counter_resets=(
*(
(_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps))
for row in team_memberships
),
*((_key_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in keys),
*((_org_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in orgs),
*((_tag_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in tags),
),
rollover_caps=rollover_caps,
cache_keys=(
*(key for row in team_memberships for key in _team_membership_cache_keys(row)),
*(key for row in keys for key in _key_cache_keys(row)),
@ -565,20 +666,18 @@ class ResetBudgetJob:
)
async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None:
enduser_ids: Final = tuple(row.user_id for row in cascade.endusers)
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids))
uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE))
uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if enduser_ids:
uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}})
_queue_budget_linked_resets(uow.team_memberships, cascade)
_queue_budget_linked_resets(uow.keys, cascade, extra=_LINKED_KEYS_WHERE)
_queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE)
_queue_budget_linked_resets(uow.tags, cascade, extra=_SPENT_ROWS_WHERE)
_queue_enduser_resets(uow.endusers, cascade)
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
for counter_key in cascade.counter_keys:
await self._invalidate_spend_counter(counter_key)
for counter_key, new_spend in cascade.counter_resets:
await self._invalidate_spend_counter(counter_key, new_spend=new_spend)
for cache_key in cascade.cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
@ -708,7 +807,11 @@ class ResetBudgetJob:
for k in updated_keys:
if k.token is None:
continue
uow.keys.queue_spend_reset(token=k.token, budget_reset_at=k.budget_reset_at)
uow.keys.queue_spend_reset(
token=k.token,
budget_reset_at=k.budget_reset_at,
spend_decrement=k.max_budget if (k.spend or 0.0) > 0.0 else None,
)
async def _write_user_reset_updates(self, updated_users: list[LiteLLM_UserTable]) -> None:
"""
@ -726,7 +829,11 @@ class ResetBudgetJob:
async def _write_user_reset_updates_once(self, updated_users: list[LiteLLM_UserTable]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for u in updated_users:
uow.users.queue_spend_reset(user_id=u.user_id, budget_reset_at=u.budget_reset_at)
uow.users.queue_spend_reset(
user_id=u.user_id,
budget_reset_at=u.budget_reset_at,
spend_decrement=u.max_budget if (u.spend or 0.0) > 0.0 else None,
)
async def _write_team_reset_updates(self, updated_teams: list[LiteLLM_TeamTable]) -> None:
"""
@ -744,7 +851,11 @@ class ResetBudgetJob:
async def _write_team_reset_updates_once(self, updated_teams: list[LiteLLM_TeamTable]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for t in updated_teams:
uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at)
uow.teams.queue_spend_reset(
team_id=t.team_id,
budget_reset_at=t.budget_reset_at,
spend_decrement=t.max_budget if (t.spend or 0.0) > 0.0 else None,
)
def _emit_phase_failure(
self,
@ -820,7 +931,7 @@ class ResetBudgetJob:
for k in updated_keys:
token = getattr(k, "token", None)
if token:
await self._invalidate_spend_counter(f"spend:key:{token}")
await self._invalidate_spend_counter(f"spend:key:{token}", new_spend=k.spend or 0.0)
end_time = time.time()
outcome: Final = _ChunkOutcome(
@ -925,7 +1036,7 @@ class ResetBudgetJob:
for u in updated_users:
user_id = getattr(u, "user_id", None)
if user_id:
await self._invalidate_spend_counter(f"spend:user:{user_id}")
await self._invalidate_spend_counter(f"spend:user:{user_id}", new_spend=u.spend or 0.0)
if user_id == LITELLM_PROXY_BUDGET_NAME:
await self._invalidate_global_proxy_spend_cache()
@ -1034,7 +1145,7 @@ class ResetBudgetJob:
for t in updated_teams:
team_id = getattr(t, "team_id", None)
if team_id:
await self._invalidate_spend_counter(f"spend:team:{team_id}")
await self._invalidate_spend_counter(f"spend:team:{team_id}", new_spend=t.spend or 0.0)
end_time = time.time()
outcome: Final = _ChunkOutcome(
@ -1107,10 +1218,11 @@ class ResetBudgetJob:
reset_at: Final = datetime.fromisoformat(reset_at_str.replace("Z", "+00:00")).replace(tzinfo=None)
if reset_at > now:
return False
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.0)
new_value: Final = await ResetBudgetJob._window_carried_spend(window, counter_key, spend_counter_cache)
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=new_value)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=0.0)
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=new_value)
except Exception as redis_err:
verbose_proxy_logger.warning("Failed to reset Redis counter %s: %s", counter_key, redis_err)
window["reset_at"] = compute_budget_reset_at(
@ -1118,6 +1230,27 @@ class ResetBudgetJob:
).isoformat()
return True
@staticmethod
async def _window_carried_spend(
window: Mapping[str, object], counter_key: str, spend_counter_cache: DualCache
) -> float:
"""Per-window spend lives only in the counter, so the carried overage is
read from it before the reset overwrites it."""
if not _rollover_enabled():
return 0.0
window_max: Final = window.get("max_budget")
cap: Final = _rollover_cap(window_max) if isinstance(window_max, (int, float)) else None
if cap is None:
return 0.0
try:
current: Final = await spend_counter_cache.async_get_cache(key=counter_key)
except Exception as e: # noqa: BLE001 # an unreadable counter falls back to a plain zero reset
verbose_proxy_logger.warning("Failed to read spend counter %s for rollover: %s", counter_key, e)
return 0.0
if not isinstance(current, (int, float)):
return 0.0
return _carried_spend(float(current), cap)
async def reset_budget_windows(self) -> None:
"""
For keys and teams with budget_limits, reset any individual windows where
@ -1222,7 +1355,7 @@ class ResetBudgetJob:
still holds the pre-reset value, admitting requests past the cap.
"""
try:
item.spend = 0.0
item.spend = _carried_spend(item.spend, _rollover_cap(item.max_budget)) if _rollover_enabled() else 0.0
if hasattr(item, "budget_duration") and item.budget_duration is not None:
item.budget_reset_at = compute_budget_reset_at(
budget_duration=item.budget_duration, settings=reset_settings

View file

@ -16415,6 +16415,13 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: Final[dict[str, GeneralSettingsUILiteLLMFie
"tab": "prompt_caching",
"description": "Empty uses Anthropic's 5m default. 1h suits long sessions but doubles the cache write cost.",
},
"budget_rollover": { # mutable-ok: registry literal, frozen with its siblings below
"type": "Boolean",
"description": (
"Carry spend beyond max_budget into the next window when budgets reset, instead of "
"forgiving it. Applies to key, user, team, team member, org, tag and end-user budgets."
),
},
"max_ui_session_budget": {
"type": "Dollar",
"default": 1.0,

View file

@ -19,32 +19,57 @@ from collections.abc import AsyncGenerator, Callable, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from typing import Final
from litellm.repositories.prisma_protocols import BatchTable, PrismaBatch
def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float | None) -> Mapping[str, object]:
spend: Final[object] = (
{"decrement": spend_decrement} # mutable-ok: prisma update payload must be a dict
if spend_decrement is not None
else 0
)
return {"spend": spend, "budget_reset_at": budget_reset_at} # mutable-ok: prisma update payload must be a dict
@dataclass(frozen=True, slots=True)
class KeySpendResetWrites:
table: BatchTable
def queue_spend_reset(self, token: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"token": token}, data={"spend": 0, "budget_reset_at": budget_reset_at})
def queue_spend_reset(
self, token: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
self.table.update(
where={"token": token}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),
)
@dataclass(frozen=True, slots=True)
class UserSpendResetWrites:
table: BatchTable
def queue_spend_reset(self, user_id: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"user_id": user_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
def queue_spend_reset(
self, user_id: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
self.table.update(
where={"user_id": user_id}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),
)
@dataclass(frozen=True, slots=True)
class TeamSpendResetWrites:
table: BatchTable
def queue_spend_reset(self, team_id: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
def queue_spend_reset(
self, team_id: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
self.table.update(
where={"team_id": team_id}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),
)
@dataclass(frozen=True, slots=True)
@ -54,6 +79,14 @@ class LinkedSpendResetWrites:
def queue_spend_zero(self, where: Mapping[str, object]) -> None:
self.table.update_many(where=where, data={"spend": 0})
def queue_spend_decrement(self, where: Mapping[str, object], amount: float) -> None:
"""``decrement`` rather than a read-then-set, so spend written between the
cascade's read and its commit survives the reset instead of being erased."""
self.table.update_many(
where=where,
data={"spend": {"decrement": amount}}, # mutable-ok: prisma update payload must be a dict
)
@dataclass(frozen=True, slots=True)
class BudgetWindowWrites:

View file

@ -1458,7 +1458,7 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo
budget = _budget_row(budget_id="budget-1", budget_duration="7d")
mock_prisma_client.data["budget"] = [budget]
mock_prisma_client.data["enduser"] = [
type("EndUser", (), {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1"})
type("EndUser", (), {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1", "budget_id": "budget-1"})
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
@ -2588,3 +2588,303 @@ def test_ambiguous_commit_replay_does_not_erase_newly_accrued_spend(
assert client.key_spend == expected_spend
assert client.commit_attempts == expected_commits
assert client.reconnect_reasons == expected_reconnects
# ---------------------------------------------------------------------------
# Budget rollover (LIT-3085): overage beyond max_budget carries into the next
# window instead of being forgiven
# ---------------------------------------------------------------------------
@pytest.fixture
def rollover_enabled(monkeypatch):
import litellm
monkeypatch.setattr(litellm, "budget_rollover", True)
@pytest.mark.parametrize(
"run_phase, table, id_field, id_value, row_factory",
[
(
lambda job: job.reset_budget_for_litellm_keys(),
"key",
"token",
"tok-roll",
lambda now: type(
"Key",
(),
{
"spend": 150.0,
"max_budget": 100.0,
"budget_duration": "1d",
"budget_reset_at": now,
"token": "tok-roll",
},
),
),
(
lambda job: job.reset_budget_for_litellm_users(),
"user",
"user_id",
"user-roll",
lambda now: type(
"User",
(),
{
"spend": 150.0,
"max_budget": 100.0,
"budget_duration": "30d",
"budget_reset_at": now,
"user_id": "user-roll",
},
),
),
(
lambda job: job.reset_budget_for_litellm_teams(),
"team",
"team_id",
"team-roll",
lambda now: type(
"Team",
(),
{
"spend": 150.0,
"max_budget": 100.0,
"budget_duration": "1mo",
"budget_reset_at": now,
"team_id": "team-roll",
},
),
),
],
)
def test_direct_reset_carries_overage_when_rollover_enabled(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch, run_phase, table, id_field, id_value, row_factory
):
"""spend=150 against max_budget=100 must decrement by the cap (leaving 50)
rather than zero the row, and the spend counter must be seeded with 50."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data[table] = [row_factory(now)]
asyncio.run(run_phase(reset_budget_job))
writes = _batch_writes(mock_prisma_client, table)
assert len(writes) == 1
assert writes[0]["where"] == {id_field: id_value}
assert writes[0]["data"]["spend"] == {"decrement": 100.0}
assert writes[0]["data"]["budget_reset_at"] > now
counter_prefix = {"key": "spend:key", "user": "spend:user", "team": "spend:team"}[table]
counter_cache.in_memory_cache.set_cache.assert_any_call(key=f"{counter_prefix}:{id_value}", value=50.0, ttl=60)
def test_direct_reset_zeroes_under_budget_row_even_with_rollover(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["key"] = [
type(
"Key",
(),
{"spend": 40.0, "max_budget": 100.0, "budget_duration": "1d", "budget_reset_at": now, "token": "tok-under"},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == 0
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:tok-under", value=0.0, ttl=60)
def test_direct_reset_zeroes_row_without_max_budget_even_with_rollover(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
"""No cap means nothing to carry against: reset to zero as before."""
_make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["key"] = [
type(
"Key",
(),
{"spend": 150.0, "max_budget": None, "budget_duration": "1d", "budget_reset_at": now, "token": "tok-nocap"},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == 0
def test_budget_cascade_carries_overage_per_tier_when_rollover_enabled(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
"""A team member 5 over the tier cap keeps a spend of 5 in the next window:
the cascade decrements over-cap rows by the cap, zeroes the rest, and seeds
the spend counter with the carried amount."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
budget = _budget_row(budget_id="budget-roll", budget_duration="7d", max_budget=10.0)
mock_prisma_client.data["budget"] = [budget]
membership = type(
"Membership",
(),
{"user_id": "member-1", "team_id": "team-1", "spend": 15.0, "budget_id": "budget-roll"},
)
mock_prisma_client.db.litellm_teammembership.set_find_many_results([membership])
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
membership_writes = _batch_writes(mock_prisma_client, "team_membership")
assert {
"table": "team_membership",
"op": "update_many",
"where": {"budget_id": "budget-roll", "spend": {"gt": 10.0}},
"data": {"spend": {"decrement": 10.0}},
} in membership_writes
assert {
"table": "team_membership",
"op": "update_many",
"where": {"budget_id": "budget-roll", "spend": {"gt": 0, "lte": 10.0}},
"data": {"spend": 0},
} in membership_writes
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team_member:member-1:team-1", value=5.0, ttl=60)
def test_budget_cascade_carries_enduser_overage_when_rollover_enabled(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
_make_counter_invalidation_job(monkeypatch)
budget = _budget_row(budget_id="budget-roll", budget_duration="1d", max_budget=10.0)
mock_prisma_client.data["budget"] = [budget]
mock_prisma_client.data["enduser"] = [
type(
"EndUser",
(),
{"spend": 15.0, "litellm_budget_table": budget, "user_id": "enduser-roll", "budget_id": "budget-roll"},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
enduser_writes = _batch_writes(mock_prisma_client, "enduser")
assert {
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"gt": 10.0}},
"data": {"spend": {"decrement": 10.0}},
} in enduser_writes
assert {
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"lte": 10.0}},
"data": {"spend": 0},
} in enduser_writes
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."""
for write in writes:
condition = write["where"].get("spend")
if isinstance(condition, dict):
if "gt" in condition and not spend > condition["gt"]:
continue
if "lte" in condition and not spend <= condition["lte"]:
continue
payload = write["data"]["spend"]
spend = payload if not isinstance(payload, dict) else spend - payload["decrement"]
return spend
@pytest.mark.parametrize("table", ["team_membership", "enduser"])
def test_cascade_rollover_writes_survive_sequential_execution(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch, table
):
"""The statements run one after another inside a transaction, so a
decrement-then-zero order would re-match the decremented row (now in the
0..cap range) and erase the carried spend. Replaying the writes in queue
order must leave the overage, for any spend between cap and twice the cap."""
_make_counter_invalidation_job(monkeypatch)
budget = _budget_row(budget_id="budget-roll", budget_duration="7d", max_budget=10.0)
mock_prisma_client.data["budget"] = [budget]
membership = type(
"Membership",
(),
{"user_id": "member-1", "team_id": "team-1", "spend": 15.0, "budget_id": "budget-roll"},
)
mock_prisma_client.db.litellm_teammembership.set_find_many_results([membership])
mock_prisma_client.data["enduser"] = [
type(
"EndUser",
(),
{"spend": 15.0, "litellm_budget_table": budget, "user_id": "enduser-roll", "budget_id": "budget-roll"},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
writes = _batch_writes(mock_prisma_client, table)
assert _replay_spend_writes(writes, 15.0) == 5.0
assert _replay_spend_writes(writes, 8.0) == 0
assert _replay_spend_writes(writes, 25.0) == 15.0
def test_budget_cascade_zeroes_everything_when_rollover_disabled(reset_budget_job, mock_prisma_client, monkeypatch):
"""Control: with the flag off the cascade keeps the plain zeroing writes."""
_make_counter_invalidation_job(monkeypatch)
budget = _budget_row(budget_id="budget-off", budget_duration="7d", max_budget=10.0)
mock_prisma_client.data["budget"] = [budget]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
membership_writes = _batch_writes(mock_prisma_client, "team_membership")
assert membership_writes == [
{
"table": "team_membership",
"op": "update_many",
"where": {"budget_id": {"in": ["budget-off"]}},
"data": {"spend": 0},
}
]
def test_window_reset_carries_counter_overage_when_rollover_enabled(rollover_enabled, monkeypatch):
"""A per-window counter at 130 against a 100 cap restarts the window at 30."""
now = datetime.utcnow()
expired = (now - timedelta(minutes=5)).isoformat() + "Z"
key_rows = [
{
"token": "sk-roll",
"budget_limits": [{"budget_duration": "1d", "reset_at": expired, "max_budget": 100.0}],
}
]
job, prisma_client, spend_counter_cache = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
spend_counter_cache.async_get_cache = AsyncMock(return_value=130.0)
asyncio.run(job.reset_budget_windows())
prisma_client.db.litellm_verificationtoken.update.assert_awaited_once()
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-roll:window:1d", value=30.0)
def test_window_reset_zeroes_counter_when_rollover_disabled(monkeypatch):
now = datetime.utcnow()
expired = (now - timedelta(minutes=5)).isoformat() + "Z"
key_rows = [
{
"token": "sk-off",
"budget_limits": [{"budget_duration": "1d", "reset_at": expired, "max_budget": 100.0}],
}
]
job, prisma_client, spend_counter_cache = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
spend_counter_cache.async_get_cache = AsyncMock(return_value=130.0)
asyncio.run(job.reset_budget_windows())
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-off:window:1d", value=0.0)
spend_counter_cache.async_get_cache.assert_not_awaited()