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:
ryan-crabbe-berri 2026-09-16 16:28:16 -07:00
parent cdf0142f4a
commit 14fbd623d7
3 changed files with 84 additions and 35 deletions

View file

@ -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."""

View file

@ -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()

View file

@ -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))