Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 22 additions & 24 deletions litellm/proxy/common_utils/reset_budget_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,33 +224,31 @@ def _queue_budget_linked_resets(


def _queue_enduser_resets(writes: LinkedSpendResetWrites, cascade: "_BudgetCascade") -> None:
"""End users are matched by id rather than budget link: rows with no
budget_id ride the default budget tier (litellm.max_end_user_budget_id).
Zero-before-decrement ordering matters here too (see
_queue_budget_linked_resets)."""
if not cascade.rollover_caps:
if cascade.endusers:
writes.queue_spend_zero(
where={"user_id": {"in": [row.user_id for row in cascade.endusers]}}
) # mutable-ok: prisma where filter must be a dict
"""End users reset on the budget link like every other gated table, plus a
NULL-budget_id branch: rows created implicitly persist no link and ride the
default tier (litellm.max_end_user_budget_id).

Matching on the link rather than enumerating user ids keeps a statement's
bind count proportional to the expiring tiers instead of the customer
population, which past ~32,700 dependents exceeds PostgreSQL's per-statement
bind ceiling and wedges the cascade permanently (#40564).
"""
_queue_budget_linked_resets(writes, cascade, extra=_SPENT_ROWS_WHERE)
default_budget_id: Final = litellm.max_end_user_budget_id
if default_budget_id is None or default_budget_id not in cascade.budget_ids:
return
tiered: Final = tuple((row.budget_id or litellm.max_end_user_budget_id, row.user_id) for row in cascade.endusers)
for budget_id, cap in cascade.rollover_caps.items():
if not (
user_ids := [uid for bid, uid in tiered if bid == budget_id]
): # mutable-ok: prisma "in" filter takes a list
continue
cap: Final = cascade.rollover_caps.get(default_budget_id)
if cap is None:
writes.queue_spend_zero(
where={"user_id": {"in": user_ids}, "spend": {"lte": cap}}
where={"budget_id": None, **_SPENT_ROWS_WHERE}
) # mutable-ok: prisma where filter must be a dict
writes.queue_spend_decrement(
where={"user_id": {"in": user_ids}, "spend": {"gt": cap}}, amount=cap
) # mutable-ok: prisma where filter must be a dict
plain: Final = [
uid for bid, uid in tiered if bid is None or bid not in cascade.rollover_caps
] # mutable-ok: prisma "in" filter takes a list
if plain:
writes.queue_spend_zero(where={"user_id": {"in": plain}}) # mutable-ok: prisma where filter must be a dict
return
writes.queue_spend_zero(
where={"budget_id": None, "spend": {"gt": 0, "lte": cap}}
) # mutable-ok: prisma where filter must be a dict
writes.queue_spend_decrement(
where={"budget_id": None, "spend": {"gt": cap}}, amount=cap
) # mutable-ok: prisma where filter must be a dict


@dataclass(frozen=True, slots=True)
Expand Down
6 changes: 3 additions & 3 deletions tests/litellm_utils_tests/test_proxy_budget_reset.py
Original file line number Diff line number Diff line change
Expand Up @@ -401,7 +401,7 @@ async def get_data_mock(table_name, *args, **kwargs):

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"]
Expand Down Expand Up @@ -602,7 +602,7 @@ async def fake_reset_team(team, current_time, reset_settings=None):
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.
Expand Down Expand Up @@ -1031,7 +1031,7 @@ async def fake_get_data(*, table_name, query_type, **kwargs):

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()
(
Expand Down
75 changes: 63 additions & 12 deletions tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,7 +403,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},
}
]
Expand Down Expand Up @@ -504,7 +504,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},
}
]
Expand All @@ -524,6 +524,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
("org", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
("tag", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
("model_access_group", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
("enduser", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}),
]


Expand Down Expand Up @@ -553,6 +554,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())
Expand Down Expand Up @@ -720,14 +763,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
Expand Down Expand Up @@ -3043,13 +3094,13 @@ def test_budget_cascade_carries_enduser_overage_when_rollover_enabled(
assert {
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"gt": 10.0}},
"where": {"budget_id": "budget-roll", "spend": {"gt": 10.0}},
"data": {"spend": {"decrement": 10.0}},
} in enduser_writes
assert {
"table": "enduser",
"op": "update_many",
"where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"lte": 10.0}},
"where": {"budget_id": "budget-roll", "spend": {"gt": 0, "lte": 10.0}},
"data": {"spend": 0},
} in enduser_writes

Expand Down
Loading