diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 040c8b687a7..e484b6a59fe 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -38,27 +38,20 @@ ScopeKey: TypeAlias = Literal["api_key", "team"] RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """ local amount = tonumber(ARGV[1]) local ttl = tonumber(ARGV[2]) -for i = 1, #KEYS do - local limit = tonumber(ARGV[2 + i]) - local current = tonumber(redis.call('GET', KEYS[i]) or '0') - if current + amount > limit then - return {0, i - 1, current} - end +local limit = tonumber(ARGV[3]) +local current = tonumber(redis.call('GET', KEYS[1]) or '0') +if current + amount > limit then + return {0, current} end -for i = 1, #KEYS do - redis.call('INCRBY', KEYS[i], amount) - redis.call('EXPIRE', KEYS[i], ttl) -end -return {1, -1, 0} +local updated = redis.call('INCRBY', KEYS[1], amount) +redis.call('EXPIRE', KEYS[1], ttl) +return {1, updated} """ REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """ -local amount = tonumber(ARGV[1]) -for i = 1, #KEYS do - local updated = redis.call('DECRBY', KEYS[i], amount) - if updated <= 0 then - redis.call('DEL', KEYS[i]) - end +local updated = redis.call('DECRBY', KEYS[1], tonumber(ARGV[1])) +if updated <= 0 then + redis.call('DEL', KEYS[1]) end return 1 """ @@ -99,7 +92,7 @@ class BatchEnqueuedTokenOverLimit: BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit _LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)]) -_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int, int]) +_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int]) _POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None) _STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None) _RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation) @@ -175,9 +168,11 @@ def batch_response_view(response: object) -> _BatchResponseView | None: class BatchEnqueuedTokenStore: """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds. - Counters and records live in Redis (via atomic Lua scripts) when Redis is - configured; otherwise a single-process in-memory fallback guarded by one - asyncio lock is used. Everything expires after + Counters and records live in Redis when Redis is configured, through + single-key Lua scripts issued one scope at a time (Redis Cluster safe: no + cross-slot commands), with an over-limit scope rolling back the scopes + reserved before it; otherwise a single-process in-memory fallback guarded + by one asyncio lock is used. Everything expires after ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the terminal-state refund can never leak tokens forever. """ @@ -215,22 +210,44 @@ class BatchEnqueuedTokenStore: ) -> BatchEnqueuedTokenOutcome: if tokens <= 0 or not scopes: return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes) - if self._reserve_script is not None: + reserve_script: Final = self._reserve_script + refund_script: Final = self._refund_script + if reserve_script is not None and refund_script is not None: try: - raw_result = await self._reserve_script( - tuple(self._counter_key(scope) for scope in scopes), - (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, *(scope.limit for scope in scopes)), - ) - result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) - if result[0] == 1: - return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) - return BatchEnqueuedTokenOverLimit(scope=scopes[result[1]], enqueued=result[2]) + return await self._reserve_via_redis(reserve_script, refund_script, tokens=tokens, scopes=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 reserve failed, falling back to in-memory: %s", str(e) ) return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span) + async def _reserve_via_redis( + self, + reserve_script: _ScriptRunner, + refund_script: _ScriptRunner, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> BatchEnqueuedTokenOutcome: + for index, scope in enumerate(scopes): + raw_result = await reserve_script( + (self._counter_key(scope),), + (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit), + ) + result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) + if result[0] != 1: + await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes[:index]) + return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1]) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + + async def _refund_via_redis( + self, + refund_script: _ScriptRunner, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> None: + for scope in scopes: + await refund_script((self._counter_key(scope),), (tokens,)) + async def _reserve_in_memory( self, tokens: int, @@ -253,12 +270,10 @@ class BatchEnqueuedTokenStore: ) -> None: if reservation.tokens <= 0 or not reservation.scopes: return - if self._refund_script is not None: + refund_script: Final = self._refund_script + if refund_script is not None: try: - await self._refund_script( - tuple(self._counter_key(scope) for scope in reservation.scopes), - (reservation.tokens,), - ) + 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) 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 a917206f33a..1bb4798eaa0 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -9,7 +9,9 @@ response-shape helpers the v3 limiter's post-call hooks rely on. import base64 import socket import uuid -from types import SimpleNamespace +from collections.abc import Mapping, Sequence +from types import MappingProxyType, SimpleNamespace +from typing import Final import pytest @@ -123,6 +125,62 @@ async def test_zero_token_reserve_charges_nothing(): assert isinstance(full, BatchEnqueuedTokenReservation) +class _SingleKeyRedisFake: + """Emulates the Redis script path one single-key call at a time, recording every call.""" + + def __init__(self) -> None: + self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () + self.counters: Mapping[str, int] = MappingProxyType({}) + + def async_register_script(self, script: str): + kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record" + + async def run(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object: + self.script_calls = (*self.script_calls, (kind, tuple(keys))) + return self._run(kind, tuple(keys), tuple(args)) + + return run + + def _run(self, kind: str, keys: tuple[str, ...], args: tuple[str | bytes | int | float, ...]) -> object: + if kind == "reserve": + amount, limit = int(args[0]), int(args[2]) + current: Final = self.counters.get(keys[0], 0) + if current + amount > limit: + return (0, current) + self.counters = MappingProxyType({**self.counters, keys[0]: current + amount}) + return (1, current + amount) + if kind == "refund": + 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]} + if remaining <= 0 + else {**self.counters, keys[0]: remaining} + ) + return 1 + raise AssertionError(f"unexpected {kind} script call for keys {keys}") + + +@pytest.mark.asyncio +async def test_redis_reserve_issues_single_key_calls_and_rolls_back_on_over_limit(): + fake = _SingleKeyRedisFake() + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund") + assert not fake.counters + + fits = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(fits, BatchEnqueuedTokenReservation) + await store.refund(fits) + assert not fake.counters + assert all(len(keys) == 1 for _, keys in fake.script_calls) + + def test_canonical_provider_batch_id_passes_raw_ids_through(): assert canonical_provider_batch_id("batch_abc123") == "batch_abc123"