From 6916432a7b2a28d43bfbb1a906eaba7a3a5cd4e0 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 22 Sep 2026 10:55:22 -0700 Subject: [PATCH] fix(proxy): seed db-down spend counter atomically and fall back to increment Set-max dropped spend that landed on the cold key before the seed, and a failed EVAL left a warm counter holding only the deltas. The seed now adds the cached spend once unless a seed already reached it, and falls back to INCRBYFLOAT when the atomic seed fails. --- litellm/caching/redis_cache.py | 19 +++ litellm/proxy/proxy_server.py | 35 +++- .../test_redis_seed_spend_counter.py | 54 ++++++ .../proxy/proxy_server/test_spend_counters.py | 161 +++++++++++------- 4 files changed, 200 insertions(+), 69 deletions(-) create mode 100644 tests/local_testing/test_redis_seed_spend_counter.py diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7cec84e0ebb..253e64029e9 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -102,7 +102,16 @@ _INCREMENT_WITH_FLOOR_LUA: Final = ( "return count" ) +_SEED_SPEND_COUNTER_LUA: Final = ( + "local cur = redis.call('GET', KEYS[1]) " + "if cur == false then redis.call('SET', KEYS[1], ARGV[1]) " + "elseif tonumber(cur) < tonumber(ARGV[1]) then redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]) end " + "if tonumber(ARGV[2]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end " + "return redis.call('GET', KEYS[1])" +) + _LUA_COUNT: Final = TypeAdapter(int) +_LUA_FLOAT: Final = TypeAdapter(float) _OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...]) @@ -1448,6 +1457,16 @@ class RedisCache(BaseCache): result = result.decode() return float(result) + @_redis_circuit_breaker_guard + async def async_seed_spend_counter(self, key: str, base: float) -> float: + """Atomically seed a spend counter with ``base``: set it when absent, add ``base`` when it holds + less (only increments that landed before any seed), and keep it when a seed already reached it.""" + _redis_client: Final = self._async_commands() + namespaced_key: Final = self.check_and_fix_namespace(key=key) + ttl: Final = int(self.get_ttl() or 0) + result: Final = await _redis_client.eval(_SEED_SPEND_COUNTER_LUA, 1, namespaced_key, str(base), str(ttl)) + return _LUA_FLOAT.validate_python(result.decode("utf-8") if isinstance(result, bytes) else result) + @_redis_circuit_breaker_guard async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int: """Async twin of ``increment_with_floor``, sharing its Lua script and its guarantees.""" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 459bf697f79..f08c2d0650a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3431,9 +3431,10 @@ async def _prepare_spend_counter_increment( 2. If not found, reseed from the DB via `SpendCounterReseed.coalesced`. Falls back to the cached object's `.spend` via user_api_key_cache only if prisma is unavailable, since that value can lag the flusher. - 3. Seed counter monotonically via `_repair_stale_spend_counter` (set-max), - so concurrent cold seeds on this pod or across pods converge on the - cached spend instead of summing it once per seeding request. + 3. Seed counter via `_seed_spend_counter_from_source_cache`, so concurrent + cold seeds on this pod or across pods add the cached spend once instead + of once per seeding request, while increments that landed on the cold + key before any seed are kept on top of it. 4. Increment is returned for the caller to apply via pipeline """ await _ensure_spend_counter_initialized( @@ -3536,7 +3537,33 @@ async def _ensure_spend_counter_initialized( # DB unavailable - fall back to in-process cache (may be stale). base_spend: Final = await _get_source_cache_base_spend(source_cache_key=source_cache_key) if base_spend > 0: - await _repair_stale_spend_counter(counter_key=counter_key, db_spend=base_spend) + await _seed_spend_counter_from_source_cache(counter_key=counter_key, base_spend=base_spend) + + +def _seeded_spend(current: object, base_spend: float) -> float: + if not isinstance(current, (int, float)): + return base_spend + return float(current) if current >= base_spend else current + base_spend + + +async def _seed_spend_counter_from_source_cache(counter_key: str, base_spend: float) -> None: + redis_cache: Final = spend_counter_cache.redis_cache + if redis_cache is None: + seeded: Final = _seeded_spend( + current=spend_counter_cache.in_memory_cache.get_cache(key=counter_key), base_spend=base_spend + ) + spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=seeded) + return + try: + current_value: Final = await redis_cache.async_seed_spend_counter(key=counter_key, base=base_spend) + except Exception: + verbose_proxy_logger.debug( + "Atomic seed of spend counter %s failed, falling back to increment", counter_key, exc_info=True + ) + await _increment_spend_counter_cache(counter_key=counter_key, increment=base_spend) + return + spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=current_value) + record_spend_counter_value(counter_key, current_value) async def _get_source_cache_base_spend( diff --git a/tests/local_testing/test_redis_seed_spend_counter.py b/tests/local_testing/test_redis_seed_spend_counter.py new file mode 100644 index 00000000000..416deaffe6f --- /dev/null +++ b/tests/local_testing/test_redis_seed_spend_counter.py @@ -0,0 +1,54 @@ +"""The spend-counter seed decides between set, add and keep inside a Lua script, so only a real +Redis can show the script itself is wrong.""" + +import asyncio +import os +import uuid +from typing import Final + +import pytest +from dotenv import load_dotenv + +load_dotenv() + +import litellm +from litellm.caching.redis_cache import RedisCache + + +@pytest.fixture +def counter(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "default_redis_ttl", 600) + cache: Final = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) + key: Final = f"spend:user:seed-{uuid.uuid4()}" + yield cache, key, cache.check_and_fix_namespace(key=key) + cache.delete_cache(key) + + +def test_a_cold_counter_is_seeded_with_the_base_and_gets_an_expiry(counter): + cache, key, namespaced_key = counter + + assert asyncio.run(cache.async_seed_spend_counter(key=key, base=6.0)) == 6.0 + assert 0 < cache.redis_client.ttl(namespaced_key) <= 600 + + +def test_concurrent_seeds_count_the_base_once(counter): + cache, key, _ = counter + + async def seed_five_times() -> tuple[float, ...]: + return tuple(await asyncio.gather(*(cache.async_seed_spend_counter(key=key, base=6.0) for _ in range(5)))) + + assert asyncio.run(seed_five_times()) == (6.0,) * 5 + + +def test_spend_that_landed_before_the_seed_is_kept_on_top_of_the_base(counter): + cache, key, namespaced_key = counter + cache.redis_client.incrbyfloat(namespaced_key, 0.5) + + assert asyncio.run(cache.async_seed_spend_counter(key=key, base=6.0)) == 6.5 + + +def test_a_counter_already_at_or_above_the_base_is_left_alone(counter): + cache, key, namespaced_key = counter + cache.redis_client.set(namespaced_key, 9.5) + + assert asyncio.run(cache.async_seed_spend_counter(key=key, base=6.0)) == 9.5 diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index d6fd056257c..6ad71839e60 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -1274,38 +1274,6 @@ async def test_ensure_spend_counter_initialized_warm_skips_reseed_and_source( } -@pytest.mark.asyncio -async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache( - monkeypatch, -): - fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=7.0) - fake_user_cache = _make_user_api_key_cache(get_value={"spend": 7.0}) - monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) - monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) - monkeypatch.setattr(ps, "prisma_client", None) - monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)) - - await ps._ensure_spend_counter_initialized( - counter_key="spend:user:u", - source_cache_key="u", - ) - - observed = { - "source_cache_called": fake_user_cache.async_get_cache.called, - "seed_set_max_value": fake_cache.redis_cache.async_set_max.call_args.kwargs["value"] == 7.0, - "seed_increment_called": fake_cache.redis_cache.async_increment.called, - "in_memory_seeded_value": fake_cache.in_memory_cache.set_cache.call_args.kwargs["value"], - "warm_check_done": fake_cache.redis_cache.async_get_cache.called, - } - assert normalize(observed) == { - "source_cache_called": True, - "seed_set_max_value": True, - "seed_increment_called": False, - "in_memory_seeded_value": 7.0, - "warm_check_done": True, - } - - class _DbDownSpendTable: def __init__(self) -> None: self.read_started: Final = asyncio.Event() @@ -1345,62 +1313,125 @@ async def test_ensure_spend_counter_initialized_concurrent_cold_seeds_converge_w assert cache.in_memory_cache.get_cache(key=counter_key) == cached_spend -class _FakeSetMaxRedis: - def __init__(self, existing: float | None = None) -> None: - self.value = existing +class _FakeSpendCounterRedis: + def __init__(self, seed_error: Exception | None = None) -> None: + self.store: dict[str, float] = {} # mutable-ok: stands in for Redis keyspace state + self._seed_error: Final = seed_error + + def get_ttl(self) -> int: + return 60 async def async_get_cache(self, key: str) -> float | None: - return self.value + return self.store.get(key) - async def async_set_max(self, key: str, value: float) -> float: - self.value = value if self.value is None else max(self.value, value) - return self.value + async def async_seed_spend_counter(self, key: str, base: float) -> float: + if self._seed_error is not None: + raise self._seed_error + current: Final = self.store.get(key) + self.store[key] = base if current is None else (current if current >= base else current + base) + return self.store[key] + + async def async_increment(self, key: str, value: float, refresh_ttl: bool = False) -> float: + self.store[key] = self.store.get(key, 0.0) + value + return self.store[key] + + async def async_increment_pipeline(self, increment_list: list[Mapping[str, object]]) -> list[float]: + return [ + await self.async_increment(key=str(op["key"]), value=float(str(op["increment_value"]))) + for op in increment_list + ] + + async def async_delete_cache(self, key: str) -> None: + self.store.pop(key, None) -class _DbDownSpendTableWhileAnotherPodSeeds: - def __init__(self, redis: _FakeSetMaxRedis, counter_key: str, other_pod_spend: float | None) -> None: +class _DbDownSpendTableWhileAnotherPodWrites: + def __init__(self, redis: _FakeSpendCounterRedis, counter_key: str, other_pod_seeds: bool, delta: float) -> None: self._redis: Final = redis self._counter_key: Final = counter_key - self._other_pod_spend: Final = other_pod_spend + self._other_pod_seeds: Final = other_pod_seeds + self._delta: Final = delta async def find_unique(self, where: Mapping[str, object]) -> SimpleNamespace: - if self._other_pod_spend is not None: - await self._redis.async_set_max(key=self._counter_key, value=self._other_pod_spend) + if self._other_pod_seeds: + await self._redis.async_seed_spend_counter(key=self._counter_key, base=6.0) + if self._delta: + await self._redis.async_increment(key=self._counter_key, value=self._delta) raise ConnectionError("database unavailable") -@pytest.mark.asyncio -@pytest.mark.parametrize( - "other_pod_spend", - [None, 9.5, 2.5], - ids=["cold", "another_pod_seeded_higher", "another_pod_seeded_lower"], -) -async def test_ensure_spend_counter_initialized_cold_seed_from_source_cache_is_monotonic_across_pods( - monkeypatch: pytest.MonkeyPatch, other_pod_spend: float | None -) -> None: - redis: Final = _FakeSetMaxRedis(existing=None) +def _db_down_spend_counter_setup( + monkeypatch: pytest.MonkeyPatch, redis: _FakeSpendCounterRedis | None, table: object, cached_spend: float +) -> DualCache: cache: Final = DualCache( in_memory_cache=InMemoryCache(), redis_cache=redis, # pyright: ignore[reportArgumentType] # duck-typed fake standing in for RedisCache ) - counter_key: Final = "spend:user:db-down-multi-pod-user" - cached_spend: Final = 6.0 - table: Final = _DbDownSpendTableWhileAnotherPodSeeds( - redis=redis, counter_key=counter_key, other_pod_spend=other_pod_spend - ) user_cache: Final = DualCache(in_memory_cache=InMemoryCache()) - user_cache.in_memory_cache.set_cache(key="db-down-multi-pod-user", value={"spend": cached_spend}) + user_cache.in_memory_cache.set_cache(key="db-down-user", value={"spend": cached_spend}) monkeypatch.setattr(ps, "spend_counter_cache", cache) monkeypatch.setattr(ps, "user_api_key_cache", user_cache) monkeypatch.setattr(ps, "prisma_client", SimpleNamespace(db=SimpleNamespace(litellm_usertable=table))) + return cache - await ps._ensure_spend_counter_initialized( - counter_key=counter_key, source_cache_key="db-down-multi-pod-user" + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("other_pod_seeds", "delta", "expected"), + [(False, 0.0, 6.0), (True, 0.0, 6.0), (True, 0.5, 6.5), (False, 0.5, 6.5)], + ids=["cold", "another_pod_already_seeded", "another_pod_seeded_then_spent", "unseeded_spend_landed_first"], +) +async def test_db_down_cold_seed_counts_cached_spend_once_and_keeps_spend_that_raced_it( + monkeypatch: pytest.MonkeyPatch, other_pod_seeds: bool, delta: float, expected: float +) -> None: + counter_key: Final = "spend:user:db-down-user" + redis: Final = _FakeSpendCounterRedis() + table: Final = _DbDownSpendTableWhileAnotherPodWrites( + redis=redis, counter_key=counter_key, other_pod_seeds=other_pod_seeds, delta=delta ) + cache: Final = _db_down_spend_counter_setup(monkeypatch, redis=redis, table=table, cached_spend=6.0) - expected: Final = cached_spend if other_pod_spend is None else max(other_pod_spend, cached_spend) - assert redis.value == expected - assert cache.in_memory_cache.get_cache(key=counter_key) == cached_spend + await ps._ensure_spend_counter_initialized(counter_key=counter_key, source_cache_key="db-down-user") + + assert (redis.store[counter_key], cache.in_memory_cache.get_cache(key=counter_key)) == (expected, expected) + + +@pytest.mark.asyncio +async def test_db_down_cold_seed_falls_back_to_increment_when_the_atomic_seed_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + counter_key: Final = "spend:user:db-down-user" + redis: Final = _FakeSpendCounterRedis(seed_error=Exception("NOPERM this user has no permissions to run 'eval'")) + table: Final = _DbDownSpendTableWhileAnotherPodWrites( + redis=redis, counter_key=counter_key, other_pod_seeds=False, delta=0.0 + ) + _db_down_spend_counter_setup(monkeypatch, redis=redis, table=table, cached_spend=90.0) + + pending: Final = await ps._prepare_spend_counter_increment( + counter_key=counter_key, source_cache_key="db-down-user", increment=1.0 + ) + await ps._apply_spend_counter_increments(pending=(pending,)) + + assert redis.store.get(counter_key) == 91.0, f"cached spend lost from the counter: {redis.store}" + + +@pytest.mark.asyncio +async def test_db_down_cold_seed_without_redis_keeps_spend_that_landed_during_the_db_read( + monkeypatch: pytest.MonkeyPatch, +) -> None: + counter_key: Final = "spend:user:db-down-user" + table: Final = _DbDownSpendTable() + cache: Final = _db_down_spend_counter_setup(monkeypatch, redis=None, table=table, cached_spend=6.0) + + seed: Final = asyncio.create_task( + ps._ensure_spend_counter_initialized(counter_key=counter_key, source_cache_key="db-down-user") + ) + await asyncio.wait_for(table.read_started.wait(), timeout=5) + cache.in_memory_cache.set_cache(key=counter_key, value=0.5) + table.resume_read.set() + await asyncio.wait_for(seed, timeout=5) + + assert cache.in_memory_cache.get_cache(key=counter_key) == 6.5 # ---------------------------------------------------------------------------