fix(proxy): issue enqueued-token Lua calls one key at a time for Redis Cluster compatibility

This commit is contained in:
mateo-berri 2026-08-19 15:40:15 -07:00
parent 7a6a677b72
commit 160d3dac42
2 changed files with 109 additions and 36 deletions

View file

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

View file

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