perf: use SQL GROUP BY for aggregated daily activity endpoints

Replace find_many + Python-side aggregation with a single SQL GROUP BY
query via query_raw in get_daily_activity_aggregated. This collapses
rows across entities (users/teams/orgs) in the database, reducing ~150k
rows to ~2-3k grouped rows before transfer to Python.

Also adds composite indexes (entity_id, date) to all 6 daily spend
tables for faster filtered queries.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-02-19 14:36:28 -08:00
parent b209b11522
commit e2e698944a
5 changed files with 212 additions and 82 deletions

View file

@ -660,7 +660,7 @@ model LiteLLM_DailyUserSpend {
@@unique([user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([user_id])
@@index([user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -691,7 +691,7 @@ model LiteLLM_DailyOrganizationSpend {
@@unique([organization_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([organization_id])
@@index([organization_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -721,7 +721,7 @@ model LiteLLM_DailyEndUserSpend {
updated_at DateTime @updatedAt
@@unique([end_user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([end_user_id])
@@index([end_user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -751,7 +751,7 @@ model LiteLLM_DailyAgentSpend {
updated_at DateTime @updatedAt
@@unique([agent_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([agent_id])
@@index([agent_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -782,7 +782,7 @@ model LiteLLM_DailyTeamSpend {
@@unique([team_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([team_id])
@@index([team_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -814,7 +814,7 @@ model LiteLLM_DailyTagSpend {
@@unique([tag, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([tag])
@@index([tag, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])

View file

@ -1,4 +1,5 @@
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union
from fastapi import HTTPException, status
@ -17,6 +18,16 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendMetrics,
)
# Mapping from Prisma accessor names to actual PostgreSQL table names.
_PRISMA_TO_PG_TABLE: Dict[str, str] = {
"litellm_dailyuserspend": "LiteLLM_DailyUserSpend",
"litellm_dailyteamspend": "LiteLLM_DailyTeamSpend",
"litellm_dailyorganizationspend": "LiteLLM_DailyOrganizationSpend",
"litellm_dailyenduserspend": "LiteLLM_DailyEndUserSpend",
"litellm_dailyagentspend": "LiteLLM_DailyAgentSpend",
"litellm_dailytagspend": "LiteLLM_DailyTagSpend",
}
def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics:
"""Update metrics with new record data."""
@ -455,6 +466,111 @@ def _build_where_conditions(
return where_conditions
def _build_aggregated_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: Optional[Union[str, List[str]]],
start_date: str,
end_date: str,
model: Optional[str],
api_key: Optional[str],
exclude_entity_ids: Optional[List[str]] = None,
timezone_offset_minutes: Optional[int] = None,
) -> Tuple[str, List[Any]]:
"""Build a parameterized SQL GROUP BY 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.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
"""
pg_table = _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
)
sql_conditions: List[str] = []
sql_params: List[Any] = []
p = 1 # parameter index (1-based for PostgreSQL $N placeholders)
# Date range (always present)
sql_conditions.append(f"date >= ${p}")
sql_params.append(adjusted_start)
p += 1
sql_conditions.append(f"date <= ${p}")
sql_params.append(adjusted_end)
p += 1
# Optional entity filter
if entity_id is not None:
if isinstance(entity_id, list):
placeholders = ", ".join(f"${p + i}" for i in range(len(entity_id)))
sql_conditions.append(f'"{entity_id_field}" IN ({placeholders})')
sql_params.extend(entity_id)
p += len(entity_id)
else:
sql_conditions.append(f'"{entity_id_field}" = ${p}')
sql_params.append(entity_id)
p += 1
# Exclude specific entities
if exclude_entity_ids:
placeholders = ", ".join(
f"${p + i}" for i in range(len(exclude_entity_ids))
)
sql_conditions.append(f'"{entity_id_field}" NOT IN ({placeholders})')
sql_params.extend(exclude_entity_ids)
p += len(exclude_entity_ids)
# Optional model filter
if model:
sql_conditions.append(f"model = ${p}")
sql_params.append(model)
p += 1
# Optional api_key filter
if api_key:
sql_conditions.append(f"api_key = ${p}")
sql_params.append(api_key)
p += 1
where_clause = " AND ".join(sql_conditions)
sql_query = f"""
SELECT
date,
api_key,
model,
model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
SUM(spend)::float AS spend,
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(api_requests)::bigint AS api_requests,
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY date, api_key, model, model_group, custom_llm_provider,
mcp_namespaced_tool_name, endpoint
ORDER BY date DESC
"""
return sql_query, sql_params
async def _aggregate_spend_records(
*,
prisma_client: PrismaClient,
@ -625,6 +741,10 @@ 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.
Matches the response model of the paginated endpoint so the UI does not need to transform.
"""
if prisma_client is None:
@ -640,7 +760,8 @@ async def get_daily_activity_aggregated(
)
try:
where_conditions = _build_where_conditions(
sql_query, sql_params = _build_aggregated_sql_query(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
start_date=start_date,
@ -651,19 +772,21 @@ async def get_daily_activity_aggregated(
timezone_offset_minutes=timezone_offset_minutes,
)
# Fetch all matching results (no pagination)
daily_spend_data = await getattr(prisma_client.db, table_name).find_many(
where=where_conditions,
order=[
{"date": "desc"},
],
)
# Execute GROUP BY query — returns pre-aggregated dicts
rows = await prisma_client.db.query_raw(sql_query, *sql_params)
if rows is None:
rows = []
# Convert dicts to objects for compatibility with _aggregate_spend_records
records = [SimpleNamespace(**row) for row in rows]
# entity_id_field=None skips entity breakdown (entity dimension was
# collapsed by the GROUP BY, so per-entity data is not available)
aggregated = await _aggregate_spend_records(
prisma_client=prisma_client,
records=daily_spend_data,
entity_id_field=entity_id_field,
entity_metadata_field=entity_metadata_field,
records=records,
entity_id_field=None,
entity_metadata_field=None,
)
return SpendAnalyticsPaginatedResponse(

View file

@ -613,7 +613,7 @@ model LiteLLM_DailyUserSpend {
@@unique([user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([user_id])
@@index([user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -644,7 +644,7 @@ model LiteLLM_DailyOrganizationSpend {
@@unique([organization_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([organization_id])
@@index([organization_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -674,7 +674,7 @@ model LiteLLM_DailyEndUserSpend {
updated_at DateTime @updatedAt
@@unique([end_user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([end_user_id])
@@index([end_user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -704,7 +704,7 @@ model LiteLLM_DailyAgentSpend {
updated_at DateTime @updatedAt
@@unique([agent_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([agent_id])
@@index([agent_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -735,7 +735,7 @@ model LiteLLM_DailyTeamSpend {
@@unique([team_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([team_id])
@@index([team_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -767,7 +767,7 @@ model LiteLLM_DailyTagSpend {
@@unique([tag, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([tag])
@@index([tag, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])

View file

@ -613,7 +613,7 @@ model LiteLLM_DailyUserSpend {
@@unique([user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([user_id])
@@index([user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -644,7 +644,7 @@ model LiteLLM_DailyOrganizationSpend {
@@unique([organization_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([organization_id])
@@index([organization_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -674,7 +674,7 @@ model LiteLLM_DailyEndUserSpend {
updated_at DateTime @updatedAt
@@unique([end_user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([end_user_id])
@@index([end_user_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -704,7 +704,7 @@ model LiteLLM_DailyAgentSpend {
updated_at DateTime @updatedAt
@@unique([agent_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([agent_id])
@@index([agent_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -735,7 +735,7 @@ model LiteLLM_DailyTeamSpend {
@@unique([team_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([team_id])
@@index([team_id, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])
@ -767,7 +767,7 @@ model LiteLLM_DailyTagSpend {
@@unique([tag, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
@@index([tag])
@@index([tag, date])
@@index([api_key])
@@index([model])
@@index([mcp_namespaced_tool_name])

View file

@ -135,36 +135,45 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
# Create mock records with endpoint fields
class MockRecord:
def __init__(self, date, endpoint, api_key, model, spend, prompt_tokens, completion_tokens):
self.date = date
self.endpoint = endpoint
self.api_key = api_key
self.model = model
self.model_group = None
self.custom_llm_provider = "openai"
self.mcp_namespaced_tool_name = None
self.spend = spend
self.prompt_tokens = prompt_tokens
self.completion_tokens = completion_tokens
self.total_tokens = prompt_tokens + completion_tokens
self.cache_read_input_tokens = 0
self.cache_creation_input_tokens = 0
self.api_requests = 1
self.successful_requests = 1
self.failed_requests = 0
mock_records = [
MockRecord("2024-01-01", "/v1/chat/completions", "key-1", "gpt-4", 10.0, 100, 50),
MockRecord("2024-01-01", "/v1/chat/completions", "key-1", "gpt-4", 5.0, 50, 25),
MockRecord("2024-01-01", "/v1/embeddings", "key-2", "text-embedding-ada-002", 3.0, 30, 0),
# query_raw returns list of dicts (pre-aggregated by GROUP BY)
mock_rows = [
{
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "key-1",
"model": "gpt-4",
"model_group": None,
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": None,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"api_requests": 2,
"successful_requests": 2,
"failed_requests": 0,
},
{
"date": "2024-01-01",
"endpoint": "/v1/embeddings",
"api_key": "key-2",
"model": "text-embedding-ada-002",
"model_group": None,
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": None,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
},
]
# Mock the table methods
mock_table = MagicMock()
mock_table.find_many = AsyncMock(return_value=mock_records)
mock_prisma.db.litellm_dailyuserspend = mock_table
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=[])
@ -210,6 +219,9 @@ 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()
@pytest.mark.asyncio
async def test_get_api_key_metadata_returns_active_key_metadata():
@ -399,33 +411,28 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
class MockRecord:
def __init__(self, date, endpoint, api_key, model, spend, prompt_tokens, completion_tokens):
self.date = date
self.endpoint = endpoint
self.api_key = api_key
self.model = model
self.model_group = None
self.custom_llm_provider = "openai"
self.mcp_namespaced_tool_name = None
self.spend = spend
self.prompt_tokens = prompt_tokens
self.completion_tokens = completion_tokens
self.total_tokens = prompt_tokens + completion_tokens
self.cache_read_input_tokens = 0
self.cache_creation_input_tokens = 0
self.api_requests = 1
self.successful_requests = 1
self.failed_requests = 0
# Records reference a deleted key
mock_records = [
MockRecord("2024-01-01", "/v1/chat/completions", "deleted-key-hash", "gpt-4", 10.0, 100, 50),
# query_raw returns list of dicts (pre-aggregated by GROUP BY)
mock_rows = [
{
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "deleted-key-hash",
"model": "gpt-4",
"model_group": None,
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": None,
"spend": 10.0,
"prompt_tokens": 100,
"completion_tokens": 50,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
},
]
mock_table = MagicMock()
mock_table.find_many = AsyncMock(return_value=mock_records)
mock_prisma.db.litellm_dailyuserspend = mock_table
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
# Active table returns nothing for this key
mock_prisma.db.litellm_verificationtoken = MagicMock()