diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py
index 52bbeaf2ad..40c8caa49e 100644
--- a/litellm/proxy/common_utils/reset_budget_job.py
+++ b/litellm/proxy/common_utils/reset_budget_job.py
@@ -414,6 +414,72 @@ class ResetBudgetJob:
)
return [LiteLLM_EndUserTable(**row.dict()) for row in rows]
+ async def _write_key_reset_updates(
+ self, updated_keys: List[LiteLLM_VerificationToken]
+ ) -> None:
+ """
+ Write per-row {spend, budget_reset_at} updates for keys.
+
+ Avoids the batched full-model update path, which trips
+ prisma.errors.DataError on any row carrying object_permission_id or
+ budget_limits (see #27730). Both fields are rejected by Prisma's
+ update input type for LiteLLM_VerificationToken, and the failure
+ aborts the entire batch — silently leaving spend over the cap and
+ budget_reset_at unchanged forever.
+ """
+ batcher = self.prisma_client.db.batch_()
+ for k in updated_keys:
+ token = getattr(k, "token", None)
+ if token is None:
+ continue
+ batcher.litellm_verificationtoken.update(
+ where={"token": token},
+ data={"spend": 0, "budget_reset_at": k.budget_reset_at},
+ )
+ await batcher.commit()
+
+ async def _write_user_reset_updates(
+ self, updated_users: List[LiteLLM_UserTable]
+ ) -> None:
+ """
+ Write per-row {spend, budget_reset_at} updates for users.
+
+ Mirrors _write_key_reset_updates — avoids the full-model update path
+ that trips Prisma's DataError on rows carrying unrecognised fields
+ (see #27730).
+ """
+ batcher = self.prisma_client.db.batch_()
+ for u in updated_users:
+ user_id = getattr(u, "user_id", None)
+ if user_id is None:
+ continue
+ batcher.litellm_usertable.update(
+ where={"user_id": user_id},
+ data={"spend": 0, "budget_reset_at": u.budget_reset_at},
+ )
+ await batcher.commit()
+
+ async def _write_team_reset_updates(
+ self, updated_teams: List[LiteLLM_TeamTable]
+ ) -> None:
+ """
+ Write per-row {spend, budget_reset_at} updates for teams.
+
+ Mirrors _write_key_reset_updates — avoids the full-model update path
+ that trips Prisma's DataError on rows carrying unrecognised fields
+ (see #27730).
+ """
+ batcher = self.prisma_client.db.batch_()
+ for t in updated_teams:
+ team_id = getattr(t, "team_id", None)
+ if team_id is None:
+ continue
+ batcher.litellm_teamtable.update(
+ where={"team_id": team_id},
+ data={"spend": 0, "budget_reset_at": t.budget_reset_at},
+ )
+ await batcher.commit()
+
async def reset_budget_for_litellm_keys(self):
"""
Resets the budget for all the litellm keys
@@ -455,11 +521,7 @@ class ResetBudgetJob:
)
if updated_keys:
- await self.prisma_client.update_data(
- query_type="update_many",
- data_list=updated_keys,
- table_name="key",
- )
+ await self._write_key_reset_updates(updated_keys=updated_keys)
for k in updated_keys:
token = getattr(k, "token", None)
if token:
@@ -544,11 +606,7 @@ class ResetBudgetJob:
"Updated users %s", json.dumps(updated_users, indent=4, default=str)
)
if updated_users:
- await self.prisma_client.update_data(
- query_type="update_many",
- data_list=updated_users,
- table_name="user",
- )
+ await self._write_user_reset_updates(updated_users=updated_users)
for u in updated_users:
user_id = getattr(u, "user_id", None)
if user_id:
@@ -641,11 +699,7 @@ class ResetBudgetJob:
"Updated teams %s", json.dumps(updated_teams, indent=4, default=str)
)
if updated_teams:
- await self.prisma_client.update_data(
- query_type="update_many",
- data_list=updated_teams,
- table_name="team",
- )
+ await self._write_team_reset_updates(updated_teams=updated_teams)
for t in updated_teams:
team_id = getattr(t, "team_id", None)
if team_id:
@@ -816,49 +870,16 @@ class ResetBudgetJob:
"""
In-place, updates spend=0, and sets budget_reset_at to current_time + budget_duration
- Common logic for resetting budget for a team, user, or key
+ Common logic for resetting budget for a team, user, or key.
+
+ Spend-counter invalidation happens in the caller, AFTER the DB write
+ commits. Zeroing the counter here would open a bypass window when the
+ DB write fails: get_current_spend reads 0 from Redis while the DB
+ still holds the pre-reset value, admitting requests past the cap.
"""
try:
item.spend = 0.0
-
- # Reset the cross-pod spend counter.
- # Reset Redis directly (not via DualCache) so a Redis failure
- # doesn't silently leave a stale counter that get_current_spend
- # would read as authoritative, permanently blocking the user.
- from litellm.proxy.proxy_server import spend_counter_cache
-
- counter_key = None
- if item_type == "key" and hasattr(item, "token") and item.token is not None: # type: ignore[union-attr]
- counter_key = f"spend:key:{item.token}" # type: ignore[union-attr]
- elif (
- item_type == "team"
- and hasattr(item, "team_id")
- and item.team_id is not None # type: ignore[union-attr]
- ):
- counter_key = f"spend:team:{item.team_id}" # type: ignore[union-attr]
-
- if counter_key is not None:
- # Always reset in-memory (local fallback)
- spend_counter_cache.in_memory_cache.set_cache(
- key=counter_key, value=0.0
- )
- # Explicitly reset Redis with warning on failure
- if spend_counter_cache.redis_cache is not None:
- try:
- await spend_counter_cache.redis_cache.async_set_cache(
- key=counter_key, value=0.0
- )
- except Exception as redis_err:
- verbose_proxy_logger.warning(
- "Failed to reset spend counter in Redis for %s key=%s: %s. "
- "Budget may be over-enforced until counter expires.",
- item_type,
- counter_key,
- redis_err,
- )
-
if hasattr(item, "budget_duration") and item.budget_duration is not None:
- # Get standardized reset time based on budget duration
from litellm.proxy.common_utils.timezone_utils import (
get_budget_reset_time,
)
diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py
index 6240bedd3e..5c96eb619b 100644
--- a/tests/litellm_utils_tests/test_proxy_budget_reset.py
+++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py
@@ -22,6 +22,60 @@ from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
# In a real-world scenario, these would be instances of LiteLLM_VerificationToken, LiteLLM_UserTable, etc.
+def _attrify(d: dict):
+ """
+ Wrap a dict so that attribute access (`.token`, `.user_id`, `.team_id`,
+ etc.) works alongside the existing item-access the fake_reset_* helpers
+ rely on. The reset job's narrow-write helpers use `getattr(item, "token",
+ None)` (et al), which returns None for plain dicts — that would silently
+ skip the row.
+ """
+ class _AttrDict(dict):
+ def __getattr__(self, k):
+ try:
+ return self[k]
+ except KeyError:
+ raise AttributeError(k)
+
+ def __setattr__(self, k, v):
+ self[k] = v
+
+ return _AttrDict(d)
+
+
+def _wire_batcher_for_test(prisma_client):
+ """
+ Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is
+ awaitable and whose per-table .update() calls get captured. The reset job
+ writes key/user/team resets via prisma.db.batch_().
.update — not via
+ prisma_client.update_data — so tests must let that batch path complete.
+
+ Returns the list that will accumulate {table, where, data} dicts from
+ each captured update call.
+ """
+ batch_calls = []
+
+ def make_batcher():
+ class _Table:
+ def __init__(self, table_name):
+ self._table_name = table_name
+
+ def update(self, where=None, data=None):
+ batch_calls.append(
+ {"table": self._table_name, "where": where, "data": data}
+ )
+
+ batcher = MagicMock()
+ batcher.litellm_verificationtoken = _Table("key")
+ batcher.litellm_usertable = _Table("user")
+ batcher.litellm_teamtable = _Table("team")
+ batcher.commit = AsyncMock(return_value=None)
+ return batcher
+
+ prisma_client.db.batch_ = MagicMock(side_effect=make_batcher)
+ return batch_calls
+
+
@pytest.mark.asyncio
async def test_reset_budget_keys_partial_failure():
"""
@@ -45,6 +99,9 @@ async def test_reset_budget_keys_partial_failure():
return_value=[key1, key2, key3, key4, key5, key6]
)
prisma_client.update_data = AsyncMock()
+ # Reset job writes key resets via prisma.db.batch_()..update — not
+ # via update_data — so wire that path.
+ batch_calls = _wire_batcher_for_test(prisma_client)
# Using a dummy logging object with async hooks mocked out.
proxy_logging_obj = MagicMock()
@@ -56,6 +113,15 @@ async def test_reset_budget_keys_partial_failure():
now = datetime.utcnow()
+ # token is needed because the new write path uses where={"token": ...}
+ # and _AttrDict makes getattr work alongside item access used by fake_reset_key.
+ for k in [key1, key2, key3, key4, key5, key6]:
+ k.setdefault("token", k["id"])
+ key1, key2, key3, key4, key5, key6 = (
+ _attrify(k) for k in [key1, key2, key3, key4, key5, key6]
+ )
+ prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6])
+
async def fake_reset_key(key, current_time):
if key["id"] == "key1":
# Simulate a failure on key1 (for example, this might be due to an invariant check)
@@ -80,17 +146,17 @@ async def test_reset_budget_keys_partial_failure():
# Assert that the helper was called for 6 keys
assert mock_reset_key.call_count == 6
- # Assert that update_data was called once with a list containing all 6 keys
- prisma_client.update_data.assert_awaited_once()
- update_call = prisma_client.update_data.call_args
- assert update_call.kwargs.get("table_name") == "key"
- updated_keys = update_call.kwargs.get("data_list", [])
- assert len(updated_keys) == 5
- assert updated_keys[0]["id"] == "key2"
- assert updated_keys[1]["id"] == "key3"
- assert updated_keys[2]["id"] == "key4"
- assert updated_keys[3]["id"] == "key5"
- assert updated_keys[4]["id"] == "key6"
+ # Assert that the new narrow write path got 5 batched updates (key1 failed).
+ # update_data must NOT have been called for keys.
+ prisma_client.update_data.assert_not_awaited()
+ key_writes = [c for c in batch_calls if c["table"] == "key"]
+ assert len(key_writes) == 5
+ written_ids = [c["where"]["token"] for c in key_writes]
+ assert written_ids == ["key2", "key3", "key4", "key5", "key6"]
+ # And every write must carry only {spend, budget_reset_at} — never the full row.
+ for c in key_writes:
+ assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
+ assert c["data"]["spend"] == 0
# Verify that the failure logging hook was scheduled (due to the failure for key1)
failure_hook_calls = (
@@ -125,6 +191,7 @@ async def test_reset_budget_users_partial_failure():
return_value=[user1, user2, user3, user4, user5, user6]
)
prisma_client.update_data = AsyncMock()
+ batch_calls = _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@@ -133,6 +200,15 @@ async def test_reset_budget_users_partial_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
+ # user_id required for the new write path's where clause; _AttrDict so
+ # getattr(u, 'user_id') works alongside the dict access fake_reset_user uses.
+ for u in [user1, user2, user3, user4, user5, user6]:
+ u.setdefault("user_id", u["id"])
+ user1, user2, user3, user4, user5, user6 = (
+ _attrify(u) for u in [user1, user2, user3, user4, user5, user6]
+ )
+ prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6])
+
async def fake_reset_user(user, current_time):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
@@ -150,16 +226,14 @@ async def test_reset_budget_users_partial_failure():
await asyncio.sleep(0.1)
assert mock_reset_user.call_count == 6
- prisma_client.update_data.assert_awaited_once()
- update_call = prisma_client.update_data.call_args
- assert update_call.kwargs.get("table_name") == "user"
- updated_users = update_call.kwargs.get("data_list", [])
- assert len(updated_users) == 5
- assert updated_users[0]["id"] == "user2"
- assert updated_users[1]["id"] == "user3"
- assert updated_users[2]["id"] == "user4"
- assert updated_users[3]["id"] == "user5"
- assert updated_users[4]["id"] == "user6"
+ prisma_client.update_data.assert_not_awaited()
+ user_writes = [c for c in batch_calls if c["table"] == "user"]
+ assert len(user_writes) == 5
+ written_ids = [c["where"]["user_id"] for c in user_writes]
+ assert written_ids == ["user2", "user3", "user4", "user5", "user6"]
+ for c in user_writes:
+ assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
+ assert c["data"]["spend"] == 0
failure_hook_calls = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
@@ -308,6 +382,7 @@ async def test_reset_budget_teams_partial_failure():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=[team1, team2])
prisma_client.update_data = AsyncMock()
+ batch_calls = _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@@ -316,6 +391,12 @@ async def test_reset_budget_teams_partial_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
+ # team_id required for the new write path's where clause; _AttrDict for getattr.
+ for t in [team1, team2]:
+ t.setdefault("team_id", t["id"])
+ team1, team2 = _attrify(team1), _attrify(team2)
+ prisma_client.get_data = AsyncMock(return_value=[team1, team2])
+
async def fake_reset_team(team, current_time):
if team["id"] == "team1":
raise Exception("Simulated failure for team1")
@@ -333,12 +414,12 @@ async def test_reset_budget_teams_partial_failure():
await asyncio.sleep(0.1)
assert mock_reset_team.call_count == 2
- prisma_client.update_data.assert_awaited_once()
- update_call = prisma_client.update_data.call_args
- assert update_call.kwargs.get("table_name") == "team"
- updated_teams = update_call.kwargs.get("data_list", [])
- assert len(updated_teams) == 1
- assert updated_teams[0]["id"] == "team2"
+ prisma_client.update_data.assert_not_awaited()
+ team_writes = [c for c in batch_calls if c["table"] == "team"]
+ assert len(team_writes) == 1
+ assert team_writes[0]["where"] == {"team_id": "team2"}
+ assert set(team_writes[0]["data"].keys()) == {"spend", "budget_reset_at"}
+ assert team_writes[0]["data"]["spend"] == 0
failure_hook_calls = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
@@ -402,6 +483,18 @@ async def test_reset_budget_continues_other_categories_on_failure():
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
+ batch_calls = _wire_batcher_for_test(prisma_client)
+ # ID fields required by the new write path's where clauses; _AttrDict
+ # lets getattr() see them alongside the item-access fake_reset_* helpers use.
+ for k in [key1, key2]:
+ k.setdefault("token", k["id"])
+ for u in [user1, user2]:
+ u.setdefault("user_id", u["id"])
+ for t in [team1, team2]:
+ t.setdefault("team_id", t["id"])
+ key1, key2 = _attrify(key1), _attrify(key2)
+ user1, user2 = _attrify(user1), _attrify(user2)
+ team1, team2 = _attrify(team1), _attrify(team2)
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
@@ -488,32 +581,29 @@ async def test_reset_budget_continues_other_categories_on_failure():
"team_membership",
}
- # Verify that update_data was called three times (one per category, enduser update includes two)
- assert prisma_client.update_data.await_count == 5
+ # After the fix, keys/users/teams write via prisma.db.batch_()..update,
+ # so only budget + enduser still go through update_data.
calls = prisma_client.update_data.await_args_list
-
- # Check keys update: both keys succeed.
- keys_call = calls[0]
- assert keys_call.kwargs.get("table_name") == "key"
- assert len(keys_call.kwargs.get("data_list", [])) == 2
-
- # Check users update: only user2 succeeded.
- users_call = calls[1]
- assert users_call.kwargs.get("table_name") == "user"
- users_updated = users_call.kwargs.get("data_list", [])
- assert len(users_updated) == 1
- assert users_updated[0]["id"] == "user2"
-
- # Check teams update: both teams succeed.
- teams_call = calls[2]
- assert teams_call.kwargs.get("table_name") == "team"
- assert len(teams_call.kwargs.get("data_list", [])) == 2
+ update_data_tables = [c.kwargs.get("table_name") for c in calls]
+ assert sorted(update_data_tables) == ["budget", "enduser"]
# Check enduser update: enduser succeed.
- enduser_call = calls[4]
- assert enduser_call.kwargs.get("table_name") == "enduser"
+ enduser_call = next(c for c in calls if c.kwargs.get("table_name") == "enduser")
assert len(enduser_call.kwargs.get("data_list", [])) == 1
+ # Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
+ key_writes = [c for c in batch_calls if c["table"] == "key"]
+ user_writes = [c for c in batch_calls if c["table"] == "user"]
+ team_writes = [c for c in batch_calls if c["table"] == "team"]
+ assert len(key_writes) == 2
+ assert len(user_writes) == 1
+ assert user_writes[0]["where"] == {"user_id": "user2"}
+ assert len(team_writes) == 2
+ # Every batched write must carry only the two reset fields, never the full row.
+ for c in key_writes + user_writes + team_writes:
+ assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
+ assert c["data"]["spend"] == 0
+
# ---------------------------------------------------------------------------
# Additional tests for service logger behavior (keys, users, teams, endusers)
@@ -527,12 +617,13 @@ async def test_service_logger_keys_success():
logger success hook is called with the correct event metadata and no exception is logged.
"""
keys = [
- {"id": "key1", "spend": 10.0, "budget_duration": 60},
- {"id": "key2", "spend": 15.0, "budget_duration": 60},
+ {"id": "key1", "spend": 10.0, "budget_duration": 60, "token": "key1"},
+ {"id": "key2", "spend": 15.0, "budget_duration": 60, "token": "key2"},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=keys)
prisma_client.update_data = AsyncMock()
+ _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@@ -644,12 +735,13 @@ async def test_service_logger_users_success():
the correct metadata and no exception is logged.
"""
users = [
- {"id": "user1", "spend": 20.0, "budget_duration": 120},
- {"id": "user2", "spend": 25.0, "budget_duration": 120},
+ {"id": "user1", "spend": 20.0, "budget_duration": 120, "user_id": "user1"},
+ {"id": "user2", "spend": 25.0, "budget_duration": 120, "user_id": "user2"},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=users)
prisma_client.update_data = AsyncMock()
+ _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@@ -756,12 +848,13 @@ async def test_service_logger_teams_success():
the proper metadata and nothing is logged as an exception.
"""
teams = [
- {"id": "team1", "spend": 30.0, "budget_duration": 180},
- {"id": "team2", "spend": 35.0, "budget_duration": 180},
+ {"id": "team1", "spend": 30.0, "budget_duration": 180, "team_id": "team1"},
+ {"id": "team2", "spend": 35.0, "budget_duration": 180, "team_id": "team2"},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=teams)
prisma_client.update_data = AsyncMock()
+ _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
index 8a47c78db0..0b68374536 100644
--- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
+++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
@@ -92,6 +92,37 @@ class MockLiteLLMEndUserTable:
return self._find_many_results
+class MockBatcher:
+ """Captures per-row update calls and exposes them after commit().
+
+ Mirrors prisma's `db.batch_()` ergonomics enough that the reset job's
+ narrow-write helpers (`_write_key_reset_updates` et al) can run against
+ the mock and the test can assert on what would have been written.
+ """
+
+ def __init__(self):
+ self.calls: List[Dict[str, Any]] = []
+ self.committed: bool = False
+
+ class _Table:
+ def __init__(_self, table_name: str, outer: "MockBatcher"):
+ _self._table_name = table_name
+ _self._outer = outer
+
+ def update(_self, where, data):
+ _self._outer.calls.append(
+ {"table": _self._table_name, "where": where, "data": data}
+ )
+
+ self.litellm_verificationtoken = _Table("key", self)
+ self.litellm_usertable = _Table("user", self)
+ self.litellm_teamtable = _Table("team", self)
+
+ async def commit(self):
+ self.committed = True
+ return self.calls
+
+
class MockDB:
def __init__(self):
self.litellm_teammembership = MockLiteLLMTeamMembership()
@@ -99,6 +130,19 @@ class MockDB:
self.litellm_endusertable = MockLiteLLMEndUserTable()
self.litellm_organizationtable = MockLiteLLMOrganizationTable()
self.litellm_tagtable = MockLiteLLMTagTable()
+ self.batch_calls: List[Dict[str, Any]] = []
+
+ def batch_(self):
+ batcher = MockBatcher()
+ # Aggregate calls across all batches so tests can assert on cumulative writes.
+ original_commit = batcher.commit
+
+ async def _record_and_commit():
+ self.batch_calls.extend(batcher.calls)
+ return await original_commit()
+
+ batcher.commit = _record_and_commit # type: ignore[assignment]
+ return batcher
class MockPrismaClient:
@@ -205,6 +249,7 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
"budget_duration": "30d",
"budget_reset_at": now,
"id": "test-key-1",
+ "token": "tok-key-1",
},
)
@@ -213,11 +258,16 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
# Run the test
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
- # Verify results
- assert len(mock_prisma_client.updated_data["key"]) == 1
- updated_key = mock_prisma_client.updated_data["key"][0]
- assert updated_key.spend == 0.0
- assert updated_key.budget_reset_at > now
+ # The reset writes only {spend, budget_reset_at} per row via batch_().
+ # Full-row writes would re-detonate the Prisma DataError on rows carrying
+ # object_permission_id / budget_limits (see #27730).
+ key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
+ assert len(key_writes) == 1
+ write = key_writes[0]
+ assert write["where"] == {"token": "tok-key-1"}
+ assert write["data"]["spend"] == 0
+ assert write["data"]["budget_reset_at"] > now
+ assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
@@ -231,6 +281,7 @@ def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
"budget_duration": "7d",
"budget_reset_at": now,
"id": "test-user-1",
+ "user_id": "uid-1",
},
)
@@ -239,11 +290,13 @@ def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
# Run the test
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
- # Verify results
- assert len(mock_prisma_client.updated_data["user"]) == 1
- updated_user = mock_prisma_client.updated_data["user"][0]
- assert updated_user.spend == 0.0
- assert updated_user.budget_reset_at > now
+ user_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "user"]
+ assert len(user_writes) == 1
+ write = user_writes[0]
+ assert write["where"] == {"user_id": "uid-1"}
+ assert write["data"]["spend"] == 0
+ assert write["data"]["budget_reset_at"] > now
+ assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
@@ -257,6 +310,7 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
"budget_duration": "1mo",
"budget_reset_at": now,
"id": "test-team-1",
+ "team_id": "tid-1",
},
)
@@ -265,11 +319,13 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
# Run the test
asyncio.run(reset_budget_job.reset_budget_for_litellm_teams())
- # Verify results
- assert len(mock_prisma_client.updated_data["team"]) == 1
- updated_team = mock_prisma_client.updated_data["team"][0]
- assert updated_team.spend == 0.0
- assert updated_team.budget_reset_at > now
+ team_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "team"]
+ assert len(team_writes) == 1
+ write = team_writes[0]
+ assert write["where"] == {"team_id": "tid-1"}
+ assert write["data"]["spend"] == 0
+ assert write["data"]["budget_reset_at"] > now
+ assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client):
@@ -324,6 +380,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
"budget_duration": "30d",
"budget_reset_at": now,
"id": "test-key-1",
+ "token": "tok-all-1",
},
)
@@ -335,6 +392,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
"budget_duration": "7d",
"budget_reset_at": now,
"id": "test-user-1",
+ "user_id": "uid-all-1",
},
)
@@ -346,6 +404,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
"budget_duration": "1mo",
"budget_reset_at": now,
"id": "test-team-1",
+ "team_id": "tid-all-1",
},
)
@@ -379,17 +438,22 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
# Run the test
asyncio.run(reset_budget_job.reset_budget())
- # Verify results
- assert len(mock_prisma_client.updated_data["key"]) == 1
- assert len(mock_prisma_client.updated_data["user"]) == 1
- assert len(mock_prisma_client.updated_data["team"]) == 1
+ # key/user/team rows are written via batch_()..update — verify each
+ # one fired exactly once with the narrow {spend, budget_reset_at} payload.
+ for table_name, where in [
+ ("key", {"token": "tok-all-1"}),
+ ("user", {"user_id": "uid-all-1"}),
+ ("team", {"team_id": "tid-all-1"}),
+ ]:
+ writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == table_name]
+ assert len(writes) == 1, f"expected 1 {table_name} write, got {len(writes)}"
+ assert writes[0]["where"] == where
+ assert writes[0]["data"]["spend"] == 0
+ assert set(writes[0]["data"].keys()) == {"spend", "budget_reset_at"}
+
+ # Enduser + budget rows still go through update_data (not narrowed; different path).
assert len(mock_prisma_client.updated_data["enduser"]) == 1
assert len(mock_prisma_client.updated_data["budget"]) == 1
-
- # Check that all spends were reset to 0
- assert mock_prisma_client.updated_data["key"][0].spend == 0.0
- assert mock_prisma_client.updated_data["user"][0].spend == 0.0
- assert mock_prisma_client.updated_data["team"][0].spend == 0.0
assert mock_prisma_client.updated_data["enduser"][0].spend == 0.0
@@ -1399,6 +1463,105 @@ def test_reset_budget_for_teams_invalidates_redis_counter(
)
+def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
+ """
+ Regression for #27730 (the bypass-half).
+
+ If the DB write inside the reset job raises (e.g. Prisma DataError on a
+ row carrying object_permission_id or budget_limits), the Redis spend
+ counter MUST NOT be zeroed — that would let get_current_spend admit
+ requests past the cap while the DB row still holds the over-budget
+ spend.
+
+ Pre-fix: _reset_budget_common pre-zeroed the counter before the DB
+ write attempt, opening the bypass window.
+ Post-fix: counter invalidation lives in the caller, AFTER the DB write
+ commits. If the write raises, the post-write invalidation never runs.
+ """
+ counter_cache = _make_counter_invalidation_job(monkeypatch)
+
+ now = datetime.now(timezone.utc)
+ prisma_client = MagicMock()
+
+ matching_key = type(
+ "Key",
+ (),
+ {
+ "spend": 100.0,
+ "budget_duration": "30d",
+ "budget_reset_at": now - timedelta(seconds=1),
+ "token": "sk-failing",
+ },
+ )
+
+ # get_data returns one key needing reset; the batched DB write then explodes.
+ async def fake_get_data(table_name, query_type, **kwargs):
+ if table_name == "key":
+ return [matching_key]
+ return []
+
+ prisma_client.get_data = fake_get_data
+
+ batcher = MagicMock()
+ batcher.litellm_verificationtoken.update = MagicMock()
+
+ async def failing_commit():
+ raise RuntimeError("simulated Prisma DataError on update")
+
+ batcher.commit = failing_commit
+ prisma_client.db.batch_ = MagicMock(return_value=batcher)
+
+ job = ResetBudgetJob(
+ proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client
+ )
+
+ asyncio.run(job.reset_budget_for_litellm_keys())
+
+ # CRITICAL: counter invalidation must NOT have been called at all —
+ # the DB write raised before the post-write invalidation loop. Using
+ # assert_not_called() instead of iterating call_args_list, because the
+ # latter is vacuously true when the list is empty (would pass even if
+ # the bypass were re-introduced via a different code path).
+ counter_cache.in_memory_cache.set_cache.assert_not_called()
+
+
+def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, mock_prisma_client):
+ """
+ Regression for #27730 (the trigger-half).
+
+ The reset job must write only {spend, budget_reset_at} per row — never
+ the full key object. Sending the full object via the old update_data
+ batcher path made Prisma reject any row carrying object_permission_id
+ or budget_limits (both became non-NULL on UI-created keys after v1.84.0).
+ """
+ now = datetime.now(timezone.utc)
+ key_with_problematic_fields = type(
+ "LiteLLM_VerificationToken",
+ (),
+ {
+ "spend": 50.0,
+ "budget_duration": "30d",
+ "budget_reset_at": now,
+ "token": "sk-problematic",
+ "object_permission_id": "perm-abc", # would be rejected on update
+ "budget_limits": [{"max_budget": 5}], # would be rejected on update
+ "metadata": {"some": "thing"},
+ },
+ )
+ mock_prisma_client.data["key"] = [key_with_problematic_fields]
+
+ asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
+
+ key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
+ assert len(key_writes) == 1
+ payload_keys = set(key_writes[0]["data"].keys())
+ assert payload_keys == {"spend", "budget_reset_at"}, (
+ f"reset payload must not include any field besides spend / budget_reset_at, "
+ f"got: {payload_keys}. Any extra field (object_permission_id, budget_limits, etc.) "
+ f"trips Prisma DataError and detonates the whole batch."
+ )
+
+
def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monkeypatch):
"""Resetting keys via budget tier must clear each linked key's counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)