Use conditional GROUP BY: skip for include_entity_id=True, keep for False

When include_entity_id=True (team endpoint), the GROUP BY columns match
the unique constraint, making aggregation a no-op. Skip it to avoid
wasted hash/sort CPU. Python _aggregate_spend_records handles rollup.

When include_entity_id=False (user/org endpoints), keep GROUP BY to
collapse rows across entities, reducing the result set. Also use
MAX(model_group) instead of including it in GROUP BY since it's
functionally dependent on model.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-06 23:07:44 -08:00
parent 38be8eaafa
commit 72a0c6c445
2 changed files with 128 additions and 72 deletions

View file

@ -481,16 +481,16 @@ def _build_aggregated_sql_query(
) -> Tuple[str, List[Any]]:
"""Build a parameterized SQL query for aggregated daily activity.
Uses a plain SELECT (no GROUP BY) because the table's unique constraint
already ensures near-uniqueness across (entity_id, date, api_key, model,
custom_llm_provider, mcp_namespaced_tool_name, endpoint). Skipping the
GROUP BY avoids the expensive hash/sort aggregation step in PostgreSQL,
which is essentially a no-op on these tables.
When include_entity_id is True (e.g. team endpoint), the entity_id column
is included in SELECT. In this case GROUP BY is skipped because the
table's unique constraint already covers (entity_id, date, api_key, model,
custom_llm_provider, mcp_namespaced_tool_name, endpoint) — making the
GROUP BY a no-op that wastes CPU on hash/sort for zero row reduction.
The Python _aggregate_spend_records function handles the final rollup.
When include_entity_id is True, the entity_id column is included in SELECT
to preserve per-entity breakdown in the results.
When include_entity_id is False (e.g. user/org endpoints), GROUP BY is
used to collapse rows across entities, which meaningfully reduces the
number of rows returned.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
@ -559,28 +559,58 @@ def _build_aggregated_sql_query(
entity_select = f'"{entity_id_field}",' if include_entity_id else ""
sql_query = f"""
SELECT
{entity_select}
date,
api_key,
model,
model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
spend::float AS spend,
prompt_tokens::bigint AS prompt_tokens,
completion_tokens::bigint AS completion_tokens,
cache_read_input_tokens::bigint AS cache_read_input_tokens,
cache_creation_input_tokens::bigint AS cache_creation_input_tokens,
api_requests::bigint AS api_requests,
successful_requests::bigint AS successful_requests,
failed_requests::bigint AS failed_requests
FROM "{pg_table}"
WHERE {where_clause}
ORDER BY date DESC
"""
if include_entity_id:
# When entity_id is in the result set, the GROUP BY columns would
# match the table's unique constraint — making aggregation a no-op.
# Skip GROUP BY entirely and let Python handle the rollup.
sql_query = f"""
SELECT
"{entity_id_field}",
date,
api_key,
model,
model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
spend::float AS spend,
prompt_tokens::bigint AS prompt_tokens,
completion_tokens::bigint AS completion_tokens,
cache_read_input_tokens::bigint AS cache_read_input_tokens,
cache_creation_input_tokens::bigint AS cache_creation_input_tokens,
api_requests::bigint AS api_requests,
successful_requests::bigint AS successful_requests,
failed_requests::bigint AS failed_requests
FROM "{pg_table}"
WHERE {where_clause}
ORDER BY date DESC
"""
else:
# When entity_id is excluded, GROUP BY collapses rows across
# entities, meaningfully reducing the result set.
sql_query = f"""
SELECT
date,
api_key,
model,
MAX(model_group) AS 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, custom_llm_provider,
mcp_namespaced_tool_name, endpoint
ORDER BY date DESC
"""
return sql_query, sql_params

View file

@ -473,43 +473,10 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
class TestBuildAggregatedSqlQuery:
"""Tests for _build_aggregated_sql_query to verify the plain SELECT (no GROUP BY) approach."""
"""Tests for _build_aggregated_sql_query conditional GROUP BY behavior."""
def test_basic_query_no_group_by(self):
"""Verify the query uses plain SELECT without GROUP BY."""
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
start_date="2024-01-01",
end_date="2024-01-31",
model=None,
api_key=None,
)
assert "GROUP BY" not in sql
assert "SUM(" not in sql
assert "ORDER BY date DESC" in sql
assert '"LiteLLM_DailyTeamSpend"' in sql
assert params == ["2024-01-01", "2024-01-31"]
def test_query_selects_columns_directly(self):
"""Verify metric columns are selected with casts, not aggregated."""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
start_date="2024-01-01",
end_date="2024-01-31",
model=None,
api_key=None,
)
assert "spend::float AS spend" in sql
assert "prompt_tokens::bigint AS prompt_tokens" in sql
assert "completion_tokens::bigint AS completion_tokens" in sql
assert "api_requests::bigint AS api_requests" in sql
def test_include_entity_id_true(self):
"""When include_entity_id=True, entity column should appear in SELECT."""
def test_include_entity_id_true_no_group_by(self):
"""When include_entity_id=True, skip GROUP BY (it's a no-op on unique constraint)."""
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
@ -520,11 +487,34 @@ class TestBuildAggregatedSqlQuery:
api_key=None,
include_entity_id=True,
)
assert '"team_id",' in sql
assert "GROUP BY" not in sql
assert "SUM(" not in sql
assert '"team_id",' in sql
assert "spend::float AS spend" in sql
assert "ORDER BY date DESC" in sql
assert '"LiteLLM_DailyTeamSpend"' in sql
assert params == ["2024-01-01", "2024-01-31"]
def test_include_entity_id_false(self):
"""When include_entity_id=False (default), entity column should not appear."""
def test_include_entity_id_false_uses_group_by(self):
"""When include_entity_id=False (default), use GROUP BY to collapse across entities."""
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
start_date="2024-01-01",
end_date="2024-01-31",
model=None,
api_key=None,
include_entity_id=False,
)
assert "GROUP BY" in sql
assert "SUM(spend)" in sql
assert '"team_id",' not in sql
assert "ORDER BY date DESC" in sql
assert params == ["2024-01-01", "2024-01-31"]
def test_group_by_excludes_model_group(self):
"""GROUP BY should not include model_group (uses MAX instead)."""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
@ -535,7 +525,27 @@ class TestBuildAggregatedSqlQuery:
api_key=None,
include_entity_id=False,
)
assert '"team_id",' not in sql
assert "MAX(model_group)" in sql
# model_group should not appear in GROUP BY clause
group_by_clause = sql.split("GROUP BY")[1].split("ORDER BY")[0]
assert "model_group" not in group_by_clause
def test_include_entity_id_true_selects_columns_directly(self):
"""With include_entity_id=True, metric columns are selected directly (no SUM)."""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
start_date="2024-01-01",
end_date="2024-01-31",
model=None,
api_key=None,
include_entity_id=True,
)
assert "spend::float AS spend" in sql
assert "prompt_tokens::bigint AS prompt_tokens" in sql
assert "completion_tokens::bigint AS completion_tokens" in sql
assert "api_requests::bigint AS api_requests" in sql
def test_entity_id_single_value_filter(self):
"""Single entity_id value should produce an = condition."""
@ -689,7 +699,7 @@ class TestBuildAggregatedSqlQuery:
)
def test_all_table_names_supported(self):
"""All known table names should produce valid queries."""
"""All known table names should produce valid queries in both modes."""
tables = [
("litellm_dailyuserspend", "user_id", "LiteLLM_DailyUserSpend"),
("litellm_dailyteamspend", "team_id", "LiteLLM_DailyTeamSpend"),
@ -703,7 +713,8 @@ class TestBuildAggregatedSqlQuery:
("litellm_dailytagspend", "tag", "LiteLLM_DailyTagSpend"),
]
for table_name, entity_field, pg_table in tables:
sql, params = _build_aggregated_sql_query(
# include_entity_id=True: no GROUP BY
sql, _ = _build_aggregated_sql_query(
table_name=table_name,
entity_id_field=entity_field,
entity_id=None,
@ -711,6 +722,21 @@ class TestBuildAggregatedSqlQuery:
end_date="2024-01-31",
model=None,
api_key=None,
include_entity_id=True,
)
assert f'FROM "{pg_table}"' in sql
assert "GROUP BY" not in sql
# include_entity_id=False: uses GROUP BY
sql, _ = _build_aggregated_sql_query(
table_name=table_name,
entity_id_field=entity_field,
entity_id=None,
start_date="2024-01-01",
end_date="2024-01-31",
model=None,
api_key=None,
include_entity_id=False,
)
assert f'FROM "{pg_table}"' in sql
assert "GROUP BY" in sql