fix(batch_enqueued_tokens): scope in-memory refunds to the granting worker

In-memory grants now record an owner token, and a refund only debits local
counters when the popping worker is the one that granted them, so a terminal
response handled elsewhere can no longer shrink another worker's unrelated
fallback reservations. A Redis-granted refund that fails no longer falls back
to decrementing local counters either: the leaked Redis increments expire
with the TTL and only tighten the allowance.
This commit is contained in:
mateo-berri 2026-08-19 16:36:27 -07:00
parent 50896f21b3
commit 4333d52813
2 changed files with 74 additions and 13 deletions

View file

@ -9,6 +9,7 @@ the reservation is refunded when the batch reaches a terminal state
"""
import asyncio
import uuid
from collections.abc import Awaitable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias
@ -85,6 +86,7 @@ class BatchEnqueuedTokenReservation:
tokens: int
scopes: tuple[BatchEnqueuedTokenScope, ...]
backend: ReservationBackend = "redis"
owner: str = ""
@dataclass(frozen=True, slots=True)
@ -177,7 +179,8 @@ class BatchEnqueuedTokenStore:
cross-slot commands), with an over-limit or failing scope rolling back the
scopes reserved before it; otherwise a single-process in-memory fallback
guarded by one asyncio lock is used. Reservations remember which backend
granted them so a refund never debits counters the grant did not charge. Everything expires after
granted them, and in-memory grants also remember the granting worker, so a
refund never debits counters the grant did not charge. Everything expires after
``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the
terminal-state refund can never leak tokens forever.
"""
@ -185,6 +188,7 @@ class BatchEnqueuedTokenStore:
def __init__(self, internal_usage_cache: "InternalUsageCache") -> None:
self.internal_usage_cache = internal_usage_cache
self._lock = asyncio.Lock()
self._owner_token = uuid.uuid4().hex
redis_cache = internal_usage_cache.dual_cache.redis_cache
self._reserve_script: _ScriptRunner | None = (
redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None
@ -300,7 +304,7 @@ class BatchEnqueuedTokenStore:
return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current)
for scope, current in zip(scopes, currents):
await self._set_local_counter(scope, current + tokens, span)
return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="memory")
return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="memory", owner=self._owner_token)
async def refund(
self,
@ -309,16 +313,14 @@ class BatchEnqueuedTokenStore:
) -> None:
if reservation.tokens <= 0 or not reservation.scopes:
return
refund_script: Final = self._refund_script
if reservation.backend == "redis" and refund_script is not None:
try:
await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes)
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters
verbose_proxy_logger.warning(
"Redis enqueued-token refund failed, falling back to in-memory: %s", str(e)
)
else:
return
if reservation.backend == "redis":
await self._refund_redis_reservation(reservation)
return
if reservation.owner != self._owner_token:
verbose_proxy_logger.warning(
"Skipping enqueued-token refund granted in another worker's memory; its counters expire with the TTL"
)
return
async with self._lock:
for scope in reservation.scopes:
current = await self._get_local_counter(scope, litellm_parent_otel_span)
@ -328,6 +330,20 @@ class BatchEnqueuedTokenStore:
else:
await self._set_local_counter(scope, remaining, litellm_parent_otel_span)
async def _refund_redis_reservation(self, reservation: BatchEnqueuedTokenReservation) -> None:
refund_script: Final = self._refund_script
if refund_script is None:
verbose_proxy_logger.warning(
"No Redis client for a Redis-granted enqueued-token refund; leaked increments expire with the TTL"
)
return
try:
await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes)
except Exception as e: # noqa: BLE001 # best-effort refund: the leak is TTL-bounded and only tightens the allowance
verbose_proxy_logger.warning(
"Redis enqueued-token refund failed; leaked increments expire with the TTL: %s", str(e)
)
async def save_reservation(
self,
batch_id: str,

View file

@ -128,10 +128,15 @@ async def test_zero_token_reserve_charges_nothing():
class _SingleKeyRedisFake:
"""Emulates the Redis script path one single-key call at a time, recording every call."""
def __init__(self, fail_reserve_keys: frozenset[str] = frozenset()) -> None:
def __init__(
self,
fail_reserve_keys: frozenset[str] = frozenset(),
fail_refund_keys: frozenset[str] = frozenset(),
) -> None:
self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = ()
self.counters: Mapping[str, int] = MappingProxyType({})
self.fail_reserve_keys = fail_reserve_keys
self.fail_refund_keys = fail_refund_keys
def async_register_script(self, script: str):
kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record"
@ -153,6 +158,8 @@ class _SingleKeyRedisFake:
self.counters = MappingProxyType({**self.counters, keys[0]: current + amount})
return (1, current + amount)
if kind == "refund":
if keys[0] in self.fail_refund_keys:
raise ConnectionError(f"simulated redis failure for {keys[0]}")
remaining: Final = self.counters.get(keys[0], 0) - int(args[0])
self.counters = MappingProxyType(
{key: value for key, value in self.counters.items() if key != keys[0]}
@ -207,6 +214,44 @@ async def test_partial_redis_reserve_failure_rolls_back_and_grants_in_memory():
assert refilled.backend == "memory"
@pytest.mark.asyncio
async def test_memory_refund_skips_reservations_granted_by_another_worker():
store = _in_memory_store()
scope = _scope(limit=100)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.backend == "memory"
assert reservation.owner
foreign: Final = BatchEnqueuedTokenReservation(
tokens=60, scopes=reservation.scopes, backend="memory", owner="another-worker"
)
await store.refund(foreign)
assert await store.reserve(tokens=50, scopes=(scope,)) == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=60)
await store.refund(reservation)
assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation)
@pytest.mark.asyncio
async def test_failed_redis_refund_leaves_local_counters_untouched():
scope = _scope(limit=100)
counter_key: Final = f"batch_enqueued_tokens:api_key:{scope.value}"
fake = _SingleKeyRedisFake(fail_refund_keys=frozenset({counter_key}))
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60))
)
reservation = await store.reserve(tokens=60, scopes=(scope,))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.backend == "redis"
store.internal_usage_cache.dual_cache.in_memory_cache.set_cache(key=counter_key, value=45)
await store.refund(reservation)
assert store.internal_usage_cache.dual_cache.in_memory_cache.get_cache(key=counter_key) == 45
assert fake.counters == {counter_key: 60}
@pytest.mark.asyncio
async def test_pop_reservation_defaults_legacy_records_to_redis_backend():
store = _in_memory_store()