perf(proxy): split aggregated usage query into key-free rollups and bounded top-N keys

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-15 21:17:00 +00:00
parent 1debb438f5
commit 92e55b3b22
5 changed files with 479 additions and 131 deletions

View file

@ -2027,6 +2027,10 @@ 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
# 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

@ -9,7 +9,7 @@ from fastapi import HTTPException, status
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.constants import PTU_SENTINEL_API_KEY
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
@ -169,6 +169,32 @@ class _EntityRollupRow(_GroupingSetsRow):
api_key_rolled: int
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]
start_date: ReadOnly[str]
end_date: ReadOnly[str]
model: ReadOnly[str | None]
api_key: ReadOnly[str | list[str] | None]
exclude_entity_ids: ReadOnly[list[str] | None]
timezone_offset_minutes: ReadOnly[int | None]
include_current_utc_day: ReadOnly[bool]
_SqlQuery = tuple[str, list[str]]
async def _query_raw_optional(
prisma_client: PrismaClient, query: _SqlQuery | None
) -> list[dict[str, object]] | None: # mutable-ok: prisma query_raw return shape
if query is None:
return None
return await prisma_client.db.query_raw(query[0], *query[1])
def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float:
"""Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled.
@ -689,6 +715,27 @@ def _ptu_flat_cost_select(table_name: str) -> str:
return "0::float AS ptu_flat_cost"
def _rollup_metric_select(table_name: str) -> str:
return f"""
SUM(spend)::float AS spend,
{_ptu_flat_cost_select(table_name)},
SUM(prompt_tokens)::bigint AS prompt_tokens,
SUM(completion_tokens)::bigint AS completion_tokens,
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens,
SUM(compression_saved_tokens)::bigint AS compression_saved_tokens,
SUM(compression_savings_spend)::float AS compression_savings_spend,
SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend,
SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend,
SUM(autorouter_savings_spend)::float AS autorouter_savings_spend,
SUM(api_requests)::bigint AS api_requests,
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests"""
_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)"
def _build_aggregated_sql_query(
*,
table_name: str,
@ -702,12 +749,16 @@ 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 a parameterized SQL GROUP BY query for aggregated daily activity.
"""Build the key-free GROUPING SETS query for aggregated daily activity.
Groups by (date, api_key, model, model_group, custom_llm_provider,
mcp_namespaced_tool_name, endpoint) with SUMs on all metric columns.
The entity_id column is intentionally omitted from GROUP BY to collapse
rows across entities — this is where the biggest row reduction comes from.
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.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
@ -730,14 +781,6 @@ def _build_aggregated_sql_query(
exclude_entity_ids=exclude_entity_ids,
)
# Postgres computes every rollup level the response needs — per-date
# totals, per-(date, model), per-(date, model, api_key), per-provider,
# etc. — in a single pass via GROUPING SETS. The GROUPING() bitmask
# encodes which level a row belongs to so Python can dispatch rows
# straight into their buckets without re-summing. The leaf grouping
# is omitted on purpose: nothing in the response shape needs it once
# all the rollups are present.
#
# 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
@ -745,44 +788,25 @@ def _build_aggregated_sql_query(
sql_query: Final = f"""
SELECT
date,
api_key,
NULL::text AS api_key,
model,
COALESCE(NULLIF(model_group, ''), model) AS model_group,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model),
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
SUM(spend)::float AS spend,
{_ptu_flat_cost_select(table_name)},
SUM(prompt_tokens)::bigint AS prompt_tokens,
SUM(completion_tokens)::bigint AS completion_tokens,
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens,
SUM(compression_saved_tokens)::bigint AS compression_saved_tokens,
SUM(compression_savings_spend)::float AS compression_savings_spend,
SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend,
SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend,
SUM(autorouter_savings_spend)::float AS autorouter_savings_spend,
SUM(api_requests)::bigint AS api_requests,
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests
(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, api_key),
(date, model),
(date, model, api_key),
(date, COALESCE(NULLIF(model_group, ''), model)),
(date, COALESCE(NULLIF(model_group, ''), model), api_key),
(date, {_MODEL_GROUP_EXPR}),
(date, custom_llm_provider),
(date, custom_llm_provider, api_key),
(date, mcp_namespaced_tool_name),
(date, mcp_namespaced_tool_name, api_key),
(date, endpoint),
(date, endpoint, api_key),
()
)
"""
@ -790,6 +814,80 @@ def _build_aggregated_sql_query(
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,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
sentinel_param: Final = f"${len(where_params) + 1}"
sql_query: Final = f"""
WITH top_api_keys AS (
SELECT api_key
FROM "{pg_table}"
WHERE {where_clause} AND api_key <> {sentinel_param}
GROUP BY api_key
ORDER BY SUM(spend) DESC, api_key
LIMIT {USAGE_TOP_API_KEYS_LIMIT}
)
SELECT
date,
api_key,
model,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
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)}
FROM "{pg_table}"
WHERE {where_clause} AND api_key IN (SELECT api_key FROM top_api_keys)
GROUP BY GROUPING SETS (
(date, api_key),
(date, model, api_key),
(date, {_MODEL_GROUP_EXPR}, api_key),
(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]
def _build_entity_rollup_sql_query(
*,
table_name: str,
@ -832,21 +930,7 @@ def _build_entity_rollup_sql_query(
"{entity_id_field}" AS entity_id,
date,
api_key,
GROUPING(api_key) AS api_key_rolled,
SUM(spend)::float AS spend,
{_ptu_flat_cost_select(table_name)},
SUM(prompt_tokens)::bigint AS prompt_tokens,
SUM(completion_tokens)::bigint AS completion_tokens,
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens,
SUM(compression_saved_tokens)::bigint AS compression_saved_tokens,
SUM(compression_savings_spend)::float AS compression_savings_spend,
SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend,
SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend,
SUM(autorouter_savings_spend)::float AS autorouter_savings_spend,
SUM(api_requests)::bigint AS api_requests,
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests
GROUPING(api_key) AS api_key_rolled,{_rollup_metric_select(table_name)}
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY GROUPING SETS (
@ -948,6 +1032,7 @@ async def _aggregate_spend_records(
# current grouping set's key), 0 when the column is part of the key.
_GROUP_GRAND_TOTAL: Final = 127 # 0b1111111 — all rolled up
_GROUP_DATE: Final = 63 # 0b0111111 — only date kept
_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000 — api_key position in the 7-bit mask
_GROUP_DATE_API_KEY: Final = 31 # 0b0011111
_GROUP_DATE_MODEL: Final = 47 # 0b0101111
_GROUP_DATE_MODEL_API_KEY: Final = 15 # 0b0001111
@ -1311,9 +1396,11 @@ async def get_daily_activity_aggregated(
) -> SpendAnalyticsPaginatedResponse:
"""Aggregated variant that returns the full result set (no pagination).
Uses SQL GROUP BY to aggregate rows in the database rather than fetching
all individual rows into Python. This collapses rows across entities
(users/teams/orgs), reducing ~150k rows to ~2-3k grouped rows.
Runs two GROUPING SETS queries in parallel: a key-free one for totals and the
model/provider/mcp/endpoint rollups (row count independent of key cardinality)
and a bounded one for the per-key rollups of the top USAGE_TOP_API_KEYS_LIMIT
keys by spend. breakdown.api_keys and every api_key_breakdown therefore list at
most that many keys, while the totals and the key-free rollups cover every key.
include_entity_breakdown runs a small companion rollup query and folds
`breakdown.entities` onto the response, as entity-scoped views like Team Usage need.
@ -1333,7 +1420,7 @@ async def get_daily_activity_aggregated(
)
try:
sql_query, sql_params = _build_aggregated_sql_query(
query_kwargs: Final = _AggregatedQueryKwargs(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
@ -1345,36 +1432,17 @@ 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)
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
entity_query: Final = (
_build_entity_rollup_sql_query(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
start_date=start_date,
end_date=end_date,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
timezone_offset_minutes=timezone_offset_minutes,
include_current_utc_day=include_current_utc_day,
)
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),
_query_raw_optional(prisma_client, entity_query),
)
# Execute the GROUPING SETS query (one row per rollup level), alongside
# the per-entity companion rollup when the caller wants entities.
raw_rows, raw_entity_rows = (
await asyncio.gather(
prisma_client.db.query_raw(sql_query, *sql_params),
prisma_client.db.query_raw(entity_query[0], *entity_query[1]),
)
if entity_query is not None
else (await prisma_client.db.query_raw(sql_query, *sql_params), None)
)
records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or [])]
records: Final = [_GroupingSetsRow(**row) for row in (*(raw_key_free_rows or ()), *(raw_top_key_rows or ()))]
# The grouping-sets dispatcher places each row directly in its bucket
# using the row's GROUPING() bitmask. No Python-side summing needed.
@ -1426,6 +1494,7 @@ async def get_daily_activity_aggregated(
page=1,
total_pages=1,
has_more=False,
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
),
)

View file

@ -96,6 +96,11 @@ class DailySpendMetadata(BaseModel):
page: int = Field(default=1)
total_pages: int = Field(default=1)
has_more: bool = Field(default=False)
api_key_limit: int | None = Field(
default=None,
description="When set, api_keys and every api_key_breakdown list at most this many keys, "
"ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.",
)
class SpendAnalyticsPaginatedResponse(BaseModel):

View file

@ -1,17 +1,24 @@
import re
from collections.abc import Sequence
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import psycopg
import pytest
from psycopg.rows import dict_row
from pytest_postgresql import factories
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
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,
@ -159,7 +166,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"autorouter_savings_spend": 0.0,
"failed_requests": 0,
}
mock_rows = [
key_free_rows = [
# (date, endpoint) — rolls up across api_keys and models
{
**base,
@ -185,31 +192,6 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"api_requests": 1,
"successful_requests": 1,
},
# (date, endpoint, api_key) — populates the per-key sub-bucket
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "key-1",
"group_level": 30,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
"api_requests": 2,
"successful_requests": 2,
},
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/embeddings",
"api_key": "key-2",
"group_level": 30,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
},
# (date) — per-date totals
{
**base,
@ -237,8 +219,35 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"successful_requests": 3,
},
]
top_key_rows = [
# (date, endpoint, api_key) — populates the per-key sub-bucket
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "key-1",
"group_level": 30,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
"api_requests": 2,
"successful_requests": 2,
},
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/embeddings",
"api_key": "key-2",
"group_level": 30,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
},
]
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows])
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -284,8 +293,11 @@ 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
# Verify query_raw was called (not find_many)
mock_prisma.db.query_raw.assert_called_once()
# 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
@pytest.mark.asyncio
@ -812,7 +824,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"autorouter_savings_spend": 0.0,
"failed_requests": 0,
}
mock_rows = [
key_free_rows = [
{
**base,
"date": "2024-01-01",
@ -825,6 +837,8 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"api_requests": 1,
"successful_requests": 1,
},
]
top_key_rows = [
{
**base,
"date": "2024-01-01",
@ -839,7 +853,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
},
]
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows])
# Active table returns nothing for this key
mock_prisma.db.litellm_verificationtoken = MagicMock()
@ -1240,17 +1254,61 @@ class TestBuildAggregatedSqlQuery:
normalized = " ".join(sql.split())
fallback = "COALESCE(NULLIF(model_group, ''), model)"
assert f"{fallback} AS model_group" in normalized
assert (
f"GROUPING(date, api_key, model, {fallback}, "
"custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized
)
assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized
assert f"GROUPING(model, {fallback}, custom_llm_provider, mcp_namespaced_tool_name, endpoint)" in normalized
assert f"(date, {fallback})," in normalized
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
is projected as a NULL literal so the dispatcher's row shape is unchanged."""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
start_date="2026-07-01",
end_date="2026-07-01",
model=None,
api_key=None,
)
normalized = " ".join(sql.split())
grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1]
assert "api_key" not in grouping_block
assert "NULL::text AS api_key" in normalized
def test_top_api_keys_query_ranks_keys_deterministically_and_shares_filters(self):
sql, params = _build_top_api_keys_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id="user-1",
start_date="2026-05-29",
end_date="2026-06-02",
model="bedrock/global.anthropic.claude-opus-4-8",
api_key="sk-test",
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]
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,
]
class TestAggregatedEmptyEntityFilter:
_BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query)
_BUILDERS: Final = (_build_aggregated_sql_query, _build_top_api_keys_sql_query, _build_entity_rollup_sql_query)
@pytest.mark.parametrize("build", _BUILDERS)
def test_empty_entity_list_emits_no_degenerate_in_clause(self, build):
@ -1267,7 +1325,8 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert "IN ()" not in normalized
assert '"team_id" IN' not in normalized
assert params == ["2026-08-01", "2026-08-19"]
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else []
assert params == ["2026-08-01", "2026-08-19", *sentinel_params]
@pytest.mark.parametrize("build", _BUILDERS)
def test_empty_entity_list_matches_nothing_rather_than_everything(self, build):
@ -1298,7 +1357,8 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert '"team_id" IN ($3, $4)' in normalized
assert "FALSE" not in normalized
assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta"]
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else []
assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta", *sentinel_params]
@pytest.mark.asyncio
@ -1313,7 +1373,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_rows = [
key_free_rows = [
{
"date": None,
"api_key": None,
@ -1338,7 +1398,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
"failed_requests": None,
}
]
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, []])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
@ -1365,6 +1425,211 @@ async def test_get_daily_activity_aggregated_empty_result_set():
assert result.metadata.total_compression_saved_tokens == 0
_aggregated_postgresql_proc: Final = factories.postgresql_proc()
_aggregated_postgresql: Final = factories.postgresql("_aggregated_postgresql_proc")
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
user_id TEXT,
date TEXT NOT NULL,
api_key TEXT NOT NULL,
model TEXT,
model_group TEXT,
custom_llm_provider TEXT,
mcp_namespaced_tool_name TEXT,
endpoint TEXT,
prompt_tokens BIGINT DEFAULT 0,
completion_tokens BIGINT DEFAULT 0,
cache_read_input_tokens BIGINT DEFAULT 0,
cache_creation_input_tokens BIGINT DEFAULT 0,
compression_saved_tokens BIGINT DEFAULT 0,
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
spend DOUBLE PRECISION DEFAULT 0,
api_requests BIGINT DEFAULT 0,
successful_requests BIGINT DEFAULT 0,
failed_requests BIGINT DEFAULT 0
)
"""
def _seed_daily_user_spend(conn: psycopg.Connection, rows: Sequence[tuple[object, ...]]) -> None:
with conn.cursor() as cur:
cur.execute(_DAILY_USER_SPEND_DDL)
cur.executemany(
"""
INSERT INTO "LiteLLM_DailyUserSpend"
(id, user_id, date, api_key, model, model_group, custom_llm_provider,
endpoint, prompt_tokens, spend, api_requests, successful_requests)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
rows,
)
conn.commit()
def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]):
"""Run the proxy's $N-parameterized SQL through psycopg, recording each result size."""
async def query_raw(sql: str, *params: str) -> list[dict[str, object]]:
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
with conn.cursor(row_factory=dict_row) as cur:
cur.execute(
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
{f"p{i}": v for i, v in enumerate(params, start=1)},
)
rows: Final = cur.fetchall()
row_counts.append(len(rows))
return rows
return query_raw
@pytest.mark.asyncio
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.
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.
"""
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 5
key_rows: Final = [
(
f"row-{i:03d}",
f"user-{i:03d}",
"2026-06-01",
f"key-{i:03d}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
6.0 if i == 4 else float(i + 1),
1,
1,
)
for i in range(n_keys)
]
sentinel_row: Final = (
"row-ptu",
None,
"2026-06-01",
PTU_SENTINEL_API_KEY,
"gpt-5",
"",
"azure",
None,
0,
1000.0,
0,
0,
)
_seed_daily_user_spend(_aggregated_postgresql, [*key_rows, sentinel_row])
key_spend: Final = sum(6.0 if i == 4 else float(i + 1) for i in range(n_keys))
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
# Key-free query: (), (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]
assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0)
assert result.metadata.total_api_requests == n_keys
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
expected_top: Final = {f"key-{i:03d}" for i in range(6, n_keys)} | {"key-004"}
day: Final = result.results[0]
assert day.metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.api_keys) == expected_top
assert day.breakdown.api_keys["key-004"].metrics.spend == 6.0
assert "key-005" not in day.breakdown.api_keys
assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys
assert day.breakdown.models["gpt-5"].metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == expected_top
assert day.breakdown.providers["openai"].metrics.spend == pytest.approx(key_spend)
assert set(day.breakdown.providers["openai"].api_key_breakdown) == expected_top
assert day.breakdown.endpoints["/v1/chat/completions"].metrics.api_requests == n_keys
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_queries(
_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."""
rows: Final = [
(
f"row-{i}",
f"user-{i}",
"2026-06-01",
f"key-{i}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
float(i + 1),
1,
1,
)
for i in range(3)
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key="key-1",
)
assert result.metadata.total_spend == 2.0
day: Final = result.results[0]
assert set(day.breakdown.api_keys) == {"key-1"}
assert day.breakdown.api_keys["key-1"].metrics.spend == 2.0
assert day.breakdown.models["gpt-5"].metrics.spend == 2.0
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"}
def _no_spend_record():
"""A rollup row for a key with no spend, where SUM() returns NULL (None)."""
return SimpleNamespace(
@ -2128,8 +2393,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()
@ -2156,7 +2421,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
},
]
mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, entity_rows])
mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, top_key_rows, entity_rows])
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -2173,9 +2438,9 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
include_entity_breakdown=True,
)
assert mock_prisma.db.query_raw.call_count == 2
assert mock_prisma.db.query_raw.call_count == 3
main_sql = mock_prisma.db.query_raw.call_args_list[0][0][0]
entity_sql = mock_prisma.db.query_raw.call_args_list[1][0][0]
entity_sql = mock_prisma.db.query_raw.call_args_list[2][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

View file

@ -27430,6 +27430,11 @@ export interface components {
};
/** DailySpendMetadata */
DailySpendMetadata: {
/**
* Api Key Limit
* @description When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.
*/
api_key_limit?: number | null;
/**
* Has More
* @default false