From 620fa2dcfdf947e01e60a2f2300d9d24cd25b48b Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 10 Feb 2026 13:23:25 -0600 Subject: [PATCH] fix(router-budget): await redis pipeline before sync reads --- litellm/router_strategy/budget_limiter.py | 29 +++--- .../router_strategy/test_budget_limiter.py | 94 +++++++++++++++++++ 2 files changed, 112 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/router_strategy/test_budget_limiter.py diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 9e4001b67b9..c8601c7e1ae 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -486,24 +486,31 @@ class RouterBudgetLimiting(CustomLogger): Only runs if Redis is initialized """ + increment_operations_to_flush: List[RedisPipelineIncrementOperation] = [] try: if not self.dual_cache.redis_cache: return # Redis is not initialized - verbose_router_logger.debug( - "Pushing Redis Increment Pipeline for queue: %s", - self.redis_increment_operation_queue, - ) - if len(self.redis_increment_operation_queue) > 0: - asyncio.create_task( - self.dual_cache.redis_cache.async_increment_pipeline( - increment_list=self.redis_increment_operation_queue, - ) - ) - + # Snapshot pending increments and clear queue before await, so new writes are queued + # for the next sync cycle while this batch is flushed to Redis. + increment_operations_to_flush = self.redis_increment_operation_queue self.redis_increment_operation_queue = [] + 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 as e: + if len(increment_operations_to_flush) > 0: + # Retry these increments on a future sync cycle. + self.redis_increment_operation_queue = ( + increment_operations_to_flush + self.redis_increment_operation_queue + ) verbose_router_logger.error( f"Error syncing in-memory cache with Redis: {str(e)}" ) diff --git a/tests/test_litellm/router_strategy/test_budget_limiter.py b/tests/test_litellm/router_strategy/test_budget_limiter.py new file mode 100644 index 00000000000..c60ece45efb --- /dev/null +++ b/tests/test_litellm/router_strategy/test_budget_limiter.py @@ -0,0 +1,94 @@ +import asyncio +from types import SimpleNamespace + +import pytest + +from litellm.router_strategy.budget_limiter import RouterBudgetLimiting +from litellm.types.utils import BudgetConfig + + +class _MockRedisCache: + def __init__(self, initial_values): + self.values = initial_values + self.events = [] + + async def async_increment_pipeline(self, increment_list, **kwargs): + self.events.append("increment_pipeline:start") + await asyncio.sleep(0.05) + for op in increment_list: + key = op["key"] + current = float(self.values.get(key, 0.0) or 0.0) + self.values[key] = current + float(op["increment_value"]) + self.events.append("increment_pipeline:done") + + async def async_batch_get_cache(self, key_list, **kwargs): + self.events.append("batch_get") + return {key: self.values.get(key) for key in key_list} + + +class _MockInMemoryCache: + def __init__(self, initial_values): + self.values = initial_values + + async def async_set_cache(self, key, value, **kwargs): + self.values[key] = float(value) + + +@pytest.mark.asyncio +async def test_should_await_redis_pipeline_before_sync_reads(): + spend_key = "provider_spend:openai:1d" + redis_cache = _MockRedisCache(initial_values={spend_key: 100.0}) + in_memory_cache = _MockInMemoryCache(initial_values={spend_key: 160.0}) + + budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting) + budget_limiter.dual_cache = SimpleNamespace( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + ) + budget_limiter.provider_budget_config = { + "openai": BudgetConfig(time_period="1d", budget_limit=500.0) + } + budget_limiter.deployment_budget_config = None + budget_limiter.tag_budget_config = None + budget_limiter.redis_increment_operation_queue = [ + { + "key": spend_key, + "increment_value": 60.0, + "ttl": 86400, + } + ] + + await budget_limiter._sync_in_memory_spend_with_redis() + + 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", + "increment_pipeline:done", + "batch_get", + ] + + +@pytest.mark.asyncio +async def test_should_requeue_increments_when_redis_pipeline_fails(): + spend_key = "provider_spend:openai:1d" + + class _FailingRedisCache: + async def async_increment_pipeline(self, increment_list, **kwargs): + raise RuntimeError("redis down") + + budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting) + budget_limiter.dual_cache = SimpleNamespace( + redis_cache=_FailingRedisCache(), + in_memory_cache=SimpleNamespace(), + ) + budget_limiter.redis_increment_operation_queue = [ + {"key": spend_key, "increment_value": 10.0, "ttl": 86400} + ] + + await budget_limiter._push_in_memory_increments_to_redis() + + assert budget_limiter.redis_increment_operation_queue == [ + {"key": spend_key, "increment_value": 10.0, "ttl": 86400} + ]