mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(budgets): page end-user cache invalidation after a budget reset
The budget-tier reset read every customer linked to an expiring tier into one result set before the write, then invalidated their caches one key at a time. Both of those scale with the customer count, so a large enough deployment can OOM the proxy pod on the read, and the tail of the population sits on a stale spend counter while the per-key invalidations drain PR #40639 moved the reset write itself to a link-based UPDATE, so that pre-commit read no longer feeds the write. It only fed cache invalidation and the service-logging counts, which means it can move after the commit. This replaces it with a keyset walk over litellm_endusertable ordered by user_id, taking RESET_BUDGET_JOB_BATCH_SIZE rows per page, the same shape _reset_windows_for_source already uses, with no per-run page cap for the same reason that walk has none: the cursor cannot survive the run, so a cap would restart at the first customer on every tick and never reach the tail Each page's counter and cache keys now go out as one batched delete through a new DualCache.async_delete_cache_keys, which drops the in-memory entries and chunks the Redis DELETE at DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE num_endusers_found and num_endusers_updated now report the customers whose caches were invalidated after the commit rather than the rows read before it, so both read 0 when the cascade write fails
This commit is contained in:
parent
b04d530ecf
commit
7ba47a5b6e
5 changed files with 262 additions and 73 deletions
|
|
@ -521,6 +521,23 @@ class DualCache(BaseCache):
|
|||
if self.redis_cache is not None:
|
||||
await self.redis_cache.async_delete_cache(key)
|
||||
|
||||
async def async_delete_cache_keys(self, keys: Sequence[str]) -> None:
|
||||
"""Batch twin of ``async_delete_cache``: one Redis round trip per chunk
|
||||
instead of one per key.
|
||||
|
||||
Chunked because Redis takes the whole list as a single DELETE command,
|
||||
and a caller holding a population-sized list would otherwise build one
|
||||
command out of it.
|
||||
"""
|
||||
if not keys:
|
||||
return
|
||||
for key in keys:
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
if self.redis_cache is None:
|
||||
return
|
||||
for start in range(0, len(keys), DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE):
|
||||
await self.redis_cache.delete_cache_keys(keys[start : start + DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE])
|
||||
|
||||
async def async_get_ttl(self, key: str) -> int | None:
|
||||
"""
|
||||
Get the remaining TTL of a key in in-memory cache or redis
|
||||
|
|
|
|||
|
|
@ -26,7 +26,6 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
|||
from litellm.proxy._types import (
|
||||
DB_RETRY_SAFE_ERROR_TYPES,
|
||||
LiteLLM_BudgetTableFull,
|
||||
LiteLLM_EndUserTable,
|
||||
Litellm_EntityType,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -193,13 +192,6 @@ 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({}),
|
||||
|
|
@ -207,6 +199,21 @@ def _budget_link_where(
|
|||
return {"budget_id": {"in": list(budget_ids)}, **extra}
|
||||
|
||||
|
||||
def _enduser_invalidation_where(budget_ids: Sequence[str]) -> dict[str, object]:
|
||||
"""Customers whose cached spend a committed reset of these tiers invalidated.
|
||||
|
||||
Mirrors ``_queue_enduser_resets``: the link, plus the NULL-budget_id rows
|
||||
that ride the default tier when that tier is one of the expiring ones. The
|
||||
write's ``spend > 0`` filter has no twin here because the commit already
|
||||
zeroed those rows, so post-commit it would match nobody.
|
||||
"""
|
||||
linked: Final = _budget_link_where(budget_ids)
|
||||
default_budget_id: Final = litellm.max_end_user_budget_id
|
||||
if default_budget_id is None or default_budget_id not in budget_ids:
|
||||
return linked
|
||||
return {"OR": [linked, {"budget_id": None}]} # mutable-ok: prisma where filter must be a dict
|
||||
|
||||
|
||||
def _queue_budget_linked_resets(
|
||||
writes: LinkedSpendResetWrites,
|
||||
cascade: "_BudgetCascade",
|
||||
|
|
@ -265,7 +272,6 @@ class _BudgetCascade:
|
|||
budgets: tuple[LiteLLM_BudgetTableFull, ...] = ()
|
||||
budget_ids: tuple[str, ...] = ()
|
||||
budget_resets: tuple[tuple[str, datetime], ...] = ()
|
||||
endusers: tuple[_EndUserRow, ...] = ()
|
||||
counter_resets: tuple[tuple[str, float], ...] = ()
|
||||
cache_keys: tuple[str, ...] = ()
|
||||
rollover_caps: Mapping[str, float] = field(default_factory=lambda: MappingProxyType({}))
|
||||
|
|
@ -275,6 +281,7 @@ class _BudgetCascade:
|
|||
class _BudgetCascadeCommitted:
|
||||
cascade: _BudgetCascade
|
||||
advanced: int
|
||||
endusers_invalidated: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -416,10 +423,10 @@ _WINDOW_SOURCES: Final[tuple[_WindowSource, ...]] = (
|
|||
)
|
||||
|
||||
|
||||
def _budget_cascade_event_metadata(cascade: _BudgetCascade) -> dict[str, object]:
|
||||
def _budget_cascade_event_metadata(cascade: _BudgetCascade, endusers_invalidated: int = 0) -> dict[str, object]:
|
||||
return {
|
||||
"num_budgets_found": len(cascade.budgets),
|
||||
"num_endusers_found": len(cascade.endusers),
|
||||
"num_endusers_found": endusers_invalidated,
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -593,6 +600,32 @@ class ResetBudgetJob:
|
|||
e,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _invalidate_caches(counter_keys: Sequence[str], cache_keys: Sequence[str]) -> None:
|
||||
"""Batch twin of ``_invalidate_spend_counter`` and
|
||||
``_invalidate_user_api_key_cache_entry``, carrying the same
|
||||
after-the-commit requirement as both.
|
||||
|
||||
One round trip per chunk rather than one per key: a tier's dependent
|
||||
population is unbounded, and awaiting each key in turn makes the last
|
||||
dependent wait out every dependent ahead of it.
|
||||
"""
|
||||
if not counter_keys and not cache_keys:
|
||||
return
|
||||
try:
|
||||
from litellm.proxy.proxy_server import spend_counter_cache, user_api_key_cache
|
||||
|
||||
await spend_counter_cache.async_delete_cache_keys(counter_keys)
|
||||
await user_api_key_cache.async_delete_cache_keys(cache_keys)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to invalidate %d spend counters and %d user_api_key_cache entries: %s. "
|
||||
"Budgets may be over-enforced until the counters expire.",
|
||||
len(counter_keys),
|
||||
len(cache_keys),
|
||||
e,
|
||||
)
|
||||
|
||||
async def _fetch_linked_rows(
|
||||
self,
|
||||
table: SpendLinkedTable[_RowT],
|
||||
|
|
@ -612,18 +645,54 @@ class ResetBudgetJob:
|
|||
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
|
||||
return ()
|
||||
|
||||
async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]:
|
||||
linked: Final[Sequence[_EndUserRow] | None] = await self._with_db_retry(
|
||||
lambda: self.prisma_client.get_data(
|
||||
table_name="enduser",
|
||||
query_type="find_all",
|
||||
budget_id_list=list(budget_ids),
|
||||
),
|
||||
reason="reset_budget_read_endusers_failure",
|
||||
)
|
||||
if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids:
|
||||
return tuple(linked or ())
|
||||
return (*(linked or ()), *await self._get_endusers_with_no_budget_id())
|
||||
async def _invalidate_enduser_caches(self, budget_ids: Sequence[str]) -> int:
|
||||
"""Drop the cached spend of every customer the committed tier reset zeroed.
|
||||
|
||||
Walked a page at a time with a keyset cursor, for the same reason
|
||||
``_reset_windows_for_source`` is: the customers sharing one tier are
|
||||
unbounded, so reading them into one result set puts a
|
||||
customer-count-sized list in the proxy's heap on every tick, and a
|
||||
deployment large enough turns that into an OOM rather than a slow tick.
|
||||
|
||||
No per-run page cap, also for that walk's reason: the position cannot
|
||||
survive the run, so a cap would restart at the first customer every tick
|
||||
and never reach the tail. The cursor strictly advances, so this
|
||||
terminates on its own.
|
||||
"""
|
||||
if not budget_ids:
|
||||
return 0
|
||||
where: Final = _enduser_invalidation_where(budget_ids)
|
||||
cursor = ""
|
||||
invalidated = 0
|
||||
while True:
|
||||
rows = await self._fetch_enduser_page(where=where, cursor=cursor)
|
||||
if not rows:
|
||||
return invalidated
|
||||
await self._invalidate_caches(
|
||||
counter_keys=tuple(_enduser_counter_key(row) for row in rows),
|
||||
cache_keys=tuple(key for row in rows for key in _enduser_cache_keys(row)),
|
||||
)
|
||||
invalidated += len(rows)
|
||||
if len(rows) < RESET_BUDGET_JOB_BATCH_SIZE:
|
||||
return invalidated
|
||||
cursor = rows[-1].user_id
|
||||
|
||||
async def _fetch_enduser_page(self, where: Mapping[str, object], cursor: str) -> tuple[_EndUserRow, ...]:
|
||||
"""One keyset page of customers, ordered by primary key so the cursor never repeats a row."""
|
||||
try:
|
||||
return tuple(
|
||||
await self._with_db_retry(
|
||||
lambda: EndUserRepository(self.prisma_client).table.find_many(
|
||||
where={**where, "user_id": {"gt": cursor}}, # mutable-ok: prisma where filter must be a dict
|
||||
order={"user_id": "asc"}, # mutable-ok: prisma order filter must be a dict
|
||||
take=RESET_BUDGET_JOB_BATCH_SIZE,
|
||||
),
|
||||
reason="reset_budget_read_endusers_failure",
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to fetch end users for cache invalidation: %s", e)
|
||||
return ()
|
||||
|
||||
async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade:
|
||||
"""Resolve every row the expiring budget tiers gate, before any write.
|
||||
|
|
@ -670,7 +739,6 @@ 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,
|
||||
|
|
@ -682,7 +750,6 @@ class ResetBudgetJob:
|
|||
for b in budgets_to_reset
|
||||
if b.budget_id is not None and b.budget_duration is not None
|
||||
),
|
||||
endusers=endusers,
|
||||
counter_resets=(
|
||||
*(
|
||||
(_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps))
|
||||
|
|
@ -695,7 +762,6 @@ 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=(
|
||||
|
|
@ -704,7 +770,6 @@ 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)),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -736,10 +801,10 @@ class ResetBudgetJob:
|
|||
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_resets:
|
||||
await self._invalidate_spend_counter(counter_key)
|
||||
for cache_key in cascade.cache_keys:
|
||||
await self._invalidate_user_api_key_cache_entry(cache_key)
|
||||
await self._invalidate_caches(
|
||||
counter_keys=tuple(counter_key for counter_key, _ in cascade.counter_resets),
|
||||
cache_keys=cascade.cache_keys,
|
||||
)
|
||||
|
||||
async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed:
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
|
|
@ -769,6 +834,7 @@ class ResetBudgetJob:
|
|||
(reset_at for _, reset_at in cascade.budget_resets),
|
||||
cutoff=datetime.now(timezone.utc),
|
||||
),
|
||||
endusers_invalidated=await self._invalidate_enduser_caches(cascade.budget_ids),
|
||||
)
|
||||
|
||||
async def reset_budget_for_litellm_budget_table(self) -> None:
|
||||
|
|
@ -788,7 +854,7 @@ class ResetBudgetJob:
|
|||
end_time: Final = time.time()
|
||||
|
||||
match outcome:
|
||||
case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced):
|
||||
case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced, endusers_invalidated=endusers_invalidated):
|
||||
asyncio.create_task(
|
||||
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
|
||||
service=ServiceTypes.RESET_BUDGET_JOB,
|
||||
|
|
@ -797,8 +863,8 @@ class ResetBudgetJob:
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
event_metadata={
|
||||
**_budget_cascade_event_metadata(cascade),
|
||||
"num_endusers_updated": len(cascade.endusers),
|
||||
**_budget_cascade_event_metadata(cascade, endusers_invalidated),
|
||||
"num_endusers_updated": endusers_invalidated,
|
||||
"num_endusers_failed": 0,
|
||||
},
|
||||
)
|
||||
|
|
@ -827,27 +893,6 @@ class ResetBudgetJob:
|
|||
case _:
|
||||
assert_never(outcome)
|
||||
|
||||
async def _get_endusers_with_no_budget_id(
|
||||
self,
|
||||
) -> list[LiteLLM_EndUserTable]:
|
||||
"""
|
||||
Fetch end users that have no explicit budget_id set (NULL) and have
|
||||
accumulated spend > 0. These are implicitly-created end users that
|
||||
rely on the default budget (litellm.max_end_user_budget_id) applied
|
||||
in-memory during auth checks.
|
||||
"""
|
||||
table: Final = EndUserRepository(self.prisma_client).table
|
||||
rows: Final = await self._with_db_retry(
|
||||
lambda: table.find_many(
|
||||
where={
|
||||
"budget_id": None,
|
||||
"spend": {"gt": 0},
|
||||
},
|
||||
),
|
||||
reason="reset_budget_read_endusers_without_budget_id_failure",
|
||||
)
|
||||
return [LiteLLM_EndUserTable.model_validate(row.model_dump()) for row in rows]
|
||||
|
||||
async def _write_key_reset_updates(self, updated_keys: Sequence[_RowReset[LiteLLM_VerificationToken]]) -> None:
|
||||
"""
|
||||
Write per-row {spend, budget_reset_at} updates for keys.
|
||||
|
|
|
|||
|
|
@ -102,21 +102,24 @@ def _wire_batcher_for_test(prisma_client, fail_commit=False):
|
|||
return batch_calls
|
||||
|
||||
|
||||
def _wire_cascade_reads_for_test(prisma_client):
|
||||
def _wire_cascade_reads_for_test(prisma_client, endusers=()):
|
||||
"""
|
||||
The budget tier's cascade reads the rows it is about to zero, so their
|
||||
spend counters can be invalidated after the commit. Give each of those
|
||||
tables an awaitable find_many so the reads resolve instead of falling into
|
||||
the job's warn-and-continue path.
|
||||
|
||||
End users are read by the post-commit invalidation walk rather than by
|
||||
``get_data``, so callers that care about customers pass them here.
|
||||
"""
|
||||
for table in (
|
||||
"litellm_teammembership",
|
||||
"litellm_verificationtoken",
|
||||
"litellm_organizationtable",
|
||||
"litellm_tagtable",
|
||||
"litellm_endusertable",
|
||||
):
|
||||
getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[])
|
||||
prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=list(endusers))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -556,7 +559,7 @@ async def test_reset_budget_continues_other_categories_on_failure():
|
|||
**{u["user_id"]: u["spend"] for u in [user2]},
|
||||
**{t["team_id"]: t["spend"] for t in [team1, team2]},
|
||||
}
|
||||
_wire_cascade_reads_for_test(prisma_client)
|
||||
_wire_cascade_reads_for_test(prisma_client, endusers=[enduser1])
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -607,7 +610,10 @@ async def test_reset_budget_continues_other_categories_on_failure():
|
|||
called_tables = {
|
||||
call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list
|
||||
}
|
||||
assert called_tables == {"key", "user", "team", "budget", "enduser"}
|
||||
assert called_tables == {"key", "user", "team", "budget"}
|
||||
# Customers are not part of that set: the cascade zeroes them by budget link
|
||||
# and reads them only afterwards, to invalidate their cached spend.
|
||||
prisma_client.db.litellm_endusertable.find_many.assert_awaited()
|
||||
|
||||
# Every category writes through the batch path now, so update_data is unused.
|
||||
prisma_client.update_data.assert_not_awaited()
|
||||
|
|
@ -1029,7 +1035,7 @@ async def test_service_logger_endusers_success():
|
|||
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
batch_calls = _wire_batcher_for_test(prisma_client)
|
||||
_wire_cascade_reads_for_test(prisma_client)
|
||||
_wire_cascade_reads_for_test(prisma_client, endusers=endusers)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -1094,7 +1100,7 @@ async def test_service_logger_endusers_failure():
|
|||
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
_wire_batcher_for_test(prisma_client, fail_commit=True)
|
||||
_wire_cascade_reads_for_test(prisma_client)
|
||||
_wire_cascade_reads_for_test(prisma_client, endusers=endusers)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -1121,7 +1127,9 @@ async def test_service_logger_endusers_failure():
|
|||
) = proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args
|
||||
event_metadata = kwargs.get("event_metadata", {})
|
||||
assert event_metadata.get("num_budgets_found") == len(budgets)
|
||||
assert event_metadata.get("num_endusers_found") == len(endusers)
|
||||
# Customers are read by the post-commit invalidation walk, which a failed
|
||||
# commit never reaches, so a failure reports none touched.
|
||||
assert event_metadata.get("num_endusers_found") == 0
|
||||
assert "endusers_found" not in event_metadata
|
||||
assert "budgets_found" not in event_metadata
|
||||
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync
|
||||
|
|
@ -759,3 +760,34 @@ async def test_redis_timeouts_falling_back_to_memory_log_once_per_interval(caplo
|
|||
" (199 more Redis timeouts since the previous Redis timeout line were logged at DEBUG)",
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_keys_drops_memory_and_chunks_redis():
|
||||
"""Batch delete clears both layers, and chunks Redis so one caller's large
|
||||
key list cannot become a single oversized DELETE command."""
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.delete_cache_keys = AsyncMock()
|
||||
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache)
|
||||
keys = [f"key-{i}" for i in range(DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE + 7)]
|
||||
for key in keys:
|
||||
dual_cache.in_memory_cache.set_cache(key=key, value=1)
|
||||
|
||||
await dual_cache.async_delete_cache_keys(keys)
|
||||
|
||||
assert all(dual_cache.in_memory_cache.get_cache(key=key) is None for key in keys)
|
||||
sent = [call.args[0] for call in redis_cache.delete_cache_keys.await_args_list]
|
||||
assert [len(chunk) for chunk in sent] == [DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE, 7]
|
||||
assert [key for chunk in sent for key in chunk] == keys
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_keys_on_empty_list_touches_no_backend():
|
||||
"""An empty page must not reach Redis: DELETE with no arguments is an error."""
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.delete_cache_keys = AsyncMock()
|
||||
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache)
|
||||
|
||||
await dual_cache.async_delete_cache_keys([])
|
||||
|
||||
redis_cache.delete_cache_keys.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -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, Final, List
|
||||
from typing import Any, Dict, Final, List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
|
|
@ -16,6 +16,7 @@ from litellm.proxy._types import LiteLLM_VerificationToken
|
|||
from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module
|
||||
from litellm.constants import (
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
RESET_BUDGET_JOB_BATCH_SIZE,
|
||||
RESET_BUDGET_JOB_LOCK_TTL_SECONDS,
|
||||
RESET_BUDGET_JOB_NAME,
|
||||
)
|
||||
|
|
@ -35,9 +36,24 @@ class MockTable:
|
|||
def set_find_many_results(self, results: List[Any]):
|
||||
self._find_many_results = results
|
||||
|
||||
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 find_many(
|
||||
self,
|
||||
where: Dict[str, Any],
|
||||
order: Optional[Dict[str, str]] = None,
|
||||
take: Optional[int] = None,
|
||||
) -> List[Any]:
|
||||
"""Replays canned rows, honouring the keyset cursor + ``take`` a paged
|
||||
caller relies on: without that a paged walk never advances and the
|
||||
test would hang instead of failing."""
|
||||
paging = {k: v for k, v in (("order", order), ("take", take)) if v is not None}
|
||||
self.find_many_calls.append({"where": where, **paging})
|
||||
rows = list(self._find_many_results)
|
||||
for field, condition in where.items():
|
||||
if isinstance(condition, dict) and "gt" in condition and field != "spend":
|
||||
rows = [row for row in rows if getattr(row, field, "") > condition["gt"]]
|
||||
for field, direction in (order or {}).items():
|
||||
rows.sort(key=lambda row: getattr(row, field, ""), reverse=direction == "desc")
|
||||
return rows[:take] if take is not None else rows
|
||||
|
||||
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})
|
||||
|
|
@ -784,10 +800,16 @@ def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock
|
|||
},
|
||||
]
|
||||
|
||||
# Verify find_many was called to fetch NULL-budget-id end users
|
||||
# The post-commit invalidation walk covers both branches, so implicitly
|
||||
# created customers on the default tier get their cached spend dropped too,
|
||||
# and it is paged rather than reading the whole customer population.
|
||||
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
|
||||
assert len(find_many_calls) == 1
|
||||
assert find_many_calls[0]["where"] == {"budget_id": None, "spend": {"gt": 0}}
|
||||
assert find_many_calls[0]["where"]["OR"] == [
|
||||
{"budget_id": {"in": [default_budget_id]}},
|
||||
{"budget_id": None},
|
||||
]
|
||||
assert find_many_calls[0]["take"] == RESET_BUDGET_JOB_BATCH_SIZE
|
||||
|
||||
litellm.max_end_user_budget_id = None
|
||||
|
||||
|
|
@ -818,9 +840,12 @@ def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured(
|
|||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
# Should NOT have queried for NULL-budget-id end users
|
||||
# The invalidation walk must not reach for NULL-budget-id customers: they
|
||||
# ride a default tier that is not expiring, so their spend stays put.
|
||||
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
|
||||
assert len(find_many_calls) == 0
|
||||
assert [call["where"] for call in find_many_calls] == [
|
||||
{"budget_id": {"in": ["some-budget"]}, "user_id": {"gt": ""}}
|
||||
]
|
||||
|
||||
litellm.max_end_user_budget_id = None
|
||||
|
||||
|
|
@ -855,9 +880,12 @@ def test_reset_budget_skips_null_budget_id_endusers_when_default_not_in_reset_li
|
|||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
# Should NOT have queried for NULL-budget-id end users
|
||||
# The invalidation walk must not reach for NULL-budget-id customers: they
|
||||
# ride a default tier that is not expiring, so their spend stays put.
|
||||
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
|
||||
assert len(find_many_calls) == 0
|
||||
assert [call["where"] for call in find_many_calls] == [
|
||||
{"budget_id": {"in": ["other-budget"]}, "user_id": {"gt": ""}}
|
||||
]
|
||||
|
||||
litellm.max_end_user_budget_id = None
|
||||
|
||||
|
|
@ -1235,6 +1263,21 @@ def _make_counter_invalidation_job(monkeypatch):
|
|||
user_api_key_cache = MagicMock()
|
||||
user_api_key_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
# Batch deletes fan out to the same per-key calls the real DualCache makes,
|
||||
# so an assertion reads "this key was invalidated" whether the caller went
|
||||
# one key at a time or a page at a time.
|
||||
async def _delete_counter_keys(keys):
|
||||
for key in keys:
|
||||
spend_counter_cache.in_memory_cache.delete_cache(key=key)
|
||||
await spend_counter_cache.redis_cache.async_delete_cache(key=key)
|
||||
|
||||
async def _delete_management_keys(keys):
|
||||
for key in keys:
|
||||
await user_api_key_cache.async_delete_cache(key=key)
|
||||
|
||||
spend_counter_cache.async_delete_cache_keys = AsyncMock(side_effect=_delete_counter_keys)
|
||||
user_api_key_cache.async_delete_cache_keys = AsyncMock(side_effect=_delete_management_keys)
|
||||
|
||||
fake_module = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_module.spend_counter_cache = spend_counter_cache
|
||||
fake_module.user_api_key_cache = user_api_key_cache
|
||||
|
|
@ -1569,7 +1612,7 @@ def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_j
|
|||
"user_id": "customer-42",
|
||||
},
|
||||
)
|
||||
mock_prisma_client.data["enduser"] = [test_enduser]
|
||||
mock_prisma_client.db.litellm_endusertable.set_find_many_results([test_enduser])
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
|
|
@ -1579,6 +1622,50 @@ def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_j
|
|||
assert "end_user_id:customer-42" in deleted
|
||||
|
||||
|
||||
def test_enduser_invalidation_is_paged_and_batched(reset_budget_job, mock_prisma_client, monkeypatch):
|
||||
"""The post-commit invalidation walk stays bounded in memory and in round trips.
|
||||
|
||||
Reading every customer on an expiring tier into one result set puts a
|
||||
customer-count-sized list in the proxy's heap on every tick, which is an OOM
|
||||
on a large enough deployment rather than a slow tick. Awaiting one cache call
|
||||
per customer makes the last customer wait out every customer ahead of it.
|
||||
Both regress silently, so pin the page size, the strictly advancing cursor,
|
||||
and one batched call per page.
|
||||
"""
|
||||
counter_cache: Final = _make_counter_invalidation_job(monkeypatch)
|
||||
mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-1")]
|
||||
population: Final = RESET_BUDGET_JOB_BATCH_SIZE * 2 + 3
|
||||
mock_prisma_client.db.litellm_endusertable.set_find_many_results(
|
||||
[
|
||||
type("EndUser", (), {"user_id": f"cust-{i:06d}", "spend": 5.0, "budget_id": "budget-1"})
|
||||
for i in range(population)
|
||||
]
|
||||
)
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
|
||||
|
||||
reads: Final = mock_prisma_client.db.litellm_endusertable.find_many_calls
|
||||
assert [read["take"] for read in reads] == [RESET_BUDGET_JOB_BATCH_SIZE] * 3
|
||||
assert [read["where"]["user_id"]["gt"] for read in reads] == [
|
||||
"",
|
||||
f"cust-{RESET_BUDGET_JOB_BATCH_SIZE - 1:06d}",
|
||||
f"cust-{RESET_BUDGET_JOB_BATCH_SIZE * 2 - 1:06d}",
|
||||
]
|
||||
|
||||
assert counter_cache.async_delete_cache_keys.await_count == 3
|
||||
assert counter_cache.user_api_key_cache.async_delete_cache_keys.await_count == 3
|
||||
counter_cache.async_delete_cache.assert_not_called()
|
||||
|
||||
invalidated: Final = {
|
||||
key for call in counter_cache.async_delete_cache_keys.await_args_list for key in call.args[0]
|
||||
}
|
||||
assert invalidated == {f"spend:end_user:cust-{i:06d}" for i in range(population)}
|
||||
evicted: Final = {
|
||||
key for call in counter_cache.user_api_key_cache.async_delete_cache_keys.await_args_list for key in call.args[0]
|
||||
}
|
||||
assert evicted == {f"end_user_id:cust-{i:06d}" for i in range(population)}
|
||||
|
||||
|
||||
|
||||
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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue