diff --git a/litellm/proxy/management_endpoints/management_v1/list_framework.py b/litellm/proxy/management_endpoints/management_v1/list_framework.py index e2bddefab83..8f25c45016b 100644 --- a/litellm/proxy/management_endpoints/management_v1/list_framework.py +++ b/litellm/proxy/management_endpoints/management_v1/list_framework.py @@ -15,6 +15,7 @@ raw-SQL executor with every caller-supplied value bound to a placeholder. from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timezone +from functools import partial, reduce from math import ceil from typing import Generic, Literal, Protocol, TypeVar @@ -242,12 +243,26 @@ def _render(predicate: Predicate, index: int) -> tuple[str, tuple[object, ...]]: assert_never(predicate) +def _render_one( + rendered: tuple[tuple[str, ...], tuple[object, ...]], + predicate: Predicate, + first_index: int, +) -> tuple[tuple[str, ...], tuple[object, ...]]: + """Append one predicate, numbering it after the binds already consumed.""" + clauses, params = rendered + clause, clause_params = _render(predicate, first_index + len(params)) + return (*clauses, clause), (*params, *clause_params) + + def _render_all(predicates: tuple[Predicate, ...], index: int) -> tuple[tuple[str, ...], tuple[object, ...]]: - if not predicates: - return (), () - head, head_params = _render(predicates[0], index) - tail, tail_params = _render_all(predicates[1:], index + len(head_params)) - return (head, *tail), head_params + tail_params + """Render every predicate, numbering placeholders continuously across them. + + Folded rather than self-recursive: walking a predicate list is a running index, and + recursing per predicate grew the stack with the filter count for nothing. `_render` + still re-enters here for `AnyOf`, whose clauses are plain `Compare`s from `?q=`, so + that nesting is one level deep and cannot be driven deeper by a caller. + """ + return reduce(partial(_render_one, first_index=index), predicates, ((), ())) def where_sql(where: tuple[Predicate, ...], first_index: int = 1) -> tuple[str, tuple[object, ...]]: diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py index f98286985b7..40473f1a25a 100644 --- a/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py @@ -382,6 +382,20 @@ def test_created_at_range_is_bound_as_a_timestamp(query_raw, as_proxy_admin): assert params[0] == datetime(2026, 7, 1, tzinfo=timezone.utc) +def test_numbers_placeholders_continuously_across_predicates(query_raw, as_proxy_admin): + """Each predicate is numbered after the binds the ones before it consumed. Restart + the count and `$1` gets read as the duration while the search string goes unbound.""" + _serve(query_raw, []) + + _get("filter[budget_duration][in]=7d,30d&filter[max_budget][gte]=5&q=prod") + + sql, *params = _select_call(query_raw) + assert '"budget_duration" IN ($1, $2)' in sql + assert '"max_budget" >= $3' in sql + assert '"budget_id" ILIKE $4' in sql + assert params[:4] == ["7d", "30d", 5.0, "%prod%"] + + def test_a_non_datetime_bind_is_not_cast(query_raw, as_proxy_admin): """Guards the cast above from being applied to every placeholder.""" _serve(query_raw, [])