fix(reset): clear team-member model spend on budget rollover

This commit is contained in:
Ishaan Jaffer 2026-04-25 10:56:59 -07:00
parent 4a87c6babe
commit e52098e1d4
No known key found for this signature in database
2 changed files with 59 additions and 0 deletions

View file

@ -63,6 +63,7 @@ class ResetBudgetJob:
for budget in budgets_to_reset
if budget.budget_id is not None
]
memberships: List[Any] = []
# Reset spend counters for affected team members.
# Reset Redis directly so a transient failure doesn't leave stale
@ -70,6 +71,7 @@ class ResetBudgetJob:
try:
from litellm.proxy.proxy_server import (
_get_team_member_counter_key,
_get_team_member_model_counter_key,
spend_counter_cache,
)
@ -98,6 +100,41 @@ class ResetBudgetJob:
counter_key,
redis_err,
)
if memberships:
member_pairs = [
{"user_id": m.user_id, "team_id": m.team_id} for m in memberships
]
model_spend_rows = (
await self.prisma_client.db.litellm_teammembermodelspend.find_many(
where={"OR": member_pairs}
)
)
for row in model_spend_rows:
model_counter_key = _get_team_member_model_counter_key(
user_id=row.user_id,
team_id=row.team_id,
model=row.model,
)
spend_counter_cache.in_memory_cache.set_cache(
key=model_counter_key, value=0.0
)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(
key=model_counter_key, value=0.0
)
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to reset team member model spend counter in Redis %s: %s. "
"Budget may be over-enforced until counter expires.",
model_counter_key,
redis_err,
)
await self.prisma_client.db.litellm_teammembermodelspend.update_many(
where={"OR": member_pairs},
data={"spend": 0},
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to reset team member spend counters: %s", e

View file

@ -836,6 +836,11 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format():
membership.user_id = "user-1"
membership.team_id = "team-1"
model_spend_row = MagicMock()
model_spend_row.user_id = "user-1"
model_spend_row.team_id = "team-1"
model_spend_row.model = "gpt-4o"
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(
return_value=[membership]
@ -843,6 +848,12 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format():
mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 1}
)
mock_prisma_client.db.litellm_teammembermodelspend.find_many = AsyncMock(
return_value=[model_spend_row]
)
mock_prisma_client.db.litellm_teammembermodelspend.update_many = AsyncMock(
return_value={"count": 1}
)
fake_module = types.ModuleType("litellm.proxy.proxy_server")
spend_counter_cache = MagicMock()
@ -852,6 +863,9 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format():
fake_module._get_team_member_counter_key = (
lambda user_id, team_id: f"spend:team_member:team_id::{team_id}::user_id::{user_id}"
)
fake_module._get_team_member_model_counter_key = (
lambda user_id, team_id, model: f"spend:team_member:team_id::{team_id}::user_id::{user_id}::model::{model}"
)
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_module}):
job = ResetBudgetJob(
@ -863,6 +877,14 @@ def test_reset_budget_for_team_members_resets_new_counter_key_format():
key="spend:team_member:team_id::team-1::user_id::user-1",
value=0.0,
)
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team_member:team_id::team-1::user_id::user-1::model::gpt-4o",
value=0.0,
)
mock_prisma_client.db.litellm_teammembermodelspend.update_many.assert_awaited_once_with(
where={"OR": [{"user_id": "user-1", "team_id": "team-1"}]},
data={"spend": 0},
)
# ---------------------------------------------------------------------------