From 2949c89bb9170fed3ca918eb0b11e46c39c64e26 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Thu, 9 Jul 2026 08:28:41 -0500 Subject: [PATCH 1/5] fix(router): await budget redis pipeline before sync reads --- .../proxy/hooks/model_max_budget_limiter.py | 2 + litellm/router_strategy/budget_limiter.py | 33 ++-- ...test_unit_test_max_model_budget_limiter.py | 18 ++ .../router_strategy/test_budget_limiter.py | 160 ++++++++++++++++++ 4 files changed, 202 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/router_strategy/test_budget_limiter.py diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 215969ef899..d0a02153ff0 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -1,3 +1,4 @@ +import asyncio import json from typing import Final @@ -28,6 +29,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): def __init__(self, dual_cache: DualCache): self.dual_cache = dual_cache self.redis_increment_operation_queue = [] + self._redis_increment_queue_lock = asyncio.Lock() self.deployment_budget_config = None async def is_key_within_model_budget( diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index d57d7da0410..c52040c8e3d 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -100,6 +100,7 @@ class RouterBudgetLimiting(CustomLogger): ): self.dual_cache = dual_cache self.redis_increment_operation_queue: list[RedisPipelineIncrementOperation] = [] + self._redis_increment_queue_lock = asyncio.Lock() 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 @@ -393,7 +394,11 @@ class RouterBudgetLimiting(CustomLogger): increment_value=response_cost, ttl=ttl, ) - self.redis_increment_operation_queue.append(increment_op) + async with self._get_redis_increment_queue_lock(): + self.redis_increment_operation_queue.append(increment_op) + + def _get_redis_increment_queue_lock(self) -> asyncio.Lock: + return self._redis_increment_queue_lock async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): """Original method now uses helper functions""" @@ -527,25 +532,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 + async with self._get_redis_increment_queue_lock(): + 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", - self.redis_increment_operation_queue, + increment_operations_to_flush, ) - 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, - ) + if len(increment_operations_to_flush) > 0: + await self.dual_cache.redis_cache.async_increment_pipeline( + increment_list=increment_operations_to_flush, ) - self.redis_increment_operation_queue = [] - - except Exception as e: - verbose_router_logger.error("Error syncing in-memory cache with Redis: %s", e) + 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 def _sync_in_memory_spend_with_redis(self): """ diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py index 55459721906..18bfa58ca8d 100644 --- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py +++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py @@ -14,6 +14,7 @@ from litellm.proxy.hooks.model_max_budget_limiter import ( _PROXY_VirtualKeyModelMaxBudgetLimiter, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import BudgetConfig as GenericBudgetInfo @@ -452,6 +453,23 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config mock_push.assert_awaited_once() +@pytest.mark.asyncio +async def test_model_budget_limiter_initializes_redis_increment_queue_lock(): + dual_cache = DualCache() + limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + spend_key = "virtual_key_spend:test-key:gpt-4:1d" + + await limiter._increment_spend_in_current_window( + spend_key=spend_key, response_cost=0.01, ttl=86400 + ) + + assert limiter.redis_increment_operation_queue == [ + RedisPipelineIncrementOperation( + key=spend_key, increment_value=0.01, ttl=86400 + ) + ] + + @pytest.mark.asyncio async def test_get_fallback_model_within_budget_returns_none_without_fallbacks( budget_limiter, 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..c5f19c260cc --- /dev/null +++ b/tests/test_litellm/router_strategy/test_budget_limiter.py @@ -0,0 +1,160 @@ +import asyncio +from types import SimpleNamespace +from typing import Optional + +import pytest + +from litellm.router_strategy.budget_limiter import RouterBudgetLimiting +from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.utils import BudgetConfig + + +class _MockRedisCache: + def __init__( + self, + initial_values: dict[str, float], + pipeline_started: Optional[asyncio.Event] = None, + allow_pipeline_to_complete: Optional[asyncio.Event] = None, + should_fail_pipeline: bool = False, + ) -> None: + self.values = initial_values + self.events: list[str] = [] + self.pipeline_started = pipeline_started + self.allow_pipeline_to_complete = allow_pipeline_to_complete + self.should_fail_pipeline = should_fail_pipeline + + async def async_increment_pipeline( + self, increment_list: list[RedisPipelineIncrementOperation], **kwargs: object + ) -> None: + self.events.append("increment_pipeline:start") + if self.pipeline_started is not None: + self.pipeline_started.set() + if self.allow_pipeline_to_complete is not None: + await self.allow_pipeline_to_complete.wait() + if self.should_fail_pipeline: + raise RuntimeError("redis down") + 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: list[str], **kwargs: object) -> dict[str, Optional[float]]: + self.events.append("batch_get") + return {key: self.values.get(key) for key in key_list} + + +class _MockInMemoryCache: + def __init__(self, initial_values: dict[str, float]) -> None: + self.values = initial_values + + async def async_increment(self, key: str, value: float, ttl: int, **kwargs: object) -> float: + current = float(self.values.get(key, 0.0) or 0.0) + self.values[key] = current + float(value) + return self.values[key] + + async def async_set_cache(self, key: str, value: float, **kwargs: object) -> None: + 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}) + + 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 = [ + RedisPipelineIncrementOperation( + key=spend_key, + increment_value=60.0, + ttl=86400, + ) + ] + budget_limiter._redis_increment_queue_lock = asyncio.Lock() + + sync_task = asyncio.create_task(budget_limiter._sync_in_memory_spend_with_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + assert "batch_get" not in redis_cache.events + 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 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() -> 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=redis_cache, + in_memory_cache=SimpleNamespace(), + ) + 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() + + assert budget_limiter.redis_increment_operation_queue == [ + RedisPipelineIncrementOperation(key=spend_key, increment_value=10.0, ttl=86400) + ] + + +@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( + initial_values={}, + pipeline_started=pipeline_started, + 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( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + ) + 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) + 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), + ] From 397b5c8bd621338dbb49d8571d4b4c5562a560cc Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Wed, 12 Aug 2026 12:50:48 -0500 Subject: [PATCH 2/5] style(router): explain mutable budget queue --- litellm/router_strategy/budget_limiter.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index c52040c8e3d..e14ecb58b68 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -532,14 +532,18 @@ class RouterBudgetLimiting(CustomLogger): Only runs if Redis is initialized """ - increment_operations_to_flush: list[RedisPipelineIncrementOperation] = [] + increment_operations_to_flush: list[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 async with self._get_redis_increment_queue_lock(): increment_operations_to_flush = self.redis_increment_operation_queue - 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", From e7e61b18108693e43f47fa54f3e392dec03f067d Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Wed, 12 Aug 2026 13:09:13 -0500 Subject: [PATCH 3/5] style(router): format budget limiter --- litellm/router_strategy/budget_limiter.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index e14ecb58b68..48b9c011b77 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -532,7 +532,9 @@ class RouterBudgetLimiting(CustomLogger): Only runs if Redis is initialized """ - increment_operations_to_flush: list[RedisPipelineIncrementOperation] = ( # mutable-ok: Redis pipeline contract requires a list batch + increment_operations_to_flush: list[ + RedisPipelineIncrementOperation + ] = ( # mutable-ok: Redis pipeline contract requires a list batch [] # mutable-ok: Redis pipeline batches are lists ) try: @@ -541,9 +543,7 @@ class RouterBudgetLimiting(CustomLogger): 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 - ) + 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", From bd34801d2c892df1eebbd2227f9eee648ee88515 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Wed, 12 Aug 2026 13:52:13 -0500 Subject: [PATCH 4/5] style(router): document Redis flush batch --- litellm/router_strategy/budget_limiter.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 48b9c011b77..007104190a2 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -532,7 +532,7 @@ class RouterBudgetLimiting(CustomLogger): Only runs if Redis is initialized """ - increment_operations_to_flush: list[ + 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 From 6498d2f93be1cc29a1f40375db229a493cbefe37 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Fri, 21 Aug 2026 20:38:53 -0500 Subject: [PATCH 5/5] 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. --- .../proxy/hooks/model_max_budget_limiter.py | 2 + litellm/router_strategy/budget_limiter.py | 117 ++++++++--- .../router_strategy/test_budget_limiter.py | 184 ++++++++++++------ 3 files changed, 215 insertions(+), 88 deletions(-) diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index d0a02153ff0..60c692d51c5 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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( diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 007104190a2..f747e99cec1 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -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 = [] diff --git a/tests/test_litellm/router_strategy/test_budget_limiter.py b/tests/test_litellm/router_strategy/test_budget_limiter.py index c5f19c260cc..9ef677cff42 100644 --- a/tests/test_litellm/router_strategy/test_budget_limiter.py +++ b/tests/test_litellm/router_strategy/test_budget_limiter.py @@ -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