fix(proxy): share per-model budget counters across replicas through the spend counter cache (#39375)

* fix(proxy): share per-model budget counters across replicas through the spend counter cache

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): keep the shared fake Redis store immutable

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-02 12:40:59 -07:00 • committed by GitHub
parent cde9d94c36
commit 987ab76921
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 132 additions and 14 deletions

View file

@ -1,6 +1,6 @@
import json
import time
from collections.abc import Iterable, Mapping
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
@ -199,18 +199,10 @@ async def build_model_max_budget_usage(
)
for budget_model, budget_config in budgets
)
batched: Final = await cache.async_batch_get_cache(
keys=list(spend_keys) # mutable-ok: async_batch_get_cache annotates keys as list, so one must exist here
)
# async_batch_get_cache returns None if it fails internally, and its result is
# index-aligned with `keys` otherwise. An unusable result reads as a miss,
# which is what a never-written counter already reads as.
current_spends: Final = (
tuple(batched) if isinstance(batched, list) and len(batched) == len(budgets) else (None,) * len(budgets)
)
current_spends: Final = await _current_window_spends(cache=cache, spend_keys=spend_keys)
return {
budget_model: {
"current_spend": round(_as_spend(current_spend), 4),
"current_spend": round(current_spend, 4),
"budget_limit": budget_config.max_budget,
"time_period": budget_config.budget_duration,
}
@ -218,6 +210,22 @@ async def build_model_max_budget_usage(
}
async def _current_window_spends(cache: DualCache, spend_keys: Sequence[str]) -> tuple[float, ...]:
"""Redis holds the window total across replicas; the in-memory copy is one replica's share."""
keys: Final = list(spend_keys) # mutable-ok: both batch readers annotate their key argument as list
redis_cache: Final = cache.redis_cache
if redis_cache is not None:
shared: Final = await redis_cache.async_batch_get_cache(key_list=keys)
return tuple(_as_spend(shared.get(key)) for key in keys)
# async_batch_get_cache returns None if it fails internally, and its result is
# index-aligned with `keys` otherwise. An unusable result reads as a miss,
# which is what a never-written counter already reads as.
batched: Final = await cache.async_batch_get_cache(keys=keys)
if not isinstance(batched, list) or len(batched) != len(keys):
return (0.0,) * len(keys)
return tuple(_as_spend(current_spend) for current_spend in batched)
def _usable_budget_config(raw_budget_config: object) -> BudgetConfig | None:
try:
budget_config: Final = BudgetConfig.model_validate(raw_budget_config)
@ -404,7 +412,10 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return current_spend + _as_spend(await self._cached_spend(legacy_spend_key))
async def _cached_spend(self, spend_key: str) -> float | None:
return await self.dual_cache.async_get_cache(key=spend_key)
redis_cache: Final = self.dual_cache.redis_cache
if redis_cache is None:
return await self.dual_cache.async_get_cache(key=spend_key)
return await redis_cache.async_get_cache(key=spend_key)
async def async_filter_deployments(
self,

View file

@ -2280,7 +2280,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache(
)
spend_counter_cache: Final = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value)
cli_sso_session_cache: Final = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS)
model_max_budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=user_api_key_cache)
model_max_budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache)
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
redis_usage_cache: RedisCache | None = None # redis cache used for tracking spend, tpm/rpm limits
polling_via_cache_enabled: Literal["all"] | list[str] | bool = False

View file

@ -1,3 +1,5 @@
import asyncio
from types import MappingProxyType
from unittest.mock import AsyncMock, patch
@ -5,6 +7,7 @@ import pytest
import litellm
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisCache
from datetime import datetime, timezone
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
@ -1332,3 +1335,85 @@ async def test_the_user_scope_has_no_pre_upgrade_counter_to_carry():
await limiter.is_user_within_model_budget(
user_id="u1", user_model_max_budget=model_max_budget, model="openai/gpt-4"
)
class _SharedFakeRedis(RedisCache):
"""Stand-in for the one Redis every replica's DualCache is attached to.
Only the methods the limiter and DualCache call are implemented, and
``super().__init__`` is skipped so no connection is opened.
"""
def __init__(self):
self._store = MappingProxyType({})
async def async_set_cache(self, key, value, **kwargs):
self._store = MappingProxyType({**self._store, key: value})
async def async_get_cache(self, key, **kwargs):
return self._store.get(key)
async def async_batch_get_cache(self, key_list, **kwargs):
return {key: self._store.get(key) for key in key_list}
async def async_increment_pipeline(self, increment_list, **kwargs):
for op in increment_list:
total = self._store.get(op["key"], 0.0) + op["increment_value"]
self._store = MappingProxyType({**self._store, op["key"]: total})
return [self._store[op["key"]] for op in increment_list]
async def _log_spend(limiter, *, key_hash, model_max_budget, response_cost):
await limiter.async_log_success_event(
_success_kwargs(
model_group="gpt-4",
response_cost=response_cost,
key_hash=key_hash,
key_model_max_budget=model_max_budget,
),
response_obj=None,
start_time=None,
end_time=None,
)
# The Redis push is scheduled as a task rather than awaited inline.
await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task()))
@pytest.mark.asyncio
async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another():
"""
Each replica increments its own in-memory copy of the per-model counter and
pushes the increment to the shared Redis, so only Redis holds the window's
total. A replica that has served part of the traffic must still enforce and
report the total, not its own share.
Regression: reads went to the in-memory tier first, so a replica whose local
copy sat under the cap kept admitting requests and /key/info on it reported
that local share, while the shared counter was already over the cap.
"""
shared_redis = _SharedFakeRedis()
replica_a = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis))
replica_b = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis))
key_hash = "vk-shared"
model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "30d"}}
user_api_key = UserAPIKeyAuth(token=key_hash, model_max_budget=model_max_budget)
await _log_spend(replica_b, key_hash=key_hash, model_max_budget=model_max_budget, response_cost=0.25)
await _log_spend(replica_a, key_hash=key_hash, model_max_budget=model_max_budget, response_cost=0.5)
await _log_spend(replica_a, key_hash=key_hash, model_max_budget=model_max_budget, response_cost=0.5)
with pytest.raises(litellm.BudgetExceededError):
await replica_b.is_key_within_model_budget(user_api_key, "gpt-4")
usage_on_b = await build_model_max_budget_usage(
entity_type=Litellm_EntityType.KEY,
entity_id=key_hash,
model_max_budget=model_max_budget,
cache=replica_b.dual_cache,
)
assert usage_on_b["gpt-4"]["current_spend"] == 1.25
# Control: a replica that never served this key reads the same total.
replica_c = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis))
with pytest.raises(litellm.BudgetExceededError):
await replica_c.is_key_within_model_budget(user_api_key, "gpt-4")

View file

@ -5,7 +5,7 @@ Verifies that _init_cache attaches Redis to user_api_key_cache only when
the flag is explicitly set to True, and leaves it in-memory-only otherwise.
"""
from contextlib import contextmanager
from contextlib import ExitStack, contextmanager
import json
from unittest.mock import MagicMock, patch
@ -167,3 +167,25 @@ class TestRedisAuthCacheFlag:
f"cli_sso_session_cache must always get Redis "
f"(enable_redis_auth_cache={flag_value!r})"
)
def test_flag_absent_still_shares_the_model_budget_counters_over_redis(self):
"""
Per-model budget counters are spend counters: the limiter must be able to
push and read them through Redis without the auth-cache opt-in, or every
worker enforces and reports its own share of a key's spend
"""
fake_redis = _FakeRedisCache()
limiter_cache = ps.model_max_budget_limiter.dual_cache
touched_caches = (
limiter_cache,
ps.spend_counter_cache,
ps.cli_sso_session_cache,
ps.user_api_key_cache,
ps.litellm_config_cache,
)
with ExitStack() as detached:
for cache in touched_caches:
detached.enter_context(patch.object(cache, "redis_cache", None))
ps._attach_redis_usage_cache(fake_redis, enable_redis_auth_cache=False)
assert limiter_cache.redis_cache is fake_redis
assert ps.user_api_key_cache.redis_cache is None