fix(reset_budget_job): reset end users by budget link, not by user id (#40639)

Adapted for stable/1.98.x: this line predates budget rollover (#38514), so the fix is applied to _commit_budget_cascade_once directly. End users reset on the budget link plus a NULL budget_id branch for the default tier, which is what upstream's _queue_enduser_resets does with rollover off. The rollover test hunk is dropped.

(cherry picked from commit 8a4fae0e17)
This commit is contained in:
ryan-crabbe-berri 2026-09-10 17:31:25 -07:00 • committed by Yuneng Jiang
parent 54fe3f5522
commit f2c9c6e818
No known key found for this signature in database
3 changed files with 68 additions and 17 deletions

View file

@ -394,14 +394,14 @@ class ResetBudgetJob:
if not cascade.budget_ids:
return
enduser_ids: Final = tuple(row.user_id for row in cascade.endusers)
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids))
uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE))
uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if enduser_ids:
uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}})
uow.endusers.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if litellm.max_end_user_budget_id in cascade.budget_ids:
uow.endusers.queue_spend_zero(where={"budget_id": None, **_SPENT_ROWS_WHERE})
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)

View file

@ -407,7 +407,7 @@ async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance()
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]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}}
assert enduser_writes[0]["data"] == {"spend": 0}
budget_writes = [c for c in batch_calls if c["table"] == "budget"]
@ -608,7 +608,7 @@ async def test_reset_budget_continues_other_categories_on_failure():
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]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}}
assert enduser_writes[0]["data"] == {"spend": 0}
# Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
@ -1038,7 +1038,7 @@ async def test_service_logger_endusers_success():
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"]}}
assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}}
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(

View file

@ -5,7 +5,7 @@ import sys
import types
from datetime import datetime, timedelta, timezone
from datetime import time as dt_time
from typing import Any, Dict, List
from typing import Any, Dict, Final, List
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -396,7 +396,7 @@ def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client):
{
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["test-enduser-1"]}},
"where": {"budget_id": {"in": ["test-budget-1"]}, "spend": {"gt": 0}},
"data": {"spend": 0},
}
]
@ -497,7 +497,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
{
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["test-enduser-1"]}},
"where": {"budget_id": {"in": ["test-budget-1"]}, "spend": {"gt": 0}},
"data": {"spend": 0},
}
]
@ -516,6 +516,7 @@ _LINKED_TABLE_CASES = [
),
("org", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
("tag", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
("enduser", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
]
@ -545,6 +546,48 @@ def test_budget_table_reset_zeroes_spend_on_every_linked_table(
assert writes[0]["data"] == {"spend": 0}
_POSTGRES_MAX_BIND_VARIABLES: Final = 32767
def _bind_count(where: Dict[str, Any]) -> int:
"""Bind variables one prisma where-clause compiles to: each scalar is one
placeholder and an ``in`` list contributes one per element."""
return sum(len(value["in"]) if isinstance(value, dict) and "in" in value else 1 for value in where.values())
@pytest.mark.parametrize("population", [3, 40_000], ids=["small", "over-pg-bind-ceiling"])
def test_enduser_reset_bind_count_does_not_scale_with_population(reset_budget_job, mock_prisma_client, population):
"""Regression for #40564.
Enumerating every dependent user id put one bind variable per customer into
a single prepared statement. Past PostgreSQL's ceiling the statement could
not be parsed at all, so the whole atomic cascade rolled back,
budget_reset_at never advanced, and the tier stayed due on every later tick
forever. Matching on the budget link keeps the statement the same size no
matter how many customers share a tier.
"""
budget = _budget_row(budget_id="shared-tier", budget_duration="1d")
mock_prisma_client.data["budget"] = [budget]
mock_prisma_client.data["enduser"] = [
types.SimpleNamespace(
spend=1.0,
litellm_budget_table=budget,
user_id=f"cust-{index:08d}",
budget_id="shared-tier",
)
for index in range(population)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
writes = _batch_writes(mock_prisma_client, "enduser")
assert [_bind_count(write["where"]) for write in writes] == [2], (
f"the cascade must not enumerate {population} user ids: past "
f"{_POSTGRES_MAX_BIND_VARIABLES} binds PostgreSQL refuses the statement, got {writes[:1]}"
)
assert _batch_writes(mock_prisma_client, "budget")[0]["data"]["budget_reset_at"] is not None
def test_budget_table_reset_writes_nothing_when_no_budget_is_due(reset_budget_job, mock_prisma_client):
"""Nothing due means no transaction is opened at all."""
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
@ -712,14 +755,22 @@ def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Both end users are zeroed by the same committed statement.
enduser_writes = _batch_writes(mock_prisma_client, "enduser")
assert len(enduser_writes) == 1, f"Expected a single enduser write, got {enduser_writes}"
assert set(enduser_writes[0]["where"]["user_id"]["in"]) == {
"enduser-explicit",
"enduser-implicit",
}
assert enduser_writes[0]["data"] == {"spend": 0}
# Both end users are zeroed: the linked rows on the tier's budget_id, the
# implicit ones on the NULL branch that stands in for the default tier.
assert _batch_writes(mock_prisma_client, "enduser") == [
{
"table": "enduser",
"op": "update_many",
"where": {"budget_id": {"in": [default_budget_id]}, "spend": {"gt": 0}},
"data": {"spend": 0},
},
{
"table": "enduser",
"op": "update_many",
"where": {"budget_id": None, "spend": {"gt": 0}},
"data": {"spend": 0},
},
]
# Verify find_many was called to fetch NULL-budget-id end users
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls