diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index e35b1c8c82b..c6060a47f84 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -245,7 +245,16 @@ class SpendCounterReseed: value=current_value, ) else: - await spend_counter_cache.async_increment_cache(key=counter_key, value=db_spend, refresh_ttl=True) + # Repair/reservations can populate the counter during the DB read. + # Seed a floor without adding the database balance again. + # No await between read/compare/write: atomic within this worker. + cached = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) + current_value = float(db_spend) + if cached is not None: + current_value = max(current_value, float(cached)) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, value=current_value + ) except Exception: verbose_proxy_logger.exception( "SpendCounterReseed.coalesced: failed to warm counter %s", @@ -438,11 +447,20 @@ class SpendCounterReseed: value=current_value, ) else: - await spend_counter_cache.async_increment_cache(key=counter_key, value=window_spend) + # Repair/reservations can populate the counter during the DB read. + # Seed a floor without adding the database balance again. + # No await between read/compare/write: atomic within this worker. + cached = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) + current_value = float(window_spend) + if cached is not None: + current_value = max(current_value, float(cached)) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, value=current_value + ) except Exception: verbose_proxy_logger.exception( "SpendCounterReseed.coalesced_window: failed to warm counter %s", counter_key, ) raise - return window_spend + return current_value diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index 0361d5cfe8f..a0cf9128cc9 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -15,6 +15,7 @@ from typing import Final import pytest from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed @@ -270,6 +271,53 @@ async def test_coalesced_window_seeds_a_cold_counter_from_the_row(): assert prisma.db.litellm_spendlogs.call_count == 0 +@pytest.mark.asyncio +@pytest.mark.parametrize("window", [False, True], ids=["primary", "window"]) +@pytest.mark.parametrize("concurrent_spend", [989.01459411, 995.0, 900.0]) +async def test_cold_reseed_does_not_add_database_spend_to_concurrent_cache( + monkeypatch: pytest.MonkeyPatch, + window: bool, + concurrent_spend: float, +): + """A repair/reservation write may populate the counter while DB read runs. + + Reseeding must establish the larger value, not increment the concurrent + value by the same authoritative spend a second time. + """ + cache = DualCache(in_memory_cache=InMemoryCache()) + counter_key = ( + "spend:team:team-1:window:1d" if window else "spend:user:user-1" + ) + db_spend = 989.01459411 + + async def read_db(*args, **kwargs): + cache.in_memory_cache.set_cache(key=counter_key, value=concurrent_spend) + return db_spend + + if window: + monkeypatch.setattr(SpendCounterReseed, "window_from_db", staticmethod(read_db)) + result = await SpendCounterReseed.coalesced_window( + prisma_client=None, + spend_counter_cache=cache, + counter_key=counter_key, + entity_type="Team", + entity_id="team-1", + window_duration="1d", + window_start=WINDOW_START, + ) + else: + monkeypatch.setattr(SpendCounterReseed, "from_db", staticmethod(read_db)) + result = await SpendCounterReseed.coalesced( + prisma_client=None, + spend_counter_cache=cache, + counter_key=counter_key, + ) + + expected = max(db_spend, concurrent_spend) + assert cache.in_memory_cache.get_cache(key=counter_key) == expected + assert result == expected + + @pytest.mark.asyncio async def test_end_user_from_db_reads_the_end_user_row_by_user_id(): prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0))