fix(router): keep budget increments if redis flush fails or is cancelled

A logging-worker timeout can cancel the awaited Redis pipeline after the
batch was already taken off the queue, which dropped spend. Shield the
pipeline, restore the batch if Redis fails, and skip the stale Redis
overwrite so in-memory spend cannot go backwards.
This commit is contained in:
Emerson Gomes 2026-08-21 20:38:53 -05:00
parent bd34801d2c
commit 6498d2f93b
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
3 changed files with 215 additions and 88 deletions

View file

@ -30,6 +30,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
self.dual_cache = dual_cache
self.redis_increment_operation_queue = []
self._redis_increment_queue_lock = asyncio.Lock()
self._redis_increment_flush_lock = asyncio.Lock()
self._detached_increment_operations = None
self.deployment_budget_config = None
async def is_key_within_model_budget(

View file

@ -26,7 +26,7 @@ from typing import Any, Final
import litellm
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
from litellm.caching.redis_cache import RedisCache, RedisPipelineIncrementOperation
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
@ -101,6 +101,8 @@ class RouterBudgetLimiting(CustomLogger):
self.dual_cache = dual_cache
self.redis_increment_operation_queue: list[RedisPipelineIncrementOperation] = []
self._redis_increment_queue_lock = asyncio.Lock()
self._redis_increment_flush_lock = asyncio.Lock()
self._detached_increment_operations: tuple[RedisPipelineIncrementOperation, ...] | None = None
asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis())
self.provider_budget_config: GenericBudgetConfigType | None = provider_budget_config
self.deployment_budget_config: GenericBudgetConfigType | None = None
@ -400,6 +402,73 @@ class RouterBudgetLimiting(CustomLogger):
def _get_redis_increment_queue_lock(self) -> asyncio.Lock:
return self._redis_increment_queue_lock
async def _detach_queued_increment_operations(self) -> tuple[RedisPipelineIncrementOperation, ...]:
async with self._get_redis_increment_queue_lock():
if self._detached_increment_operations is not None:
return self._detached_increment_operations
increment_operations_to_flush: Final = tuple(self.redis_increment_operation_queue)
self.redis_increment_operation_queue = [] # mutable-ok: emptied queue must stay appendable
self._detached_increment_operations = increment_operations_to_flush
return increment_operations_to_flush
async def _clear_detached_increment_operations(self) -> None:
async with self._get_redis_increment_queue_lock():
self._detached_increment_operations = None
async def _requeue_detached_increment_operations(self) -> None:
async with self._get_redis_increment_queue_lock():
detached_increment_operations: Final = self._detached_increment_operations
if detached_increment_operations is None:
return
self.redis_increment_operation_queue = (
list( # mutable-ok: restored flush batch must stay appendable
detached_increment_operations
)
+ self.redis_increment_operation_queue
)
self._detached_increment_operations = None
async def _finish_increment_pipeline_after_cancellation(
self,
pipeline_task: asyncio.Task[object],
) -> None:
try:
await pipeline_task
except Exception:
await self._requeue_detached_increment_operations()
verbose_router_logger.exception("Error pushing queued Redis increment operations to Redis")
return
await self._clear_detached_increment_operations()
async def _flush_queued_increment_operations(self, redis_cache: RedisCache) -> bool:
increment_operations_to_flush: Final = await self._detach_queued_increment_operations()
if len(increment_operations_to_flush) == 0:
return True
verbose_router_logger.debug(
"Pushing Redis Increment Pipeline for queue: %s",
increment_operations_to_flush,
)
increment_list: Final = list( # mutable-ok: Redis pipeline contract requires a list
increment_operations_to_flush
)
pipeline_task: Final = asyncio.create_task(
redis_cache.async_increment_pipeline(
increment_list=increment_list,
)
)
try:
await asyncio.shield(pipeline_task)
except Exception:
await asyncio.shield(self._requeue_detached_increment_operations())
verbose_router_logger.exception("Error pushing queued Redis increment operations to Redis")
return False
except asyncio.CancelledError:
await asyncio.shield(self._finish_increment_pipeline_after_cancellation(pipeline_task))
raise
await self._clear_detached_increment_operations()
return True
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""Original method now uses helper functions"""
verbose_router_logger.debug("in RouterBudgetLimiting.async_log_success_event")
@ -524,43 +593,25 @@ class RouterBudgetLimiting(CustomLogger):
DEFAULT_REDIS_SYNC_INTERVAL
) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying
async def _push_in_memory_increments_to_redis(self):
async def _push_in_memory_increments_to_redis(self) -> bool:
"""
How this works:
- async_log_success_event collects all provider spend increments in `redis_increment_operation_queue`
- This function pushes all increments to Redis in a batched pipeline to optimize performance
Only runs if Redis is initialized
Only runs if Redis is initialized. Returns False when the detached batch could not be
written, so callers must not treat Redis as up to date.
"""
increment_operations_to_flush: list[ # mutable-ok: Redis pipeline contract requires a concrete mutable batch
RedisPipelineIncrementOperation
] = ( # mutable-ok: Redis pipeline contract requires a list batch
[] # mutable-ok: Redis pipeline batches are lists
)
try:
if not self.dual_cache.redis_cache:
return # Redis is not initialized
redis_cache: Final = self.dual_cache.redis_cache
if redis_cache is None:
return True
async with self._get_redis_increment_queue_lock():
increment_operations_to_flush = self.redis_increment_operation_queue
self.redis_increment_operation_queue = [] # mutable-ok: the emptied queue must remain appendable by logging callbacks
verbose_router_logger.debug(
"Pushing Redis Increment Pipeline for queue: %s",
increment_operations_to_flush,
)
if len(increment_operations_to_flush) > 0:
await self.dual_cache.redis_cache.async_increment_pipeline(
increment_list=increment_operations_to_flush,
)
except Exception:
if len(increment_operations_to_flush) > 0:
async with self._get_redis_increment_queue_lock():
self.redis_increment_operation_queue = (
increment_operations_to_flush + self.redis_increment_operation_queue
)
verbose_router_logger.exception("Error pushing queued Redis increment operations to Redis")
async with self._redis_increment_flush_lock:
try:
return await self._flush_queued_increment_operations(redis_cache)
except asyncio.CancelledError:
await asyncio.shield(self._requeue_detached_increment_operations())
raise
async def _sync_in_memory_spend_with_redis(self):
"""
@ -581,7 +632,9 @@ class RouterBudgetLimiting(CustomLogger):
return
# 1. Push all provider spend increments to Redis
await self._push_in_memory_increments_to_redis()
flush_succeeded: Final = await self._push_in_memory_increments_to_redis()
if not flush_succeeded:
return
# 2. Fetch all current provider spend from Redis to update in-memory cache
cache_keys: Final = []

View file

@ -1,6 +1,5 @@
import asyncio
from types import SimpleNamespace
from typing import Optional
import pytest
@ -8,13 +7,19 @@ from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import BudgetConfig
_SPEND_KEY = "provider_spend:openai:1d"
def _increment(increment_value: float) -> RedisPipelineIncrementOperation:
return RedisPipelineIncrementOperation(key=_SPEND_KEY, increment_value=increment_value, ttl=86400)
class _MockRedisCache:
def __init__(
self,
initial_values: dict[str, float],
pipeline_started: Optional[asyncio.Event] = None,
allow_pipeline_to_complete: Optional[asyncio.Event] = None,
pipeline_started: asyncio.Event | None = None,
allow_pipeline_to_complete: asyncio.Event | None = None,
should_fail_pipeline: bool = False,
) -> None:
self.values = initial_values
@ -39,7 +44,7 @@ class _MockRedisCache:
self.values[key] = current + float(op["increment_value"])
self.events.append("increment_pipeline:done")
async def async_batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, Optional[float]]:
async def async_batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, float | None]:
self.events.append("batch_get")
return {key: self.values.get(key) for key in key_list}
@ -57,34 +62,46 @@ class _MockInMemoryCache:
self.values[key] = float(value)
@pytest.mark.asyncio
async def test_should_await_redis_pipeline_before_sync_reads() -> None:
spend_key = "provider_spend:openai:1d"
pipeline_started = asyncio.Event()
allow_pipeline_to_complete = asyncio.Event()
redis_cache = _MockRedisCache(
initial_values={spend_key: 100.0},
pipeline_started=pipeline_started,
allow_pipeline_to_complete=allow_pipeline_to_complete,
)
in_memory_cache = _MockInMemoryCache(initial_values={spend_key: 160.0})
def _new_router_budget_limiter(
*,
redis_cache: object,
in_memory_cache: object | None = None,
redis_increment_operation_queue: list[RedisPipelineIncrementOperation] | None = None,
provider_budget_config: dict[str, BudgetConfig] | None = None,
) -> RouterBudgetLimiting:
budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting)
budget_limiter.dual_cache = SimpleNamespace(
redis_cache=redis_cache,
in_memory_cache=in_memory_cache,
in_memory_cache=in_memory_cache if in_memory_cache is not None else SimpleNamespace(),
)
budget_limiter.provider_budget_config = {"openai": BudgetConfig(time_period="1d", budget_limit=500.0)}
budget_limiter.provider_budget_config = provider_budget_config
budget_limiter.deployment_budget_config = None
budget_limiter.tag_budget_config = None
budget_limiter.redis_increment_operation_queue = [
RedisPipelineIncrementOperation(
key=spend_key,
increment_value=60.0,
ttl=86400,
)
]
budget_limiter.redis_increment_operation_queue = (
list(redis_increment_operation_queue) if redis_increment_operation_queue is not None else []
)
budget_limiter._redis_increment_queue_lock = asyncio.Lock()
budget_limiter._redis_increment_flush_lock = asyncio.Lock()
budget_limiter._detached_increment_operations = None
return budget_limiter
@pytest.mark.asyncio
async def test_should_await_redis_pipeline_before_sync_reads() -> None:
pipeline_started = asyncio.Event()
allow_pipeline_to_complete = asyncio.Event()
redis_cache = _MockRedisCache(
initial_values={_SPEND_KEY: 100.0},
pipeline_started=pipeline_started,
allow_pipeline_to_complete=allow_pipeline_to_complete,
)
in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 160.0})
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
in_memory_cache=in_memory_cache,
redis_increment_operation_queue=[_increment(60.0)],
provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=500.0)},
)
sync_task = asyncio.create_task(budget_limiter._sync_in_memory_spend_with_redis())
await asyncio.wait_for(pipeline_started.wait(), timeout=1)
@ -92,8 +109,8 @@ async def test_should_await_redis_pipeline_before_sync_reads() -> None:
allow_pipeline_to_complete.set()
await sync_task
assert redis_cache.values[spend_key] == 160.0
assert in_memory_cache.values[spend_key] == 160.0
assert redis_cache.values[_SPEND_KEY] == 160.0
assert in_memory_cache.values[_SPEND_KEY] == 160.0
assert budget_limiter.redis_increment_operation_queue == []
assert redis_cache.events == [
"increment_pipeline:start",
@ -104,31 +121,21 @@ async def test_should_await_redis_pipeline_before_sync_reads() -> None:
@pytest.mark.asyncio
async def test_should_requeue_increments_when_redis_pipeline_fails() -> None:
spend_key = "provider_spend:openai:1d"
redis_cache = _MockRedisCache(
initial_values={},
should_fail_pipeline=True,
)
budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting)
budget_limiter.dual_cache = SimpleNamespace(
redis_cache = _MockRedisCache(initial_values={}, should_fail_pipeline=True)
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
in_memory_cache=SimpleNamespace(),
redis_increment_operation_queue=[_increment(10.0)],
)
budget_limiter.redis_increment_operation_queue = [
RedisPipelineIncrementOperation(key=spend_key, increment_value=10.0, ttl=86400)
]
budget_limiter._redis_increment_queue_lock = asyncio.Lock()
await budget_limiter._push_in_memory_increments_to_redis()
flush_succeeded = await budget_limiter._push_in_memory_increments_to_redis()
assert budget_limiter.redis_increment_operation_queue == [
RedisPipelineIncrementOperation(key=spend_key, increment_value=10.0, ttl=86400)
]
assert flush_succeeded is False
assert budget_limiter.redis_increment_operation_queue == [_increment(10.0)]
assert budget_limiter._detached_increment_operations is None
@pytest.mark.asyncio
async def test_should_keep_new_increments_when_pipeline_flush_fails() -> None:
spend_key = "provider_spend:openai:1d"
pipeline_started = asyncio.Event()
allow_pipeline_to_complete = asyncio.Event()
redis_cache = _MockRedisCache(
@ -137,24 +144,89 @@ async def test_should_keep_new_increments_when_pipeline_flush_fails() -> None:
allow_pipeline_to_complete=allow_pipeline_to_complete,
should_fail_pipeline=True,
)
in_memory_cache = _MockInMemoryCache(initial_values={spend_key: 0.0})
budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting)
budget_limiter.dual_cache = SimpleNamespace(
in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 0.0})
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
in_memory_cache=in_memory_cache,
redis_increment_operation_queue=[_increment(10.0)],
)
budget_limiter.redis_increment_operation_queue = [
RedisPipelineIncrementOperation(key=spend_key, increment_value=10.0, ttl=86400)
]
budget_limiter._redis_increment_queue_lock = asyncio.Lock()
push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis())
await asyncio.wait_for(pipeline_started.wait(), timeout=1)
await budget_limiter._increment_spend_in_current_window(spend_key=spend_key, response_cost=20.0, ttl=86400)
await budget_limiter._increment_spend_in_current_window(spend_key=_SPEND_KEY, response_cost=20.0, ttl=86400)
allow_pipeline_to_complete.set()
await push_task
assert budget_limiter.redis_increment_operation_queue == [
RedisPipelineIncrementOperation(key=spend_key, increment_value=10.0, ttl=86400),
RedisPipelineIncrementOperation(key=spend_key, increment_value=20.0, ttl=86400),
]
assert budget_limiter.redis_increment_operation_queue == [_increment(10.0), _increment(20.0)]
@pytest.mark.asyncio
async def test_should_keep_in_memory_spend_when_redis_pipeline_fails() -> None:
redis_cache = _MockRedisCache(initial_values={_SPEND_KEY: 100.0}, should_fail_pipeline=True)
in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 160.0})
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
in_memory_cache=in_memory_cache,
redis_increment_operation_queue=[_increment(60.0)],
provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=500.0)},
)
await budget_limiter._sync_in_memory_spend_with_redis()
assert in_memory_cache.values[_SPEND_KEY] == 160.0
assert redis_cache.values[_SPEND_KEY] == 100.0
assert budget_limiter.redis_increment_operation_queue == [_increment(60.0)]
assert "batch_get" not in redis_cache.events
@pytest.mark.asyncio
async def test_should_keep_increments_when_flush_is_cancelled_after_success() -> None:
pipeline_started = asyncio.Event()
allow_pipeline_to_complete = asyncio.Event()
redis_cache = _MockRedisCache(
initial_values={_SPEND_KEY: 0.0},
pipeline_started=pipeline_started,
allow_pipeline_to_complete=allow_pipeline_to_complete,
)
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
redis_increment_operation_queue=[_increment(10.0)],
)
push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis())
await asyncio.wait_for(pipeline_started.wait(), timeout=1)
push_task.cancel()
allow_pipeline_to_complete.set()
with pytest.raises(asyncio.CancelledError):
await push_task
assert redis_cache.values[_SPEND_KEY] == 10.0
assert budget_limiter.redis_increment_operation_queue == []
assert budget_limiter._detached_increment_operations is None
@pytest.mark.asyncio
async def test_should_requeue_increments_when_flush_is_cancelled_and_redis_fails() -> None:
pipeline_started = asyncio.Event()
allow_pipeline_to_complete = asyncio.Event()
redis_cache = _MockRedisCache(
initial_values={_SPEND_KEY: 0.0},
pipeline_started=pipeline_started,
allow_pipeline_to_complete=allow_pipeline_to_complete,
should_fail_pipeline=True,
)
budget_limiter = _new_router_budget_limiter(
redis_cache=redis_cache,
redis_increment_operation_queue=[_increment(10.0)],
)
push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis())
await asyncio.wait_for(pipeline_started.wait(), timeout=1)
push_task.cancel()
allow_pipeline_to_complete.set()
with pytest.raises(asyncio.CancelledError):
await push_task
assert redis_cache.values[_SPEND_KEY] == 0.0
assert budget_limiter.redis_increment_operation_queue == [_increment(10.0)]
assert budget_limiter._detached_increment_operations is None