perf(proxy): serve key-free rollups and top-N keys from one UNION ALL statement

Both arms now run in a single query_raw call so totals and per-key
breakdowns come from the same snapshot. USAGE_TOP_API_KEYS_LIMIT can be
raised via env for deployments that need every key in the response.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-15 21:52:37 +00:00
parent a0311dddf7
commit f94a40f841
3 changed files with 87 additions and 150 deletions

View file

@ -2027,10 +2027,7 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: "
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
# Per-api_key rollups on the aggregated usage endpoint cover only the top N keys
# by spend so the result set stops growing with key count. Totals and the
# model/provider/endpoint rollups still cover every key.
USAGE_TOP_API_KEYS_LIMIT: Final[int] = 100
USAGE_TOP_API_KEYS_LIMIT: Final[int] = int(os.getenv("USAGE_TOP_API_KEYS_LIMIT", "100"))
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
# declares no ptu_effective_from, bounding the scan for an open-ended window.
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90

View file

@ -170,8 +170,6 @@ class _EntityRollupRow(_GroupingSetsRow):
class _AggregatedQueryKwargs(TypedDict):
"""Filter arguments shared by the three aggregated SQL builders."""
table_name: ReadOnly[str]
entity_id_field: ReadOnly[str]
entity_id: ReadOnly[str | list[str] | None]
@ -749,16 +747,14 @@ def _build_aggregated_sql_query(
timezone_offset_minutes: int | None = None,
include_current_utc_day: bool = False,
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
"""Build the key-free GROUPING SETS query for aggregated daily activity.
"""Build the GROUPING SETS query for aggregated daily activity.
Emits the grand total, per-date totals and the per-(date, model), model_group,
provider, mcp tool and endpoint rollups. api_key is never a grouping column here,
so the row count is bounded by dates x distinct models/providers/endpoints and
does not grow with the number of keys. Per-key rollups come from
_build_top_api_keys_sql_query. Both queries emit the same 7-bit group_level
bitmask (date, api_key, model, model_group, provider, mcp, endpoint); this one
hard-codes the api_key bit to "rolled up" so the dispatcher can consume the two
result sets as one stream.
One statement, two UNION ALL arms over the same WHERE clause. The first arm is
key-free: grand total, per-date totals and the (date, model / model_group /
provider / mcp / endpoint) rollups, so its row count never grows with the number
of keys. The second arm emits the (date, <dimension>, api_key) rollups for the
USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit
group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint).
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
@ -771,77 +767,6 @@ def _build_aggregated_sql_query(
start_date, end_date, timezone_offset_minutes, include_current_utc_day
)
where_clause, sql_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
# TODO: drop the successful_requests/failed_requests aggregates (and the
# total_successful_requests metadata they feed) once the admin UI reads SGR
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
# api_requests rollups are still served from here.
sql_query: Final = f"""
SELECT
date,
NULL::text AS api_key,
model,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
(GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT}
| GROUPING(model, {_MODEL_GROUP_EXPR},
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,{_rollup_metric_select(table_name)}
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY GROUPING SETS (
(date),
(date, model),
(date, {_MODEL_GROUP_EXPR}),
(date, custom_llm_provider),
(date, mcp_namespaced_tool_name),
(date, endpoint),
()
)
"""
return sql_query, sql_params
def _build_top_api_keys_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
start_date: str,
end_date: str,
model: str | None,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None = None,
include_current_utc_day: bool = False,
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
"""Per-key companion to _build_aggregated_sql_query.
Ranks keys by spend over the same WHERE clause, keeps the top
USAGE_TOP_API_KEYS_LIMIT (ties broken by api_key so the set is stable across
refreshes) and emits the six (date, <dimension>, api_key) rollups for those keys
only. The PTU flat-cost sentinel never ranks, so it cannot occupy a visible slot.
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
raise ValueError(f"Unknown table name: {table_name}")
adjusted_start, adjusted_end = _adjust_dates_for_timezone(
start_date, end_date, timezone_offset_minutes, include_current_utc_day
)
where_clause, where_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
@ -852,9 +777,38 @@ def _build_top_api_keys_sql_query(
exclude_entity_ids=exclude_entity_ids,
)
sentinel_param: Final = f"${len(where_params) + 1}"
metric_select: Final = _rollup_metric_select(table_name)
# TODO: drop the successful_requests/failed_requests aggregates (and the
# total_successful_requests metadata they feed) once the admin UI reads SGR
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
# api_requests rollups are still served from here.
sql_query: Final = f"""
WITH top_api_keys AS (
(SELECT
date,
NULL::text AS api_key,
model,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
(GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT}
| GROUPING(model, {_MODEL_GROUP_EXPR},
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,{metric_select}
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY GROUPING SETS (
(date),
(date, model),
(date, {_MODEL_GROUP_EXPR}),
(date, custom_llm_provider),
(date, mcp_namespaced_tool_name),
(date, endpoint),
()
))
UNION ALL
(WITH top_api_keys AS (
SELECT api_key
FROM "{pg_table}"
WHERE {where_clause} AND api_key <> {sentinel_param}
@ -872,7 +826,7 @@ def _build_top_api_keys_sql_query(
endpoint,
GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR},
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,{_rollup_metric_select(table_name)}
endpoint) AS group_level,{metric_select}
FROM "{pg_table}"
WHERE {where_clause} AND api_key IN (SELECT api_key FROM top_api_keys)
GROUP BY GROUPING SETS (
@ -882,7 +836,7 @@ def _build_top_api_keys_sql_query(
(date, custom_llm_provider, api_key),
(date, mcp_namespaced_tool_name, api_key),
(date, endpoint, api_key)
)
))
"""
return sql_query, [*where_params, PTU_SENTINEL_API_KEY]
@ -1432,17 +1386,15 @@ async def get_daily_activity_aggregated(
timezone_offset_minutes=timezone_offset_minutes,
include_current_utc_day=include_current_utc_day,
)
key_free_sql, key_free_params = _build_aggregated_sql_query(**query_kwargs)
top_keys_sql, top_keys_params = _build_top_api_keys_sql_query(**query_kwargs)
sql_query, sql_params = _build_aggregated_sql_query(**query_kwargs)
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
raw_key_free_rows, raw_top_key_rows, raw_entity_rows = await asyncio.gather(
prisma_client.db.query_raw(key_free_sql, *key_free_params),
prisma_client.db.query_raw(top_keys_sql, *top_keys_params),
raw_rows, raw_entity_rows = await asyncio.gather(
prisma_client.db.query_raw(sql_query, *sql_params),
_query_raw_optional(prisma_client, entity_query),
)
records: Final = [_GroupingSetsRow(**row) for row in (*(raw_key_free_rows or ()), *(raw_top_key_rows or ()))]
records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or ())]
# The grouping-sets dispatcher places each row directly in its bucket
# using the row's GROUPING() bitmask. No Python-side summing needed.

View file

@ -18,7 +18,6 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
_adjust_dates_for_timezone,
_build_aggregated_sql_query,
_build_entity_rollup_sql_query,
_build_top_api_keys_sql_query,
_is_user_agent_tag,
_record_to_spend_metrics,
get_api_key_metadata,
@ -166,7 +165,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"autorouter_savings_spend": 0.0,
"failed_requests": 0,
}
key_free_rows = [
mock_rows = [
# (date, endpoint) — rolls up across api_keys and models
{
**base,
@ -218,8 +217,6 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"api_requests": 3,
"successful_requests": 3,
},
]
top_key_rows = [
# (date, endpoint, api_key) — populates the per-key sub-bucket
{
**base,
@ -247,7 +244,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
},
]
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows])
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -293,11 +290,8 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
assert "key-2" in embeddings_endpoint.api_key_breakdown
assert embeddings_endpoint.api_key_breakdown["key-2"].metrics.spend == 3.0
# One key-free rollup query plus one bounded per-key query, no find_many
assert mock_prisma.db.query_raw.call_count == 2
key_free_sql, top_keys_sql = (call.args[0] for call in mock_prisma.db.query_raw.call_args_list)
assert "top_api_keys" not in key_free_sql
assert "WITH top_api_keys AS" in top_keys_sql
# Verify query_raw was called (not find_many)
mock_prisma.db.query_raw.assert_called_once()
@pytest.mark.asyncio
@ -484,9 +478,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash(
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com")]
)
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"}
]
return_value=[{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"}]
)
result = await get_api_key_metadata(
@ -824,7 +816,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"autorouter_savings_spend": 0.0,
"failed_requests": 0,
}
key_free_rows = [
mock_rows = [
{
**base,
"date": "2024-01-01",
@ -837,8 +829,6 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"api_requests": 1,
"successful_requests": 1,
},
]
top_key_rows = [
{
**base,
"date": "2024-01-01",
@ -853,7 +843,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
},
]
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows])
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
# Active table returns nothing for this key
mock_prisma.db.litellm_verificationtoken = MagicMock()
@ -1226,6 +1216,7 @@ class TestBuildAggregatedSqlQuery:
"user-1",
"bedrock/global.anthropic.claude-opus-4-8",
"sk-test",
PTU_SENTINEL_API_KEY,
]
assert "model = $4" in sql
assert "api_key = $5" in sql
@ -1259,9 +1250,9 @@ class TestBuildAggregatedSqlQuery:
assert "(date, model_group)" not in normalized
assert "COALESCE(model_group, model)" not in normalized
def test_key_free_query_never_groups_by_api_key(self):
"""The main rollup query must not emit one row per key, that is what blew up
the query engine at 3k+ keys. Every grouping set stays key-free and api_key
def test_totals_arm_never_groups_by_api_key(self):
"""The totals arm must not emit one row per key, that is what blew up the
query engine at 3k+ keys. Every grouping set there stays key-free and api_key
is projected as a NULL literal so the dispatcher's row shape is unchanged."""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
@ -1273,13 +1264,15 @@ class TestBuildAggregatedSqlQuery:
api_key=None,
)
normalized = " ".join(sql.split())
grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1]
totals_arm, _ = " ".join(sql.split()).split("UNION ALL")
grouping_block = totals_arm.split("GROUP BY GROUPING SETS (", 1)[1]
assert "api_key" not in grouping_block
assert "NULL::text AS api_key" in normalized
assert "NULL::text AS api_key" in totals_arm
def test_top_api_keys_query_ranks_keys_deterministically_and_shares_filters(self):
sql, params = _build_top_api_keys_sql_query(
def test_per_key_arm_ranks_keys_deterministically_and_shares_filters(self):
"""Both arms sit in one statement so totals and per-key rows come from the
same snapshot, and the per-key arm reuses the caller's filter params."""
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id="user-1",
@ -1290,25 +1283,20 @@ class TestBuildAggregatedSqlQuery:
timezone_offset_minutes=-330,
)
normalized = " ".join(sql.split())
assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in normalized
assert "api_key IN (SELECT api_key FROM top_api_keys)" in normalized
assert "api_key <> $6" in normalized
grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1]
totals_arm, per_key_arm = " ".join(sql.split()).split("UNION ALL")
assert "top_api_keys" not in totals_arm
assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in per_key_arm
assert "api_key IN (SELECT api_key FROM top_api_keys)" in per_key_arm
assert "api_key <> $6" in per_key_arm
assert per_key_arm.count("model = $4 AND api_key = $5") == 2
grouping_block = per_key_arm.split("GROUP BY GROUPING SETS (", 1)[1]
assert grouping_block.count(", api_key)") == 6
assert grouping_block.count("(date") == 6
assert params == [
"2026-05-29",
"2026-06-02",
"user-1",
"bedrock/global.anthropic.claude-opus-4-8",
"sk-test",
PTU_SENTINEL_API_KEY,
]
assert params[-1] == PTU_SENTINEL_API_KEY
class TestAggregatedEmptyEntityFilter:
_BUILDERS: Final = (_build_aggregated_sql_query, _build_top_api_keys_sql_query, _build_entity_rollup_sql_query)
_BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query)
@pytest.mark.parametrize("build", _BUILDERS)
def test_empty_entity_list_emits_no_degenerate_in_clause(self, build):
@ -1325,7 +1313,7 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert "IN ()" not in normalized
assert '"team_id" IN' not in normalized
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else []
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else []
assert params == ["2026-08-01", "2026-08-19", *sentinel_params]
@pytest.mark.parametrize("build", _BUILDERS)
@ -1357,7 +1345,7 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert '"team_id" IN ($3, $4)' in normalized
assert "FALSE" not in normalized
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else []
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else []
assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta", *sentinel_params]
@ -1373,7 +1361,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
key_free_rows = [
mock_rows = [
{
"date": None,
"api_key": None,
@ -1398,7 +1386,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
"failed_requests": None,
}
]
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, []])
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
@ -1492,13 +1480,13 @@ def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]):
async def test_get_daily_activity_aggregated_bounds_api_key_rollups(
_aggregated_postgresql: psycopg.Connection,
):
"""Run both GROUPING SETS queries against real Postgres with more keys than the cap.
"""Run the GROUPING SETS statement against real Postgres with more keys than the cap.
key-004 and key-005 tie on spend exactly at the USAGE_TOP_API_KEYS_LIMIT
cutoff; the api_key tiebreaker must keep key-004 and drop key-005. The PTU
sentinel outspends every key but must not take a slot. Excluded keys and the
sentinel still count toward the totals and the model rollup, which come from
the key-free query.
the key-free arm.
"""
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 5
key_rows: Final = [
@ -1554,10 +1542,10 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups(
api_key=None,
)
# Key-free query: (), (date), (date, model), (date, model_group), two providers,
# Key-free arm: (), (date), (date, model), (date, model_group), two providers,
# one mcp NULL bucket, endpoint plus its NULL bucket = 9 rows regardless of key count.
# Top-key query: six per-key grouping sets, each capped at the limit.
assert row_counts == [9, 6 * USAGE_TOP_API_KEYS_LIMIT]
# Per-key arm: six per-key grouping sets, each capped at the limit.
assert row_counts == [9 + 6 * USAGE_TOP_API_KEYS_LIMIT]
assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0)
assert result.metadata.total_api_requests == n_keys
@ -1579,11 +1567,11 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups(
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_queries(
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_arms(
_aggregated_postgresql: psycopg.Connection,
):
"""An explicit api_key filter must scope the key-free totals and the per-key
rollups to that key alone, so the two result sets never disagree."""
rollups to that key alone, so the two arms never disagree."""
rows: Final = [
(
f"row-{i}",
@ -2358,7 +2346,7 @@ def test_entity_rollup_sql_query_and_api_key_list_filter():
api_key=[],
)
assert "FALSE" in empty_sql
assert empty_params == ["2024-01-01", "2024-01-31"]
assert empty_params == ["2024-01-01", "2024-01-31", PTU_SENTINEL_API_KEY]
@pytest.mark.asyncio
@ -2393,8 +2381,8 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
{**base, "date": None, "group_level": 127, "spend": 18.0},
{**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0},
{**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0},
{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0},
]
top_key_rows = [{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}]
entity_base = {
key: value
for key, value in base.items()
@ -2421,7 +2409,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
},
]
mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, top_key_rows, entity_rows])
mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, entity_rows])
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -2438,9 +2426,9 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
include_entity_breakdown=True,
)
assert mock_prisma.db.query_raw.call_count == 3
assert mock_prisma.db.query_raw.call_count == 2
main_sql = mock_prisma.db.query_raw.call_args_list[0][0][0]
entity_sql = mock_prisma.db.query_raw.call_args_list[2][0][0]
entity_sql = mock_prisma.db.query_raw.call_args_list[1][0][0]
assert "entity_id" not in main_sql
assert '"team_id" AS entity_id' in entity_sql
assert '(date, "team_id"),' in entity_sql