mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Fix Redis spend counter reseed race
Co-authored-by: Krrish Dholakia <krrish-berri-2@users.noreply.github.com>
This commit is contained in:
parent
79b4578671
commit
8e92cce27f
3 changed files with 153 additions and 19 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue