From 222f51f8ef07269c6df1a848ca5f2b716f61c57c Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 27 Jul 2026 15:45:33 +0000 Subject: [PATCH] fix(router): reset budget windows atomically so concurrent spend is not lost --- litellm/router_strategy/budget_limiter.py | 98 +++++++-- .../router_strategy/test_budget_limiter.py | 199 ++++++++++++++++++ 2 files changed, 274 insertions(+), 23 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 067f38ab11c..c59cfec5ed7 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -20,7 +20,7 @@ anthropic: import asyncio from datetime import datetime, timedelta, timezone -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, Union import litellm from litellm._logging import verbose_router_logger @@ -43,6 +43,31 @@ from litellm.types.utils import GenericBudgetConfigType, StandardLoggingPayload DEFAULT_REDIS_SYNC_INTERVAL = 1 +BUDGET_WINDOW_RESET_SCRIPT = """ +local start_time_key = KEYS[1] +local spend_key = KEYS[2] +local current_time = ARGV[1] +local response_cost = ARGV[2] +local ttl = tonumber(ARGV[3]) + +local window_start = redis.call('GET', start_time_key) +if window_start == false or (tonumber(current_time) - tonumber(window_start)) > ttl then + redis.call('SET', start_time_key, current_time, 'EX', ttl) + redis.call('SET', spend_key, response_cost, 'EX', ttl) + return {current_time, response_cost} +end + +local new_spend = redis.call('INCRBYFLOAT', spend_key, response_cost) +if redis.call('TTL', spend_key) < 0 then + redis.call('EXPIRE', spend_key, ttl) +end +return {window_start, new_spend} +""" + + +def _as_float(value: Union[bytes, str, float, int]) -> float: + return float(value.decode() if isinstance(value, bytes) else value) + class _LiteLLMParamsDictView: """ @@ -99,6 +124,12 @@ class RouterBudgetLimiting(CustomLogger): model_list: Optional[List[Union[DeploymentTypedDict, Dict[str, Any]]]] = None, ): self.dual_cache = dual_cache + self._budget_window_reset_script: Callable[..., Awaitable[Any]] | None = ( + dual_cache.redis_cache.async_register_script(BUDGET_WINDOW_RESET_SCRIPT) + if dual_cache.redis_cache is not None + else None + ) + self._local_budget_window_lock = asyncio.Lock() self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = [] asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis()) self.provider_budget_config: Optional[GenericBudgetConfigType] = provider_budget_config @@ -359,20 +390,50 @@ class RouterBudgetLimiting(CustomLogger): ttl_seconds: int, ) -> float: """ - Handle start of new budget window by resetting spend and start time + Start a new budget window, or join one another caller just started. - Enters this when: - - The budget does not exist in cache, so we need to set it - - The budget window has expired, so we need to reset everything + Enters this when the budget window has expired, which several concurrent + responses can observe at the same time. Only the first of them may reset + the spend to its own cost; the others have to add their cost on top, else + every reset would drop the spend written by the resets racing with it. - Does 2 things: - - stores key: `provider_spend:{provider}:1d`, value: response_cost - - stores key: `provider_budget_start_time:{provider}`, value: current_time. - This stores the start time of the new budget window + Redis does the window check, reset and increment in one atomic script so + the winner is decided server-side, and the resulting window is mirrored + into the in-memory cache. Without Redis a local lock is enough, since + there is a single instance reading and writing the spend. + + Returns the start time of the window this spend was recorded against. """ - await self.dual_cache.async_set_cache(key=spend_key, value=response_cost, ttl=ttl_seconds) - await self.dual_cache.async_set_cache(key=start_time_key, value=current_time, ttl=ttl_seconds) - return current_time + if self._budget_window_reset_script is not None: + try: + raw_window_start, raw_spend = await self._budget_window_reset_script( + keys=[start_time_key, spend_key], + args=[str(current_time), str(response_cost), ttl_seconds], + ) + window_start = _as_float(raw_window_start) + await self.dual_cache.in_memory_cache.async_set_cache( + key=start_time_key, value=window_start, ttl=ttl_seconds + ) + await self.dual_cache.in_memory_cache.async_set_cache( + key=spend_key, value=_as_float(raw_spend), ttl=ttl_seconds + ) + return window_start + except Exception as e: + verbose_router_logger.warning( + "Atomic budget window reset failed for %s, falling back to local reset: %s", + spend_key, + str(e), + ) + + async with self._local_budget_window_lock: + existing_start = await self.dual_cache.async_get_cache(start_time_key) + if existing_start is not None and (current_time - float(existing_start)) <= ttl_seconds: + await self.dual_cache.async_increment_cache(key=spend_key, value=response_cost, ttl=ttl_seconds) + return float(existing_start) + + await self.dual_cache.async_set_cache(key=spend_key, value=response_cost, ttl=ttl_seconds) + await self.dual_cache.async_set_cache(key=start_time_key, value=current_time, ttl=ttl_seconds) + return current_time async def _increment_spend_in_current_window(self, spend_key: str, response_cost: float, ttl: int): """ @@ -471,19 +532,10 @@ class RouterBudgetLimiting(CustomLogger): ttl_seconds=ttl_seconds, ) - if budget_start is None: - # First spend for this provider - budget_start = await self._handle_new_budget_window( - spend_key=spend_key, - start_time_key=start_time_key, - current_time=current_time, - response_cost=response_cost, - ttl_seconds=ttl_seconds, - ) - elif (current_time - budget_start) > ttl_seconds: + if (current_time - budget_start) > ttl_seconds: # Budget window expired - reset everything verbose_router_logger.debug("Budget window expired - resetting everything") - budget_start = await self._handle_new_budget_window( + await self._handle_new_budget_window( spend_key=spend_key, start_time_key=start_time_key, current_time=current_time, 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..def8cf9ced8 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_budget_limiter.py @@ -0,0 +1,199 @@ +import asyncio +from typing import Any, Dict, List, Optional, Sequence + +import pytest + +from litellm.caching.caching import DualCache +from litellm.router_strategy.budget_limiter import RouterBudgetLimiting +from litellm.types.utils import BudgetConfig + +TTL_SECONDS = 86400 + + +class YieldingDualCache(DualCache): + """DualCache that suspends on every read/write, so concurrent callers interleave.""" + + async def async_get_cache(self, key, parent_otel_span=None, local_only: bool = False, **kwargs): + await asyncio.sleep(0) + return await super().async_get_cache(key, parent_otel_span, local_only, **kwargs) + + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): + await asyncio.sleep(0) + return await super().async_set_cache(key, value, local_only, **kwargs) + + +class FakeRedisCacheWithAtomicScripts: + """ + Stand-in for RedisCache that runs registered scripts atomically over a local dict. + + Mirrors what Redis guarantees for a Lua script: the body observes and mutates the + store without another caller interleaving, while callers still race to enter it. + """ + + def __init__(self) -> None: + self.store: Dict[str, str] = {} + self.registered_scripts: List[str] = [] + + def async_register_script(self, script: str): + self.registered_scripts.append(script) + + async def run_script(keys: Sequence[str], args: Sequence[Any], client: Optional[Any] = None) -> List[bytes]: + await asyncio.sleep(0) + start_time_key, spend_key = keys + current_time, response_cost, ttl = str(args[0]), str(args[1]), float(args[2]) + + window_start = self.store.get(start_time_key) + if window_start is None or (float(current_time) - float(window_start)) > ttl: + self.store[start_time_key] = current_time + self.store[spend_key] = response_cost + return [current_time.encode(), response_cost.encode()] + + new_spend = str(float(self.store.get(spend_key, "0")) + float(response_cost)) + self.store[spend_key] = new_spend + return [window_start.encode(), new_spend.encode()] + + return run_script + + +@pytest.fixture +def disable_budget_sync(monkeypatch): + async def noop(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis", + noop, + ) + + +@pytest.mark.asyncio +async def test_concurrent_expired_window_resets_keep_every_response_cost(disable_budget_sync): + """Every response crossing an expired window boundary must be counted, not just the last one.""" + budget_limiter = RouterBudgetLimiting( + dual_cache=YieldingDualCache(), + provider_budget_config={"openai": BudgetConfig(budget_duration="1d", max_budget=100)}, + ) + spend_key = "provider_spend:openai:1d" + start_time_key = "provider_budget_start_time:openai" + now = 1_000_000.0 + + await budget_limiter.dual_cache.async_set_cache( + key=start_time_key, value=now - (2 * TTL_SECONDS), ttl=10 * TTL_SECONDS + ) + await budget_limiter.dual_cache.async_set_cache(key=spend_key, value=7.0, ttl=10 * TTL_SECONDS) + + costs = (0.5, 0.25, 0.125) + await asyncio.gather( + *[ + budget_limiter._handle_new_budget_window( + spend_key=spend_key, + start_time_key=start_time_key, + current_time=now, + response_cost=cost, + ttl_seconds=TTL_SECONDS, + ) + for cost in costs + ] + ) + + spend = await budget_limiter.dual_cache.async_get_cache(spend_key) + assert float(spend) == pytest.approx(sum(costs)) + + window_start = await budget_limiter.dual_cache.async_get_cache(start_time_key) + assert float(window_start) == now + + +@pytest.mark.asyncio +async def test_concurrent_expired_window_resets_keep_every_response_cost_with_redis(disable_budget_sync): + """With Redis the window reset and the increment are delegated to one atomic script.""" + fake_redis = FakeRedisCacheWithAtomicScripts() + fake_redis.store["provider_budget_start_time:openai"] = str(1_000_000.0 - (2 * TTL_SECONDS)) + fake_redis.store["provider_spend:openai:1d"] = "7.0" + + budget_limiter = RouterBudgetLimiting( + dual_cache=DualCache(redis_cache=fake_redis), + provider_budget_config={"openai": BudgetConfig(budget_duration="1d", max_budget=100)}, + ) + spend_key = "provider_spend:openai:1d" + start_time_key = "provider_budget_start_time:openai" + now = 1_000_000.0 + + costs = (0.5, 0.25, 0.125) + window_starts = await asyncio.gather( + *[ + budget_limiter._handle_new_budget_window( + spend_key=spend_key, + start_time_key=start_time_key, + current_time=now, + response_cost=cost, + ttl_seconds=TTL_SECONDS, + ) + for cost in costs + ] + ) + + assert float(fake_redis.store[spend_key]) == pytest.approx(sum(costs)) + assert float(fake_redis.store[start_time_key]) == now + assert window_starts == [now, now, now] + + in_memory_spend = await budget_limiter.dual_cache.in_memory_cache.async_get_cache(spend_key) + assert float(in_memory_spend) == pytest.approx(sum(costs)) + + +@pytest.mark.asyncio +async def test_expired_window_reset_drops_previous_window_spend(disable_budget_sync): + """A single response crossing the boundary still starts the new window from its own cost.""" + budget_limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": BudgetConfig(budget_duration="1d", max_budget=100)}, + ) + spend_key = "provider_spend:openai:1d" + start_time_key = "provider_budget_start_time:openai" + now = 1_000_000.0 + + await budget_limiter.dual_cache.async_set_cache( + key=start_time_key, value=now - (2 * TTL_SECONDS), ttl=10 * TTL_SECONDS + ) + await budget_limiter.dual_cache.async_set_cache(key=spend_key, value=7.0, ttl=10 * TTL_SECONDS) + + window_start = await budget_limiter._handle_new_budget_window( + spend_key=spend_key, + start_time_key=start_time_key, + current_time=now, + response_cost=0.5, + ttl_seconds=TTL_SECONDS, + ) + + assert window_start == now + spend = await budget_limiter.dual_cache.async_get_cache(spend_key) + assert float(spend) == pytest.approx(0.5) + + +@pytest.mark.asyncio +async def test_concurrent_success_events_across_window_boundary_keep_every_response_cost(disable_budget_sync): + """End-to-end through _increment_spend_for_key: expired window, concurrent responses.""" + budget_limiter = RouterBudgetLimiting( + dual_cache=YieldingDualCache(), + provider_budget_config={"openai": BudgetConfig(budget_duration="1d", max_budget=100)}, + ) + budget_config = BudgetConfig(budget_duration="1d", max_budget=100) + spend_key = "provider_spend:openai:1d" + start_time_key = "provider_budget_start_time:openai" + + await budget_limiter.dual_cache.async_set_cache(key=start_time_key, value=0.0, ttl=10 * TTL_SECONDS) + + costs = (0.5, 0.25) + await asyncio.gather( + *[ + budget_limiter._increment_spend_for_key( + budget_config=budget_config, + spend_key=spend_key, + start_time_key=start_time_key, + response_cost=cost, + ) + for cost in costs + ] + ) + + spend = await budget_limiter.dual_cache.async_get_cache(spend_key) + assert float(spend) == pytest.approx(sum(costs))