mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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
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:
parent
37e187ac52
commit
6916432a7b
4 changed files with 200 additions and 69 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
54
tests/local_testing/test_redis_seed_spend_counter.py
Normal file
54
tests/local_testing/test_redis_seed_spend_counter.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue