fix(proxy): renew budget reservation counter TTL while the request is in flight

A reservation lives inside spend counter keys that expire on the Redis idle
TTL (60s). A stream that outlives the TTL dropped its reservation, so a
concurrent request on any worker was admitted against the DB floor until the
stream reconciled. Renew the counter TTL with EXPIRE every ttl/2 while the
reservation is open and stop once it is finalized, so an idle counter still
expires on its own if the worker dies.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-08 23:10:03 +00:00
parent 5b2b5420af
commit a558a0b6a9
4 changed files with 123 additions and 2 deletions

View file

@ -76,6 +76,8 @@ class _AsyncRedisCommands(Protocol):
def ttl(self, name: str) -> Awaitable[int]: ...
def expire(self, name: str, time: int) -> Awaitable[bool]: ...
def rpush(self, name: str, *values: str | bytes | float) -> Awaitable[int]: ...
def lpop(self, name: str, count: int | None = None) -> Awaitable[object]: ...
@ -1795,6 +1797,19 @@ class RedisCache(BaseCache):
_record_swallowed_redis_failure(self._circuit_breaker, e)
return None
@_redis_circuit_breaker_guard
async def async_refresh_ttl(self, key: str, ttl: int | None = None) -> bool:
"""EXPIRE an existing key without touching its value. False when the key is absent or Redis failed."""
_used_ttl: Final = self.get_ttl(ttl=ttl)
if _used_ttl is None:
return False
try:
return await self._async_commands().expire(self.check_and_fix_namespace(key=key), _used_ttl)
except Exception as e:
verbose_logger.debug("Redis EXPIRE Error: %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
return False
@_redis_circuit_breaker_guard
async def async_rpush(
self,

View file

@ -3222,6 +3222,12 @@ async def increment_spend_counter(counter_key: str, increment: float):
return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def refresh_spend_counter_ttl(counter_key: str) -> bool:
if spend_counter_cache.redis_cache is None:
return False
return await spend_counter_cache.redis_cache.async_refresh_ttl(key=counter_key)
async def _increment_spend_counter_cache(counter_key: str, increment: float):
if spend_counter_cache.redis_cache is not None:
try:

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio
import json
import time
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@ -99,6 +100,42 @@ def get_reserved_counter_keys(budget_reservation: dict | None) -> set:
}
_lease_renewals: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio only weak-refs pending tasks
def _start_reservation_lease_renewal(budget_reservation: Mapping[str, object], counter_keys: frozenset[str]) -> None:
"""A reservation lives inside spend counter keys that expire on their Redis TTL. Renew the TTL
while the request is in flight so a request longer than the TTL does not drop its
reservation and admit concurrent requests against the DB floor on any worker."""
from litellm.proxy.proxy_server import spend_counter_cache
if spend_counter_cache.redis_cache is None or not counter_keys:
return
task: Final = asyncio.create_task(
_renew_reservation_lease(
budget_reservation=budget_reservation,
counter_keys=counter_keys,
interval=spend_counter_cache.redis_cache.default_ttl / 2,
)
)
_lease_renewals.add(task)
task.add_done_callback(_lease_renewals.discard)
async def _renew_reservation_lease(
budget_reservation: Mapping[str, object], counter_keys: frozenset[str], interval: float
) -> None:
from litellm.proxy.proxy_server import refresh_spend_counter_ttl
deadline: Final = time.monotonic() + litellm.request_timeout
while time.monotonic() < deadline:
await asyncio.sleep(interval)
if budget_reservation.get("finalized") is True:
return
for counter_key in counter_keys:
await refresh_spend_counter_ttl(counter_key=counter_key)
def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: UserAPIKeyAuth | None) -> bool:
"""
Whether an over-budget key's own ``max_budget`` reservation should be
@ -294,12 +331,17 @@ async def reserve_budget_for_request(
llm_router=llm_router,
input_token_counts=input_token_counts,
)
return {
budget_reservation: Final = {
"reserved_cost": reservation_cost,
"entries": applied_entries,
"finalized": False,
"input_cost": min(float(input_cost or 0.0), reservation_cost),
}
_start_reservation_lease_renewal(
budget_reservation=budget_reservation,
counter_keys=frozenset(get_reserved_counter_keys(budget_reservation=budget_reservation)),
)
return budget_reservation
async def reconcile_budget_reservation(

View file

@ -1,5 +1,6 @@
import asyncio
import threading
import time
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
@ -2199,26 +2200,83 @@ async def test_release_non_numeric_counter_reseeds_from_db(spend_counter_state):
class _ExpiringRedisCache:
def __init__(self) -> None:
"""In-memory stand-in for RedisCache with real wall-clock key expiry."""
def __init__(self, default_ttl: float = 60.0) -> None:
self.default_ttl = default_ttl
self.store: dict[str, float] = {}
self.expires_at: dict[str, float] = {}
self.refresh_count = 0
def _evict_expired(self, key: str) -> None:
if self.expires_at.get(key, float("inf")) <= time.monotonic():
self.store.pop(key, None)
self.expires_at.pop(key, None)
async def async_get_cache(self, key: str, *args: object, **kwargs: object) -> float | None:
self._evict_expired(key)
return self.store.get(key)
async def async_increment(self, key: str, value: float, **kwargs: object) -> float:
self._evict_expired(key)
self.store[key] = self.store.get(key, 0.0) + float(value)
self.expires_at[key] = time.monotonic() + self.default_ttl
return self.store[key]
async def async_set_max(self, key: str, value: float, **kwargs: object) -> float:
self._evict_expired(key)
self.store[key] = max(self.store.get(key, float("-inf")), float(value))
self.expires_at[key] = time.monotonic() + self.default_ttl
return self.store[key]
async def async_set_cache(self, key: str, value: float, *args: object, **kwargs: object) -> bool:
self.store[key] = float(value)
self.expires_at[key] = time.monotonic() + self.default_ttl
return True
async def async_delete_cache(self, key: str, *args: object, **kwargs: object) -> None:
self.store.pop(key, None)
self.expires_at.pop(key, None)
async def async_refresh_ttl(self, key: str, ttl: int | None = None) -> bool:
self._evict_expired(key)
if key not in self.store:
return False
self.refresh_count += 1
self.expires_at[key] = time.monotonic() + (ttl if ttl is not None else self.default_ttl)
return True
@pytest.mark.asyncio
async def test_reservation_survives_redis_counter_ttl_while_request_in_flight(
spend_counter_state,
):
"""A request that runs longer than the counter TTL must keep its reservation in Redis
(so a concurrent request on any worker still sees it), and renewal must stop once the
reservation is reconciled so an idle counter still expires on its own."""
counter_cache, key_cache = spend_counter_state
redis_cache = _ExpiringRedisCache(default_ttl=0.2)
counter_cache.redis_cache = redis_cache
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(token="key-lease", spend=0.0, max_budget=1.0)
counter_key = "spend:key:key-lease"
reservation = await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj)
assert reservation is not None
await asyncio.sleep(0.5)
assert await redis_cache.async_get_cache(key=counter_key) == pytest.approx(0.6)
concurrent = await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj)
assert concurrent is not None
assert concurrent["reserved_cost"] == pytest.approx(0.4)
await release_budget_reservation(reservation)
await release_budget_reservation(concurrent)
await asyncio.sleep(0.15)
refreshes_after_release = redis_cache.refresh_count
await asyncio.sleep(0.35)
assert redis_cache.refresh_count == refreshes_after_release
assert await redis_cache.async_get_cache(key=counter_key) is None
@pytest.mark.asyncio