mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
50896f21b3
commit
4333d52813
2 changed files with 74 additions and 13 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue