litellm/tests/litellm_utils_tests/test_proxy_budget_reset.py
ryan-crabbe-berri c40828509b
fix(reset_budget_job): atomic budget cascade with chunked reset scans (#36287)
* fix(reset_budget_job): advance budget_reset_at atomically with the spend cascade

A postgres timeout mid-cascade previously left LiteLLM_BudgetTable rows
stamped for the next window while team member, enduser, org and tag spend
stayed at cap, so every later tick skipped them until the window rolled
over. All cascade writes and the budget_reset_at advance now share one
prisma batch transaction; a failed run persists nothing and the rows stay
due for the next ~10 minute tick. Cache and counter invalidation runs only
after commit, and the catch-all enduser log line now names the cascade.

* fix(reset_budget_job): elect one runner per tick and chunk the reset scans

Every pod and worker previously ran the reset job every ~10 minutes,
each fetching every expired row with no limit and writing one giant
transaction at the same calendar-aligned boundary; that concurrency is
what piled up postgres lock contention and timeouts. The job now takes
the shared PodLockManager redis lock (no redis keeps the old behavior),
and each phase walks its due rows in 500-row chunks, one transaction per
chunk, stopping when a chunk is short, makes no forward progress, or
hits the per-run cap; leftovers wait for the next tick.

* chore(lint): ratchet budget ceilings down for fixed violations

* fix(reset_budget_job): harden chunk loop, fail open on redis errors, heartbeat the lock

Review fixes on the two prior commits. Reset scans now skip rows with no
budget_duration, so permanently due rows can neither starve a phase nor
have a lifetime cap zeroed every tick. Chunk progress counts rows whose
new budget_reset_at actually cleared the cutoff, so a zero-length
duration cannot burn the per-run chunk cap. A failed lock acquire only
skips the run when another pod verifiably holds the lock; a broken redis
runs unguarded instead of silently disabling resets fleet-wide. Partial
row failures report real progress and fire the failure hook without
killing the phase. The leader re-asserts the lock between phases and
stops if another pod took over, and the budget window advance uses
update_many so a tier deleted mid-chunk cannot abort the transaction.
Lint budget ceilings re-ratcheted for the net-fixed violations.

* fix(reset_budget_job): renew the leader lease and reject non-positive budget durations

Bot review follow-ups. PodLockManager now extends the lock TTL when the
holding pod re-acquires, via an atomic compare-and-expire script with a
plain SET fallback, so a run longer than the TTL keeps its lease instead
of silently sharing the job with another pod. The positive-duration
validation that team member endpoints already had is hoisted to
management common_utils and applied to key, internal user, budget,
customer and team intake, so a tenant can no longer create zero-duration
budgets whose permanently due rows starve other tenants' resets. Such
durations now return 400 at intake; existing rows are untouched.

* refactor(reset_budget_job): defer leader election to a follow-up PR

* fix(reset_budget_job): satisfy strict lint gates

String defaults for the two getenv calls (PLW1508) and the chunk
outcome returns moved to try/else (TRY300).
2026-08-10 14:42:36 -07:00

1168 lines
44 KiB
Python

import asyncio
import os
import sys
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from dotenv import load_dotenv
load_dotenv()
import os
from litellm.proxy._types import LiteLLM_BudgetTableFull
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
# Note: In our "fake" items we use dicts with fields that our fake reset functions modify.
# 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, fail_commit=False):
"""
Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is
awaitable and whose per-table .update()/.update_many() calls get captured.
The reset job writes every reset through prisma.db.batch_() — key/user/team
rows one by one, and the budget tier's cascade as a single transaction — so
tests must let that batch path complete.
Only committed batches contribute to the returned list, mirroring prisma:
with fail_commit=True the transaction blows up and must persist nothing.
Returns the list that will accumulate {table, op, where, data} dicts from
each captured write.
"""
batch_calls = []
def make_batcher():
queued = []
class _Table:
def __init__(self, table_name):
self._table_name = table_name
def update(self, where=None, data=None):
queued.append(
{
"table": self._table_name,
"op": "update",
"where": where,
"data": data,
}
)
def update_many(self, where=None, data=None):
queued.append(
{
"table": self._table_name,
"op": "update_many",
"where": where,
"data": data,
}
)
async def commit():
if fail_commit:
raise RuntimeError("simulated Postgres failure committing the batch")
batch_calls.extend(queued)
batcher = MagicMock()
batcher.litellm_verificationtoken = _Table("key")
batcher.litellm_usertable = _Table("user")
batcher.litellm_teamtable = _Table("team")
batcher.litellm_budgettable = _Table("budget")
batcher.litellm_teammembership = _Table("team_membership")
batcher.litellm_organizationtable = _Table("org")
batcher.litellm_tagtable = _Table("tag")
batcher.litellm_endusertable = _Table("enduser")
batcher.commit = commit
return batcher
prisma_client.db.batch_ = MagicMock(side_effect=make_batcher)
return batch_calls
def _wire_cascade_reads_for_test(prisma_client):
"""
The budget tier's cascade reads the rows it is about to zero, so their
spend counters can be invalidated after the commit. Give each of those
tables an awaitable find_many so the reads resolve instead of falling into
the job's warn-and-continue path.
"""
for table in (
"litellm_teammembership",
"litellm_verificationtoken",
"litellm_organizationtable",
"litellm_tagtable",
"litellm_endusertable",
):
getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[])
@pytest.mark.asyncio
async def test_reset_budget_keys_partial_failure():
"""
Test that if one key fails to reset, the failure for that key does not block processing of the other keys.
We simulate two keys where the first fails and the second succeeds.
"""
# Arrange
key1 = {
"id": "key1",
"spend": 10.0,
"budget_duration": 60,
} # Will trigger simulated failure
key2 = {"id": "key2", "spend": 15.0, "budget_duration": 60} # Should be updated
key3 = {"id": "key3", "spend": 20.0, "budget_duration": 60} # Should be updated
key4 = {"id": "key4", "spend": 25.0, "budget_duration": 60} # Should be updated
key5 = {"id": "key5", "spend": 30.0, "budget_duration": 60} # Should be updated
key6 = {"id": "key6", "spend": 35.0, "budget_duration": 60} # Should be updated
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(
return_value=[key1, key2, key3, key4, key5, key6]
)
prisma_client.update_data = AsyncMock()
# Reset job writes key resets via prisma.db.batch_().<table>.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()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
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, reset_settings=None):
if key["id"] == "key1":
# Simulate a failure on key1 (for example, this might be due to an invariant check)
raise Exception("Simulated failure for key1")
else:
# Simulate successful reset modification
key["spend"] = 0.0
# Compute a new reset time based on the budget duration
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
).isoformat()
return key
with patch.object(
ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key
) as mock_reset_key:
# Call the method; even though one key fails, the loop should process both
await job.reset_budget_for_litellm_keys()
# Allow any created tasks (logging hooks) to schedule
await asyncio.sleep(0.1)
# Assert that the helper was called for 6 keys
assert mock_reset_key.call_count == 6
# 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 = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
)
# There should be one failure hook call for keys (with call_type "reset_budget_keys")
assert any(
call.kwargs.get("call_type") == "reset_budget_keys"
for call in failure_hook_calls
)
@pytest.mark.asyncio
async def test_reset_budget_users_partial_failure():
"""
Test that if one user fails to reset, the reset loop still processes the other users.
We simulate two users where the first fails and the second is updated.
"""
user1 = {
"id": "user1",
"spend": 20.0,
"budget_duration": 120,
} # Will trigger simulated failure
user2 = {"id": "user2", "spend": 25.0, "budget_duration": 120} # Should be updated
user3 = {"id": "user3", "spend": 30.0, "budget_duration": 120} # Should be updated
user4 = {"id": "user4", "spend": 35.0, "budget_duration": 120} # Should be updated
user5 = {"id": "user5", "spend": 40.0, "budget_duration": 120} # Should be updated
user6 = {"id": "user6", "spend": 45.0, "budget_duration": 120} # Should be updated
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(
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()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
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, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
else:
user["spend"] = 0.0
user["budget_reset_at"] = (
current_time + timedelta(seconds=user["budget_duration"])
).isoformat()
return user
with patch.object(
ResetBudgetJob, "_reset_budget_for_user", side_effect=fake_reset_user
) as mock_reset_user:
await job.reset_budget_for_litellm_users()
await asyncio.sleep(0.1)
assert mock_reset_user.call_count == 6
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
)
assert any(
call.kwargs.get("call_type") == "reset_budget_users"
for call in failure_hook_calls
)
@pytest.mark.asyncio
async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing():
"""
A failure anywhere in the budget-tier cascade must persist nothing, so the
tier stays due and the next scheduler tick retries it. Before the fix the
job committed the new budget_reset_at first and zeroed the dependent spend
afterwards, so a failure here left the tier stamped for the next window
while every end user stayed at the cap.
"""
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
prisma_client = MagicMock()
async def get_data_mock(table_name, *args, **kwargs):
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
batch_calls = _wire_batcher_for_test(prisma_client, fail_commit=True)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert batch_calls == [], "a failed cascade must not persist any write"
assert (
prisma_client.update_data.await_count == 0
), "budget_reset_at must not be advanced outside the cascade transaction"
failure_hook_calls = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
)
assert any(
call.kwargs.get("call_type") == "reset_budget_endusers"
for call in failure_hook_calls
)
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance():
"""
The happy path: every end user the tier gates is zeroed and the tier's
budget_reset_at advances, all inside one transaction.
"""
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
prisma_client = MagicMock()
async def get_data_mock(table_name, *args, **kwargs):
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
batch_calls = _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction"
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"]["user_id"]["in"] == [f"user{i}" for i in range(1, 7)]
assert enduser_writes[0]["data"] == {"spend": 0}
budget_writes = [c for c in batch_calls if c["table"] == "budget"]
assert len(budget_writes) == 1
assert budget_writes[0]["where"] == {"budget_id": "budget1"}
assert budget_writes[0]["data"]["budget_reset_at"] > datetime.now(timezone.utc)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
async def test_reset_budget_teams_partial_failure():
"""
Test that if one team fails to reset, the loop processes both teams and only updates the ones that succeeded.
We simulate two teams where the first fails and the second is updated.
"""
team1 = {
"id": "team1",
"spend": 30.0,
"budget_duration": 180,
} # Will trigger simulated failure
team2 = {"id": "team2", "spend": 35.0, "budget_duration": 180} # Should be updated
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()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
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, reset_settings=None):
if team["id"] == "team1":
raise Exception("Simulated failure for team1")
else:
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
).isoformat()
return team
with patch.object(
ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team
) as mock_reset_team:
await job.reset_budget_for_litellm_teams()
await asyncio.sleep(0.1)
assert mock_reset_team.call_count == 2
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
)
assert any(
call.kwargs.get("call_type") == "reset_budget_teams"
for call in failure_hook_calls
)
@pytest.mark.asyncio
async def test_reset_budget_continues_other_categories_on_failure():
"""
Test that executing the overall reset_budget() method continues to process keys, users, and teams,
even if one of the sub-categories (here, users) experiences a partial failure.
In this simulation:
- All keys are processed successfully.
- One of the two users fails.
- All teams are processed successfully.
We then assert that:
- update_data is called for each category with the correctly updated items.
- Each get_data call is made (indicating that one failing category did not abort the others).
"""
# Arrange dummy items for each table
key1 = {"id": "key1", "spend": 10.0, "budget_duration": 60}
key2 = {"id": "key2", "spend": 15.0, "budget_duration": 60}
user1 = {
"id": "user1",
"spend": 20.0,
"budget_duration": 120,
} # Will fail in user reset
user2 = {"id": "user2", "spend": 25.0, "budget_duration": 120} # Succeeds
team1 = {"id": "team1", "spend": 30.0, "budget_duration": 180}
team2 = {"id": "team2", "spend": 35.0, "budget_duration": 180}
enduser1 = {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
prisma_client = MagicMock()
async def fake_get_data(*, table_name, query_type, **kwargs):
if table_name == "key":
return [key1, key2]
elif table_name == "user":
return [user1, user2]
elif table_name == "team":
return [team1, team2]
elif table_name == "budget":
return [budget1]
elif table_name == "enduser":
return [enduser1]
return []
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)
enduser1 = _attrify(enduser1)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time, reset_settings=None):
key["spend"] = 0.0
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
).isoformat()
return key
async def fake_reset_user(user, current_time, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
user["spend"] = 0.0
user["budget_reset_at"] = (
current_time + timedelta(seconds=user["budget_duration"])
).isoformat()
return user
async def fake_reset_team(team, current_time, reset_settings=None):
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
).isoformat()
return team
with (
patch.object(
ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key
) as mock_reset_key,
patch.object(
ResetBudgetJob, "_reset_budget_for_user", side_effect=fake_reset_user
) as mock_reset_user,
patch.object(
ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team
) as mock_reset_team,
):
# Call the overall reset_budget method.
await job.reset_budget()
await asyncio.sleep(0.1)
# Verify that get_data was called for each table. We can check the table names across calls.
called_tables = {
call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list
}
assert called_tables == {"key", "user", "team", "budget", "enduser"}
# Every category writes through the batch path now, so update_data is unused.
prisma_client.update_data.assert_not_awaited()
# The budget tier's cascade still ran despite the failing user category.
assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1"]}}
assert enduser_writes[0]["data"] == {"spend": 0}
# Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
# `op` separates the per-row resets from the cascade sweep, which also
# targets the key table.
key_writes = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"]
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)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_service_logger_keys_success():
"""
Test that when resetting keys succeeds (all keys are updated) the service
logger success hook is called with the correct event metadata and no exception is logged.
"""
keys = [
_attrify(
{"id": "key1", "spend": 10.0, "budget_duration": 60, "token": "key1"}
),
_attrify(
{"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()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time, reset_settings=None):
key["spend"] = 0.0
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
).isoformat()
return key
with patch.object(
ResetBudgetJob,
"_reset_budget_for_key",
side_effect=fake_reset_key,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_keys()
# Allow async logging task to complete
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
# Verify success hook call
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_success_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_keys_found") == len(keys)
assert event_metadata.get("num_keys_updated") == len(keys)
assert event_metadata.get("num_keys_failed") == 0
# Failure hook should not be executed.
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_keys_failure():
"""
Test that when a key reset fails the service logger failure hook is called,
the event metadata reflects the number of keys processed, and that the verbose
logger exception is called.
"""
keys = [
{"id": "key1", "spend": 10.0, "budget_duration": 60},
{"id": "key2", "spend": 15.0, "budget_duration": 60},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=keys)
prisma_client.update_data = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time, reset_settings=None):
if key["id"] == "key1":
raise Exception("Simulated failure for key1")
key["spend"] = 0.0
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
).isoformat()
return key
with patch.object(
ResetBudgetJob,
"_reset_budget_for_key",
side_effect=fake_reset_key,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_keys()
await asyncio.sleep(0.1)
# Expect at least one exception logged (the inner error and the outer catch)
assert mock_verbose_exc.call_count >= 1
# Verify exception was logged with correct message
assert any(
"Failed to reset budget for key" in str(call.args)
for call in mock_verbose_exc.call_args_list
)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_keys_found") == len(keys)
keys_found_str = event_metadata.get("keys_found", "")
assert "key1" in keys_found_str
# Success hook should not be called.
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_users_success():
"""
Test that when resetting users succeeds the service logger success hook is called with
the correct metadata and no exception is logged.
"""
users = [
_attrify(
{"id": "user1", "spend": 20.0, "budget_duration": 120, "user_id": "user1"}
),
_attrify(
{"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()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_user(user, current_time, reset_settings=None):
user["spend"] = 0.0
user["budget_reset_at"] = (
current_time + timedelta(seconds=user["budget_duration"])
).isoformat()
return user
with patch.object(
ResetBudgetJob,
"_reset_budget_for_user",
side_effect=fake_reset_user,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_users()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_success_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_users_found") == len(users)
assert event_metadata.get("num_users_updated") == len(users)
assert event_metadata.get("num_users_failed") == 0
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_users_failure():
"""
Test that a failure during user reset calls the failure hook with appropriate metadata,
logs the exception, and does not call the success hook.
"""
users = [
{"id": "user1", "spend": 20.0, "budget_duration": 120},
{"id": "user2", "spend": 25.0, "budget_duration": 120},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=users)
prisma_client.update_data = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_user(user, current_time, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
user["spend"] = 0.0
user["budget_reset_at"] = (
current_time + timedelta(seconds=user["budget_duration"])
).isoformat()
return user
with patch.object(
ResetBudgetJob,
"_reset_budget_for_user",
side_effect=fake_reset_user,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_users()
await asyncio.sleep(0.1)
# Verify exception logging
assert mock_verbose_exc.call_count >= 1
# Verify exception was logged with correct message
assert any(
"Failed to reset budget for user" in str(call.args)
for call in mock_verbose_exc.call_args_list
)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_users_found") == len(users)
users_found_str = event_metadata.get("users_found", "")
assert "user1" in users_found_str
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_teams_success():
"""
Test that when resetting teams is successful the service logger success hook is called with
the proper metadata and nothing is logged as an exception.
"""
teams = [
_attrify(
{"id": "team1", "spend": 30.0, "budget_duration": 180, "team_id": "team1"}
),
_attrify(
{"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()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_team(team, current_time, reset_settings=None):
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
).isoformat()
return team
with patch.object(
ResetBudgetJob,
"_reset_budget_for_team",
side_effect=fake_reset_team,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_teams()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_success_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_teams_found") == len(teams)
assert event_metadata.get("num_teams_updated") == len(teams)
assert event_metadata.get("num_teams_failed") == 0
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_teams_failure():
"""
Test that a failure during team reset triggers the failure hook with proper metadata,
results in an exception log and no success hook call.
"""
teams = [
{"id": "team1", "spend": 30.0, "budget_duration": 180},
{"id": "team2", "spend": 35.0, "budget_duration": 180},
]
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(return_value=teams)
prisma_client.update_data = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_team(team, current_time, reset_settings=None):
if team["id"] == "team1":
raise Exception("Simulated failure for team1")
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
).isoformat()
return team
with patch.object(
ResetBudgetJob,
"_reset_budget_for_team",
side_effect=fake_reset_team,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_teams()
await asyncio.sleep(0.1)
# Verify exception logging
assert mock_verbose_exc.call_count >= 1
# Verify exception was logged with correct message
assert any(
"Failed to reset budget for team" in str(call.args)
for call in mock_verbose_exc.call_args_list
)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_teams_found") == len(teams)
teams_found_str = event_metadata.get("teams_found", "")
assert "team1" in teams_found_str
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_endusers_success():
"""
Test that when the budget-tier cascade commits, the service logger success
hook is called with the correct metadata and no exception is logged.
"""
endusers = [
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
]
async def fake_get_data(*, table_name, query_type, **kwargs):
if table_name == "budget":
return budgets
elif table_name == "enduser":
return endusers
return []
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1", "user2"]}}
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_success_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_budgets_found") == len(budgets)
assert event_metadata.get("num_endusers_found") == len(endusers)
assert event_metadata.get("num_endusers_updated") == len(endusers)
assert event_metadata.get("num_endusers_failed") == 0
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
async def test_service_logger_endusers_failure():
"""
Test that a failed cascade calls the failure hook with the rows it had
found, logs the exception, and does not call the success hook.
"""
endusers = [
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
]
async def fake_get_data(*, table_name, query_type, **kwargs):
if table_name == "budget":
return budgets
elif table_name == "enduser":
return endusers
return []
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
_wire_batcher_for_test(prisma_client, fail_commit=True)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
# The log must name the whole cascade, not just end users: the write
# that failed could have been any of team member / enduser / org / tag
# spend or the budget_reset_at advance.
assert mock_verbose_exc.call_count == 1
assert "budget table cascade" in str(mock_verbose_exc.call_args.args[0])
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
args,
kwargs,
) = proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args
event_metadata = kwargs.get("event_metadata", {})
assert event_metadata.get("num_budgets_found") == len(budgets)
assert event_metadata.get("num_endusers_found") == len(endusers)
endusers_found_str = event_metadata.get("endusers_found", "")
assert "user1" in endusers_found_str
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_reset_budget_for_litellm_team_members_called():
"""
Test that when reset_budget_for_litellm_budget_table is called, team
members' spend is zeroed as part of the cascade transaction.
"""
# Arrange
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 100.0,
"budget_duration": "1d",
"created_at": datetime.now(timezone.utc) - timedelta(days=2),
}
)
enduser1 = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"})
prisma_client = MagicMock()
async def fake_get_data(*, table_name, query_type, **kwargs):
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return [enduser1]
return []
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
prisma_client.db = MagicMock()
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
# Act
await job.reset_budget_for_litellm_budget_table()
# Assert
team_member_writes = [c for c in batch_calls if c["table"] == "team_membership"]
assert len(team_member_writes) == 1
assert team_member_writes[0]["where"]["budget_id"]["in"] == ["budget1"]
assert team_member_writes[0]["data"] == {"spend": 0}