mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): seed cold spend counter monotonically when the db is unavailable
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5fc510a6fd
commit
01284d4893
2 changed files with 110 additions and 7 deletions
|
|
@ -3301,11 +3301,9 @@ 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 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 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.
|
||||
4. Increment is returned for the caller to apply via pipeline
|
||||
"""
|
||||
await _ensure_spend_counter_initialized(
|
||||
|
|
@ -3408,7 +3406,7 @@ 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 _increment_spend_counter_cache(counter_key=counter_key, increment=base_spend)
|
||||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=base_spend)
|
||||
|
||||
|
||||
async def _get_source_cache_base_spend(
|
||||
|
|
|
|||
|
|
@ -22,13 +22,17 @@ Pins covered:
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
|
||||
from .conftest import normalize
|
||||
|
||||
|
|
@ -1288,16 +1292,117 @@ async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache(
|
|||
|
||||
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_increment_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()
|
||||
self.resume_read: Final = asyncio.Event()
|
||||
|
||||
async def find_unique(self, where: Mapping[str, object]) -> SimpleNamespace:
|
||||
self.read_started.set()
|
||||
await self.resume_read.wait()
|
||||
raise ConnectionError("database unavailable")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_spend_counter_initialized_concurrent_cold_seeds_converge_when_db_is_unavailable(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
cache: Final = DualCache(in_memory_cache=InMemoryCache())
|
||||
counter_key: Final = "spend:user:db-down-user"
|
||||
cached_spend: Final = 6.0
|
||||
table: Final = _DbDownSpendTable()
|
||||
user_cache: Final = DualCache(in_memory_cache=InMemoryCache())
|
||||
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)))
|
||||
|
||||
first: 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)
|
||||
second: Final = asyncio.create_task(
|
||||
ps._ensure_spend_counter_initialized(counter_key=counter_key, source_cache_key="db-down-user")
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
table.resume_read.set()
|
||||
await asyncio.wait_for(asyncio.gather(first, second), timeout=5)
|
||||
|
||||
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
|
||||
|
||||
async def async_get_cache(self, key: str) -> float | None:
|
||||
return self.value
|
||||
|
||||
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
|
||||
|
||||
|
||||
class _DbDownSpendTableWhileAnotherPodSeeds:
|
||||
def __init__(self, redis: _FakeSetMaxRedis, counter_key: str, other_pod_spend: float | None) -> None:
|
||||
self._redis: Final = redis
|
||||
self._counter_key: Final = counter_key
|
||||
self._other_pod_spend: Final = other_pod_spend
|
||||
|
||||
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)
|
||||
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)
|
||||
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})
|
||||
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)))
|
||||
|
||||
await ps._ensure_spend_counter_initialized(
|
||||
counter_key=counter_key, source_cache_key="db-down-multi-pod-user"
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_source_cache_base_spend
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue