fix(router): reset budget windows atomically so concurrent spend is not lost

This commit is contained in:
Devin AI 2026-07-27 15:45:33 +00:00
parent 24123269cc
commit 222f51f8ef
2 changed files with 274 additions and 23 deletions

View file

@ -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,

View file

@ -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))