mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(router): reset budget windows atomically so concurrent spend is not lost
This commit is contained in:
parent
24123269cc
commit
222f51f8ef
2 changed files with 274 additions and 23 deletions
|
|
@ -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,
|
||||
|
|
|
|||
199
tests/test_litellm/router_strategy/test_budget_limiter.py
Normal file
199
tests/test_litellm/router_strategy/test_budget_limiter.py
Normal 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))
|
||||
Loading…
Add table
Reference in a new issue