fix(proxy): seed db-down spend counter atomically and fall back to increment
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

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.
This commit is contained in:
ryan-crabbe-berri 2026-09-22 10:55:22 -07:00
parent 37e187ac52
commit 6916432a7b
4 changed files with 200 additions and 69 deletions

View file

@ -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."""

View file

@ -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(

View file

@ -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

View file

@ -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
# ---------------------------------------------------------------------------