mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): issue enqueued-token Lua calls one key at a time for Redis Cluster compatibility
This commit is contained in:
parent
7a6a677b72
commit
160d3dac42
2 changed files with 109 additions and 36 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue