From 8e92cce27f68c0c8bed6aab32689ceb00c7b25c0 Mon Sep 17 00:00:00 2001 From: "{ \"message\": \"Bad credentials\", \"documentation_url\": \"https://docs.github.com/rest\", \"status\": \"401\" }" <{ "message": "Bad credentials", "documentation_url": "https://docs.github.com/rest", "status": "401" }+{ "message": "Bad credentials", "documentation_url": "https://docs.github.com/rest", "status": "401" }@users.noreply.github.com> Date: Thu, 21 May 2026 23:00:17 +0000 Subject: [PATCH] Fix Redis spend counter reseed race Co-authored-by: Krrish Dholakia --- litellm/proxy/db/spend_counter_reseed.py | 31 +++- litellm/proxy/proxy_server.py | 8 +- tests/test_litellm/proxy/test_proxy_server.py | 133 +++++++++++++++++- 3 files changed, 153 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 19ec6699390..d7af209c704 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -177,16 +177,33 @@ class SpendCounterReseed: db_spend = await SpendCounterReseed.from_db(prisma_client, counter_key) if db_spend is None: return None - # Warm even when 0 so subsequent reads hit cache, not DB. + # Warm even when 0 so subsequent reads hit cache, not DB. Use SET + # NX for Redis so concurrent cold pods do not each add db_spend. try: if spend_counter_cache.redis_cache is not None: - current_value = ( - await spend_counter_cache.redis_cache.async_increment( - key=counter_key, - value=db_spend, - refresh_ttl=True, - ) + seeded = await spend_counter_cache.redis_cache.async_set_cache( + key=counter_key, + value=db_spend, + nx=True, ) + if seeded: + current_value = db_spend + else: + current_cached_value = ( + await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key + ) + ) + if current_cached_value is None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=db_spend, + refresh_ttl=True, + ) + ) + else: + current_value = float(current_cached_value) spend_counter_cache.in_memory_cache.set_cache( key=counter_key, value=current_value, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 879914e5ac6..4f62695f11f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2214,11 +2214,9 @@ async def _init_and_increment_spend_counter( 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 via async_increment_cache (not async_set_cache) to avoid a - check-then-set race: if two pods cold-start simultaneously, both may see - the counter as absent and seed it. Using increment means the worst case - is over-counting (conservative, blocks slightly early) rather than - under-counting (would allow overspend). + 3. Seed Redis via atomic SET NX so simultaneous cold pods do not each add + the DB spend. Fall back to the cached winner before incrementing the + request cost. 4. Increment atomically (both in-memory + Redis) """ await _ensure_spend_counter_initialized( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 6d10d2a6353..cef2e598c23 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5699,13 +5699,19 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( from litellm.caching.dual_cache import DualCache counter_cache = DualCache() + recorded_sets: list = [] recorded_increments: list = [] + async def record_set(key, value, **kwargs): + recorded_sets.append({"key": key, "value": value, "nx": kwargs.get("nx")}) + return True + async def record_increment(key, value, ttl=None, **kwargs): recorded_increments.append({"key": key, "value": value, "ttl": ttl}) return value fake_redis = AsyncMock() + fake_redis.async_set_cache = AsyncMock(side_effect=record_set) fake_redis.async_increment = AsyncMock(side_effect=record_increment) fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing counter_cache.redis_cache = fake_redis @@ -5744,10 +5750,11 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( where={"team_id": "team-9"} ) - # Two increments keyed on the counter: seed ($42) then request ($1.50). - writes = [(c["key"], c["value"]) for c in recorded_increments] - assert ("spend:team:team-9", 42.0) in writes - assert ("spend:team:team-9", 1.5) in writes + # Seed uses SET NX ($42) and request cost uses increment ($1.50). + seed_writes = [(c["key"], c["value"], c["nx"]) for c in recorded_sets] + assert ("spend:team:team-9", 42.0, True) in seed_writes + increments = [(c["key"], c["value"]) for c in recorded_increments] + assert ("spend:team:team-9", 1.5) in increments finally: ps.user_api_key_cache = orig_user ps.spend_counter_cache = orig_counter @@ -5877,8 +5884,15 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): redis_store[key] = (redis_store.get(key) or 0.0) + value return redis_store[key] + async def redis_set_cache(key, value, **kwargs): + if kwargs.get("nx") and key in redis_store: + return False + redis_store[key] = value + return True + fake_redis = AsyncMock() fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis @@ -6299,12 +6313,13 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing(): counter_cache = DualCache() recorded_warms: list = [] - async def record_increment(key, value, ttl=None, **kwargs): + async def record_set(key, value, **kwargs): recorded_warms.append({"key": key, "value": value}) - return value + return True fake_redis = AsyncMock() - fake_redis.async_increment = AsyncMock(side_effect=record_increment) + fake_redis.async_set_cache = AsyncMock(side_effect=record_set) + fake_redis.async_increment = AsyncMock() fake_redis.async_get_cache = AsyncMock(return_value=None) counter_cache.redis_cache = fake_redis @@ -6408,7 +6423,14 @@ async def test_get_current_spend_coalesces_concurrent_reseeds(): redis_store[key] = (redis_store.get(key) or 0.0) + value return redis_store[key] + async def redis_set_cache(key, value, **kwargs): + if kwargs.get("nx") and key in redis_store: + return False + redis_store[key] = value + return True + fake_redis.async_get_cache = AsyncMock(side_effect=redis_get) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis @@ -6516,8 +6538,15 @@ async def test_concurrent_read_and_write_paths_share_one_db_query(): redis_store[key] = (redis_store.get(key) or 0.0) + value return redis_store[key] + async def redis_set_cache(key, value, **kwargs): + if kwargs.get("nx") and key in redis_store: + return False + redis_store[key] = value + return True + fake_redis = AsyncMock() fake_redis.async_get_cache = AsyncMock(side_effect=redis_get) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis @@ -6560,6 +6589,89 @@ async def test_concurrent_read_and_write_paths_share_one_db_query(): ps.user_api_key_cache = orig_user +@pytest.mark.asyncio +async def test_spend_counter_reseed_redis_concurrent_pods_do_not_double_seed( + monkeypatch, +): + """ + Pod-local singleflight locks cannot protect Redis from simultaneous cold + reseeds across pods. Redis seeding must therefore be idempotent: two pods + that both observe an initial miss should leave the counter at db_spend, + not 2 * db_spend. + """ + import asyncio as _asyncio + + from litellm.caching.dual_cache import DualCache + from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed + + counter_key = "spend:team:team-multi-pod-race" + redis_store: dict = {} + initial_miss_reads = 0 + both_pods_observed_miss = _asyncio.Event() + + async def redis_get_cache(key, **_): + nonlocal initial_miss_reads + observed = redis_store.get(key) + if observed is None: + initial_miss_reads += 1 + if initial_miss_reads >= 2: + both_pods_observed_miss.set() + await both_pods_observed_miss.wait() + return observed + + async def redis_set_cache(key, value, **kwargs): + if kwargs.get("nx") and key in redis_store: + return False + redis_store[key] = value + return True + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + + db_row = MagicMock() + db_row.spend = 42.0 + fake_prisma = MagicMock() + fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) + + async def independent_pod_lock(_counter_key): + return _asyncio.Lock() + + monkeypatch.setattr( + SpendCounterReseed, + "_get_lock", + staticmethod(independent_pod_lock), + ) + + pod_caches = [] + for _ in range(2): + counter_cache = DualCache() + counter_cache.redis_cache = fake_redis + pod_caches.append(counter_cache) + + await _asyncio.gather( + *[ + SpendCounterReseed.coalesced( + prisma_client=fake_prisma, + spend_counter_cache=counter_cache, + counter_key=counter_key, + ) + for counter_cache in pod_caches + ] + ) + + assert initial_miss_reads == 2 + assert fake_prisma.db.litellm_teamtable.find_unique.await_count == 2 + assert fake_redis.async_set_cache.await_count == 2 + fake_redis.async_increment.assert_not_awaited() + assert redis_store[counter_key] == pytest.approx(42.0) + + @pytest.mark.asyncio async def test_reseed_locks_dict_is_bounded(): """ @@ -6621,8 +6733,15 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): redis_store[key] = (redis_store.get(key) or 0.0) + value return redis_store[key] + async def redis_set_cache(key, value, **kwargs): + if kwargs.get("nx") and key in redis_store: + return False + redis_store[key] = value + return True + fake_redis = AsyncMock() fake_redis.async_get_cache = AsyncMock(side_effect=redis_get) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis