fix(budget): reseed model counters on miss and avoid baseline double-count

This commit is contained in:
Ishaan Jaffer 2026-04-25 13:59:43 -07:00
parent 7ea753e415
commit 41d2cc1012
No known key found for this signature in database
4 changed files with 68 additions and 32 deletions

View file

@ -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(

View file

@ -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)

View file

@ -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,
)

View file

@ -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