From 01284d489319e817459ee1d4a84212c618568f24 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 19 Sep 2026 08:57:32 +0000 Subject: [PATCH] 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> --- litellm/proxy/proxy_server.py | 10 +- .../proxy/proxy_server/test_spend_counters.py | 107 +++++++++++++++++- 2 files changed, 110 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 36fdea605c2..5b38c6ac11e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index 0731c233fef..d6fd056257c 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -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 # ---------------------------------------------------------------------------