mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(budget): reseed model counters on miss and avoid baseline double-count
This commit is contained in:
parent
7ea753e415
commit
41d2cc1012
4 changed files with 68 additions and 32 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue