mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): clear both cache partitions and carry the walk position as a value
UserApiKeyCache's batch delete ran the two partitions in sequence, so a Redis failure on the hashed token partition returned before the ordinary management keys were touched. Both partitions are attempted now and the first failure is re-raised for the caller to report. The customer walk kept its position in two locals it reassigned each page. It now mirrors the window walk in the same file: a page helper returns where the walk goes next, and the driver rebinds one value.
This commit is contained in:
parent
cdf0142f4a
commit
14fbd623d7
3 changed files with 84 additions and 35 deletions
|
|
@ -278,22 +278,24 @@ class _BudgetCascade:
|
|||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _EndUserInvalidation:
|
||||
"""How far the post-commit customer walk got, and whether a failed page read
|
||||
cut it short of the tail."""
|
||||
class _EndUserWalk:
|
||||
"""Where the post-commit customer walk stands: the keyset cursor its next
|
||||
page resumes from, None once there is no next page, how many customers it
|
||||
has reached, and whether a failed page read cut it short of the tail."""
|
||||
|
||||
cursor: str | None = ""
|
||||
invalidated: int = 0
|
||||
truncated: bool = False
|
||||
|
||||
|
||||
_NO_ENDUSERS_INVALIDATED: Final = _EndUserInvalidation()
|
||||
_ENDUSER_WALK_DONE: Final = _EndUserWalk(cursor=None)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BudgetCascadeCommitted:
|
||||
cascade: _BudgetCascade
|
||||
advanced: int
|
||||
endusers: _EndUserInvalidation
|
||||
endusers: _EndUserWalk
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -440,7 +442,7 @@ _WINDOW_SOURCES: Final[tuple[_WindowSource, ...]] = (
|
|||
|
||||
|
||||
def _budget_cascade_event_metadata(
|
||||
cascade: _BudgetCascade, endusers: _EndUserInvalidation = _NO_ENDUSERS_INVALIDATED
|
||||
cascade: _BudgetCascade, endusers: _EndUserWalk = _ENDUSER_WALK_DONE
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"num_budgets_found": len(cascade.budgets),
|
||||
|
|
@ -674,7 +676,7 @@ class ResetBudgetJob:
|
|||
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
|
||||
return ()
|
||||
|
||||
async def _invalidate_enduser_caches(self, budget_ids: Sequence[str]) -> _EndUserInvalidation:
|
||||
async def _invalidate_enduser_caches(self, budget_ids: Sequence[str]) -> _EndUserWalk:
|
||||
"""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
|
||||
|
|
@ -694,32 +696,38 @@ class ResetBudgetJob:
|
|||
passing the part it managed off as the whole.
|
||||
"""
|
||||
if not budget_ids:
|
||||
return _NO_ENDUSERS_INVALIDATED
|
||||
return _ENDUSER_WALK_DONE
|
||||
where: Final = _enduser_invalidation_where(budget_ids)
|
||||
cursor = ""
|
||||
invalidated = 0
|
||||
while True:
|
||||
try:
|
||||
rows = await self._fetch_enduser_page(where=where, cursor=cursor)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to fetch end users for cache invalidation after %s customers (cursor %r): %s. "
|
||||
"The customers past that page keep their cached spend until it expires.",
|
||||
invalidated,
|
||||
cursor,
|
||||
e,
|
||||
)
|
||||
return _EndUserInvalidation(invalidated=invalidated, truncated=True)
|
||||
if not rows:
|
||||
return _EndUserInvalidation(invalidated=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)),
|
||||
walk = _EndUserWalk()
|
||||
while walk.cursor is not None:
|
||||
walk = await self._invalidate_enduser_page(where=where, cursor=walk.cursor, reached=walk.invalidated)
|
||||
return walk
|
||||
|
||||
async def _invalidate_enduser_page(
|
||||
self, where: Mapping[str, object], cursor: str, reached: int
|
||||
) -> _EndUserWalk:
|
||||
"""Invalidate one page of customers and say where the walk goes next."""
|
||||
try:
|
||||
rows: Final = await self._fetch_enduser_page(where=where, cursor=cursor)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to fetch end users for cache invalidation after %s customers (cursor %r): %s. "
|
||||
"The customers past that page keep their cached spend until it expires.",
|
||||
reached,
|
||||
cursor,
|
||||
e,
|
||||
)
|
||||
invalidated += len(rows)
|
||||
if len(rows) < RESET_BUDGET_JOB_BATCH_SIZE:
|
||||
return _EndUserInvalidation(invalidated=invalidated)
|
||||
cursor = rows[-1].user_id
|
||||
return _EndUserWalk(cursor=None, invalidated=reached, truncated=True)
|
||||
if not rows:
|
||||
return _EndUserWalk(cursor=None, invalidated=reached)
|
||||
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)),
|
||||
)
|
||||
walked: Final = reached + len(rows)
|
||||
if len(rows) < RESET_BUDGET_JOB_BATCH_SIZE:
|
||||
return _EndUserWalk(cursor=None, invalidated=walked)
|
||||
return _EndUserWalk(cursor=rows[-1].user_id, invalidated=walked)
|
||||
|
||||
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."""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast, overload
|
||||
|
|
@ -222,12 +223,24 @@ class UserApiKeyCache(DualCache):
|
|||
await super().async_delete_cache(key)
|
||||
|
||||
async def async_delete_cache_keys(self, keys: Sequence[str]) -> None:
|
||||
"""Batch twin of ``async_delete_cache``, partitioned the way
|
||||
``async_set_cache_pipeline`` partitions its writes.
|
||||
|
||||
Both partitions are cleared even when one of them raises: a caller
|
||||
batching these has already committed the rows they cache, so a partition
|
||||
left holding pre-reset spend goes on being authorized against until the
|
||||
entry expires. The first failure is re-raised for the caller to report.
|
||||
"""
|
||||
key_object_keys: Final = tuple(key for key in keys if is_user_key_cache_key(key))
|
||||
other_keys: Final = tuple(key for key in keys if not is_user_key_cache_key(key))
|
||||
if key_object_keys:
|
||||
await self.key_object_cache.async_delete_cache_keys(key_object_keys)
|
||||
if other_keys:
|
||||
await super().async_delete_cache_keys(other_keys)
|
||||
outcomes: Final = await asyncio.gather(
|
||||
self.key_object_cache.async_delete_cache_keys(key_object_keys),
|
||||
super().async_delete_cache_keys(other_keys),
|
||||
return_exceptions=True,
|
||||
)
|
||||
failed: Final = tuple(outcome for outcome in outcomes if isinstance(outcome, BaseException))
|
||||
if failed:
|
||||
raise failed[0]
|
||||
|
||||
def flush_cache(self) -> None:
|
||||
super().flush_cache()
|
||||
|
|
|
|||
|
|
@ -87,6 +87,15 @@ class FakeRedisCache(RedisCache):
|
|||
self._store.pop(key, None)
|
||||
|
||||
|
||||
class PartitionFailingRedisCache(FakeRedisCache):
|
||||
"""Fails the batch delete for the key-object partition and no other."""
|
||||
|
||||
async def delete_cache_keys(self, keys): # type: ignore[override]
|
||||
if any(is_user_key_cache_key(key) for key in keys):
|
||||
raise ConnectionError("redis unavailable")
|
||||
await super().delete_cache_keys(keys)
|
||||
|
||||
|
||||
def _make_key_obj(token: str = "tok") -> UserAPIKeyAuth:
|
||||
# Minimal object (UserAPIKeyAuth inherits token from base view).
|
||||
return UserAPIKeyAuth(token=token)
|
||||
|
|
@ -356,6 +365,25 @@ class TestUserKeyObjectPartition:
|
|||
assert await redis.async_get_cache(HASHED_TOKEN) is None
|
||||
assert await redis.async_get_cache(end_user_cache_key("u1")) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_delete_clears_the_other_partition_when_one_fails(self):
|
||||
"""One partition failing must not cost the other its deletions.
|
||||
|
||||
A caller batching these has already committed the rows they cache, so a
|
||||
partition that is skipped keeps authorizing against pre-reset spend until
|
||||
the entry expires. The failure is still raised for the caller to report.
|
||||
"""
|
||||
redis = PartitionFailingRedisCache()
|
||||
cache = UserApiKeyCache(redis_cache=redis)
|
||||
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
|
||||
await cache.async_set_cache(end_user_cache_key("u1"), {"user_id": "u1"})
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await cache.async_delete_cache_keys([HASHED_TOKEN, end_user_cache_key("u1")])
|
||||
|
||||
assert await cache.async_get_cache(end_user_cache_key("u1")) is None
|
||||
assert await redis.async_get_cache(end_user_cache_key("u1")) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pipeline_write_routes_each_entry_to_its_partition(self):
|
||||
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue