From 41d2cc101284430f42d6e9d646ce13a2abebf3ba Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 25 Apr 2026 13:59:43 -0700 Subject: [PATCH] fix(budget): reseed model counters on miss and avoid baseline double-count --- litellm/proxy/auth/auth_checks.py | 30 ++++-------- litellm/proxy/proxy_server.py | 11 ++--- .../proxy/auth/test_auth_checks.py | 47 +++++++++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 12 +++-- 4 files changed, 68 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 275d278c14e..3e24f50ed97 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1296,19 +1296,6 @@ async def get_team_membership( return None -async def get_team_member_model_spend( - user_id: str, - team_id: str, - model: str, - user_api_key_cache: DualCache, -) -> float: - _key = f"team_member_model_spend:{user_id}:{team_id}:{model}" - cached_spend = await user_api_key_cache.async_get_cache(key=_key) - if cached_spend is not None: - return float(cached_spend) - return 0.0 - - def model_in_access_group( model: str, team_models: Optional[List[str]], llm_router: Optional[Router] ) -> bool: @@ -3398,7 +3385,9 @@ async def _check_team_member_model_budget( from litellm.proxy.proxy_server import ( _get_team_member_model_counter_key, + _reseed_spend_from_db, get_current_spend, + spend_counter_cache, ) for requested_model_str in models_to_check: @@ -3440,18 +3429,15 @@ async def _check_team_member_model_budget( team_id=team_object.team_id, model=counter_model, ) - fallback_spend = 0.0 - if user_api_key_cache is not None: - fallback_spend = await get_team_member_model_spend( - user_id=valid_token.user_id, - team_id=team_object.team_id, - model=counter_model, - user_api_key_cache=user_api_key_cache, - ) model_spend = await get_current_spend( counter_key=counter_key, - fallback_spend=fallback_spend, + fallback_spend=-1.0, ) + if model_spend < 0: + model_spend = await _reseed_spend_from_db(counter_key) + await spend_counter_cache.async_set_cache( + key=counter_key, value=model_spend, nx=True + ) if model_spend >= max_budget: raise litellm.BudgetExceededError( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1da6bfe3bec..0dca9832593 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2039,11 +2039,8 @@ async def _init_and_increment_spend_counter( 2. If not found, reseed from the DB (`_reseed_spend_from_db`). 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 via async_set_cache(nx=True) to avoid double-counting the + DB baseline when multiple pods cold-start simultaneously. 4. Increment atomically (both in-memory + Redis) """ current = await spend_counter_cache.async_get_cache(key=counter_key) @@ -2059,8 +2056,8 @@ async def _init_and_increment_spend_counter( else: base_spend = getattr(source, "spend", 0.0) or 0.0 if base_spend > 0: - await spend_counter_cache.async_increment_cache( - key=counter_key, value=base_spend + await spend_counter_cache.async_set_cache( + key=counter_key, value=base_spend, nx=True ) await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 1064f4614bb..c8af66954e4 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2387,3 +2387,50 @@ async def test_team_member_model_budget_accepts_numeric_string_max_budget(): ) assert exc_info.value.current_cost == 11.0 assert exc_info.value.max_budget == 10.0 + + +@pytest.mark.asyncio +async def test_team_member_model_budget_reseeds_on_counter_miss(): + team_object = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_model_max_budget": {"gpt-4o": {"max_budget": 10.0}}}, + ) + valid_token = UserAPIKeyAuth( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + return fallback_spend + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.proxy_server._reseed_spend_from_db", + new_callable=AsyncMock, + return_value=11.0, + ) as mock_reseed, + patch( + "litellm.proxy.proxy_server.spend_counter_cache.async_set_cache", + new_callable=AsyncMock, + ) as mock_set_cache, + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_team_member_model_budget( + team_object=team_object, + valid_token=valid_token, + model="gpt-4o", + llm_router=None, + user_api_key_cache=None, + ) + assert exc_info.value.current_cost == 11.0 + assert exc_info.value.max_budget == 10.0 + mock_reseed.assert_awaited_once_with( + "spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o" + ) + mock_set_cache.assert_awaited_once_with( + key="spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o", + value=11.0, + nx=True, + ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9bd39db3f6f..d4367b91548 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4982,14 +4982,20 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( counter_cache = DualCache() recorded_increments: list = [] + recorded_sets: list = [] async def record_increment(key, value, ttl=None, **kwargs): recorded_increments.append({"key": key, "value": value, "ttl": ttl}) return value + async def record_set(key, value, **kwargs): + recorded_sets.append({"key": key, "value": value, **kwargs}) + return True + fake_redis = AsyncMock() fake_redis.async_increment = AsyncMock(side_effect=record_increment) fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing + fake_redis.async_set_cache = AsyncMock(side_effect=record_set) counter_cache.redis_cache = fake_redis # Prisma returns spend=42.0 (authoritative) while the stale cached @@ -5026,10 +5032,10 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( where={"team_id": "team-9"} ) - # Two increments keyed on the counter: seed ($42) then request ($1.50). + seed_writes = [(c["key"], c["value"]) for c in recorded_sets] + assert ("spend:team:team-9", 42.0) in seed_writes writes = [(c["key"], c["value"]) for c in recorded_increments] - assert ("spend:team:team-9", 42.0) in writes - assert ("spend:team:team-9", 1.5) in writes + assert writes == [("spend:team:team-9", 1.5)] finally: ps.user_api_key_cache = orig_user ps.spend_counter_cache = orig_counter