diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 570291aa452..bbc007bb00d 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index bc924c32ba3..d6ae200dfaf 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -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()