From 46ef016f2fc43daf75fb718a7f0478e68b01d02d Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Fri, 11 Sep 2026 10:46:32 -0400 Subject: [PATCH 1/3] perf(proxy): make /spend/logs/ui session listing use an index and run its queries concurrently --- .../migration.sql | 14 ++++ .../litellm_proxy_extras/schema.prisma | 1 + litellm/proxy/schema.prisma | 1 + .../spend_management_endpoints.py | 65 ++++++++++--------- schema.prisma | 1 + .../test_spend_management_endpoints.py | 10 +++ 6 files changed, 63 insertions(+), 29 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260911000000_spend_logs_session_window_index/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260911000000_spend_logs_session_window_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260911000000_spend_logs_session_window_index/migration.sql new file mode 100644 index 00000000000..ccf14492cae --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260911000000_spend_logs_session_window_index/migration.sql @@ -0,0 +1,14 @@ +-- CreateIndex (CONCURRENTLY) +-- +-- Covering index for the session-grouped /spend/logs/ui listing. The page, +-- count, and representative queries all GROUP BY +-- COALESCE(NULLIF(session_id, ''), request_id), api_key over a startTime +-- window; with only the plain startTime index they aggregate full heap rows +-- (hundreds of KB each from messages/response), which is the 40s+ page load. +-- This index makes those window scans index-only. +-- +-- CREATE INDEX CONCURRENTLY cannot run inside a transaction, so this +-- migration must stay a single statement (see +-- 20260415120000_health_check_latest_per_model_index for the full +-- disclaimer). +CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_session_window_idx" ON "LiteLLM_SpendLogs"("startTime", session_id, request_id, api_key, status); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 817df082d8c..0a990c05c56 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -665,6 +665,7 @@ model LiteLLM_SpendLogs { @@index([startTime, request_id]) @@index([end_user]) @@index([session_id]) + @@index([startTime, session_id, request_id, api_key, status], map: "LiteLLM_SpendLogs_session_window_idx") } model LiteLLM_BudgetWindowSpend { diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 817df082d8c..0a990c05c56 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -665,6 +665,7 @@ model LiteLLM_SpendLogs { @@index([startTime, request_id]) @@index([end_user]) @@index([session_id]) + @@index([startTime, session_id, request_id, api_key, status], map: "LiteLLM_SpendLogs_session_window_idx") } model LiteLLM_BudgetWindowSpend { diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index da79328fa59..c2e404deddd 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1,8 +1,9 @@ #### SPEND MANAGEMENT ##### +import asyncio import collections import json import os -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Mapping, Sequence from datetime import date, datetime, timedelta, timezone from itertools import groupby from types import MappingProxyType @@ -2821,13 +2822,6 @@ async def ui_view_spend_logs( LIMIT ${p} ) AS bounded_matches """ - count_rows: Final[Sequence[_SpendLogsCountRow] | None] = await _query_raw_or_none( - prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 - ) - raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 - total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP - total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total - sql_query: Final = ( f""" SELECT * FROM ( @@ -2850,17 +2844,20 @@ async def ui_view_spend_logs( LIMIT ${p} OFFSET ${p + 1} """ ) - sql_params.extend([page_size, skip]) - - data: Final = await prisma_client.db.query_raw(sql_query, *sql_params) + count_task: Final[Awaitable[Sequence[_SpendLogsCountRow] | None]] = _query_raw_or_none( + prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 + ) + data_task: Final = prisma_client.db.query_raw(sql_query, *sql_params, page_size, skip) + count_rows, data = await asyncio.gather(count_task, data_task) + raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 + total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP + total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total _hydrate_spend_log_metadata(data) # Calculate total pages total_pages: Final = (total_records + page_size - 1) // page_size - verbose_proxy_logger.debug("data= %s", json.dumps(data, indent=4, default=str)) - return await _build_ui_spend_logs_response( prisma_client, data, @@ -2899,15 +2896,25 @@ async def _fetch_session_representatives( next_param_index: int, session_keys: Sequence[tuple[str, str]], ) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row - """Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order.""" + """Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order. + + The SQL matches candidates on the plain ``session_id`` / ``request_id`` / + ``api_key`` columns so the ``(startTime, session_id, request_id, api_key)`` + index applies; a row-tuple ``IN`` on the COALESCE group key forces a full + window scan. That predicate admits a superset of the requested pairs (the + cross product of session ids and api keys, plus rows whose request_id + collides with a requested key), so the exact ``(session_key, api_key)`` + pairing is restored by the ``rep_by_key`` lookup below. + """ rep_query: Final = f""" SELECT * FROM ( SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL}) {_SPEND_LOG_LIST_COLUMNS} FROM "LiteLLM_SpendLogs" WHERE {where_clause} - AND ({_SESSION_GROUP_KEY_SQL}) IN ( - SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[]) + AND ( + (session_id = ANY(${next_param_index}::text[]) AND api_key = ANY(${next_param_index + 1}::text[])) + OR request_id = ANY(${next_param_index}::text[]) ) ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC ) AS session_representatives @@ -2975,18 +2982,6 @@ async def _ui_session_grouped_spend_logs( ORDER BY MAX("startTime") {direction}, {_SESSION_KEY_EXPR} {direction}, api_key {direction} LIMIT ${limit_index} """ - page_rows: Final[Sequence[_SessionPageRow]] = await _query_raw( - prisma_client, page_query, *sql_params, *cursor_params, page_size + 1 - ) - - has_more: Final = len(page_rows) > page_size - visible_rows: Final = page_rows[:page_size] - next_cursor: Final = ( - f"{visible_rows[-1]['last_activity']}|{visible_rows[-1]['api_key']}|{visible_rows[-1]['session_key']}" - if has_more and visible_rows - else None - ) - count_query: Final = f""" SELECT COUNT(*) AS total_count FROM ( @@ -2997,9 +2992,21 @@ async def _ui_session_grouped_spend_logs( LIMIT ${next_param_index} ) AS bounded_sessions """ - count_rows: Final[Sequence[_SpendLogsCountRow]] = await _query_raw( + page_rows_task: Final[Awaitable[Sequence[_SessionPageRow]]] = _query_raw( + prisma_client, page_query, *sql_params, *cursor_params, page_size + 1 + ) + count_rows_task: Final[Awaitable[Sequence[_SpendLogsCountRow]]] = _query_raw( prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 ) + page_rows, count_rows = await asyncio.gather(page_rows_task, count_rows_task) + + has_more: Final = len(page_rows) > page_size + visible_rows: Final = page_rows[:page_size] + next_cursor: Final = ( + f"{visible_rows[-1]['last_activity']}|{visible_rows[-1]['api_key']}|{visible_rows[-1]['session_key']}" + if has_more and visible_rows + else None + ) raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total diff --git a/schema.prisma b/schema.prisma index 817df082d8c..0a990c05c56 100644 --- a/schema.prisma +++ b/schema.prisma @@ -665,6 +665,7 @@ model LiteLLM_SpendLogs { @@index([startTime, request_id]) @@index([end_user]) @@index([session_id]) + @@index([startTime, session_id, request_id, api_key, status], map: "LiteLLM_SpendLogs_session_window_idx") } model LiteLLM_BudgetWindowSpend { diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 671a8ae63fc..12419a61c99 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -6640,6 +6640,10 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc reps = [ _session_representative_row("req-solo", None), _session_representative_row("req-1", "sess-1"), + # The plain-column predicate admits the cross product of session ids and + # api keys, so the endpoint must drop rows whose (session_key, api_key) + # pair was never requested. + _session_representative_row("req-cross-product", "sess-unrequested"), ] mock_prisma = _session_grouped_mock_prisma(page_rows, 3, reps) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -6689,6 +6693,12 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc f"ORDER BY {SESSION_GROUP_KEY_SQL}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in rep_call[0] ), "the session representative must prefer the newest non-MCP call" + assert "session_id = ANY(" in rep_call[0] and "OR request_id = ANY(" in rep_call[0], ( + "the representative lookup must filter on the plain session_id/request_id columns so the " + "(startTime, session_id, request_id, api_key, status) index applies; a row-tuple IN on the " + "COALESCE group key forces a full window scan (the 40s+ logs page)" + ) + assert f"({SESSION_GROUP_KEY_SQL}) IN" not in rep_call[0] assert rep_call[-2] == ["sess-1", "req-solo"] assert rep_call[-1] == ["hashed-key", "hashed-key"] finally: From 3001970e366e2ed39a3c0931b22acdc6b44f1f50 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 19 Sep 2026 14:02:16 -0400 Subject: [PATCH 2/3] fix(proxy): type the concurrent flat-page count as a task and unpin SQL emission order in the offset test Co-Authored-By: Claude Fable 5 --- .../proxy/spend_tracking/spend_management_endpoints.py | 10 +++++----- .../spend_tracking/test_spend_management_endpoints.py | 5 +++-- 2 files changed, 8 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 63e5cc4672a..0da193fd904 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3,7 +3,7 @@ import asyncio import collections import json import os -from collections.abc import Awaitable, Mapping, Sequence +from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import date, datetime, timedelta, timezone from itertools import groupby @@ -2880,11 +2880,11 @@ async def ui_view_spend_logs( LIMIT ${p} OFFSET ${p + 1} """ ) - count_task: Final[Awaitable[Sequence[_SpendLogsCountRow] | None]] = _query_raw_or_none( - prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 + count_task: Final[asyncio.Task[Sequence[_SpendLogsCountRow] | None]] = asyncio.create_task( + _query_raw_or_none(prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1) ) - data_task: Final = prisma_client.db.query_raw(sql_query, *sql_params, page_size, skip) - count_rows, data = await asyncio.gather(count_task, data_task) + data: Final = await prisma_client.db.query_raw(sql_query, *sql_params, page_size, skip) + count_rows: Final = await count_task raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 344dcabe8a8..445ea205881 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -7717,8 +7717,9 @@ async def test_ui_view_spend_logs_group_by_session_offset_for_non_starttime_sort assert "next_session_cursor" not in data emitted_sql = [call.args[0] for call in mock_prisma.db.query_raw.await_args_list] assert "HAVING" not in " ".join(emitted_sql) - assert f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in emitted_sql[1] - assert "OFFSET" in emitted_sql[1] + page_sql = [sql for sql in emitted_sql if f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in sql] + assert len(page_sql) == 1 + assert "OFFSET" in page_sql[0] finally: app.dependency_overrides.pop(ps.user_api_key_auth, None) From 5ebcf8ece161d6dfc5e4c49017d130dc8c6147f3 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 19 Sep 2026 16:31:07 -0400 Subject: [PATCH 3/3] fix(tests): dispatch spend logs UI mocks on SQL content instead of call order The count and page queries run concurrently, so mocks and assertions keyed on positional call order break; ruff format on the merged endpoint. Co-Authored-By: Claude Fable 5 --- .../spend_management_endpoints.py | 4 +- .../test_spend_query_optimization.py | 63 +++++++++---------- 2 files changed, 29 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0da193fd904..cc23dd8eda7 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3085,9 +3085,7 @@ async def _ui_session_grouped_spend_logs( page_ends_the_list: Final = cursor is None and page_limit > 0 and not has_more and page_starts_inside_the_list if page_ends_the_list: count_task.cancel() - total_records, total_is_capped = ( - (offset + len(page_rows), False) if page_ends_the_list else await count_task - ) + total_records, total_is_capped = (offset + len(page_rows), False) if page_ends_the_list else await count_task session_keys: Final = tuple((row["session_key"], row["api_key"]) for row in visible_rows) data: Final[list[dict[str, object]]] = ( # mutable-ok: _build_ui_spend_logs_response writes onto each row diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py index 54e5a6d5385..fd4fe757af7 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py @@ -203,20 +203,31 @@ async def test_spend_logs_ui_wraps_params_in_at_time_zone_utc(monkeypatch): def _make_ui_spend_logs_mock(count_total, page_rows): """ - Build a prisma mock whose first `query_raw` (the bounded count) returns - `count_total` and whose second `query_raw` (the page data) returns - `page_rows`. + Build a prisma mock whose bounded-count `query_raw` returns `count_total` + and whose page-data `query_raw` returns `page_rows`, dispatching on the + SQL so the two queries may run in either order (they run concurrently). """ + + async def mock_query_raw(sql_query, *params): + if "COUNT(*) AS total_count" in sql_query: + return [{"total_count": count_total}] + return page_rows + mock_prisma = MagicMock() mock_prisma.db = MagicMock() - mock_prisma.db.query_raw = AsyncMock( - side_effect=[[{"total_count": count_total}], page_rows] - ) + mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw) mock_prisma.db.litellm_spendlogs = MagicMock() mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) return mock_prisma +def _query_raw_call_matching(mock_prisma, sql_fragment): + """Return the single query_raw call whose SQL contains ``sql_fragment``.""" + matches = [call for call in mock_prisma.db.query_raw.call_args_list if sql_fragment in call[0][0]] + assert len(matches) == 1, f"expected exactly one query containing {sql_fragment!r}, got {len(matches)}" + return matches[0] + + @pytest.mark.asyncio async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch): """ @@ -260,21 +271,17 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch): mock_prisma.db.litellm_spendlogs.count.assert_not_called() - count_call = mock_prisma.db.query_raw.call_args_list[0] + count_call = _query_raw_call_matching(mock_prisma, "COUNT(*) AS total_count") count_sql = count_call[0][0] assert "COUNT(*) OVER ()" not in count_sql assert "LIMIT" in count_sql and "FROM (" in count_sql, ( - "the total must come from a bounded subquery count, not a full-window " - f"scan. SQL was:\n{count_sql}" - ) - assert count_call[0][-1] == SPEND_LOGS_PAGINATION_COUNT_CAP + 1, ( - "the bounded count must probe at most cap+1 rows" + f"the total must come from a bounded subquery count, not a full-window scan. SQL was:\n{count_sql}" ) + assert count_call[0][-1] == SPEND_LOGS_PAGINATION_COUNT_CAP + 1, "the bounded count must probe at most cap+1 rows" - page_sql = mock_prisma.db.query_raw.call_args_list[1][0][0] + page_sql = _query_raw_call_matching(mock_prisma, "ORDER BY")[0][0] assert "COUNT(*) OVER ()" not in page_sql, ( - "the page query must not carry a window count that forces a full-window " - f"scan. SQL was:\n{page_sql}" + f"the page query must not carry a window count that forces a full-window scan. SQL was:\n{page_sql}" ) assert "GROUP BY" not in count_sql and "DISTINCT ON" not in page_sql, ( "without group_by_session the endpoint must keep raw per-call pagination" @@ -302,9 +309,7 @@ async def test_spend_logs_ui_caps_total_for_large_result_sets(monkeypatch): ) page_rows = [{"request_id": "req-1", "metadata": "{}", "session_id": None}] - mock_prisma = _make_ui_spend_logs_mock( - count_total=SPEND_LOGS_PAGINATION_COUNT_CAP + 1, page_rows=page_rows - ) + mock_prisma = _make_ui_spend_logs_mock(count_total=SPEND_LOGS_PAGINATION_COUNT_CAP + 1, page_rows=page_rows) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") @@ -343,13 +348,7 @@ async def test_spend_logs_ui_empty_page_reports_zero_total(monkeypatch): ui_view_spend_logs, ) - # First query_raw call is the bounded count (0 matches), second is the empty - # page. - mock_prisma = MagicMock() - mock_prisma.db = MagicMock() - mock_prisma.db.query_raw = AsyncMock(side_effect=[[{"total_count": 0}], []]) - mock_prisma.db.litellm_spendlogs = MagicMock() - mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) + mock_prisma = _make_ui_spend_logs_mock(count_total=0, page_rows=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -391,13 +390,7 @@ async def test_spend_logs_ui_out_of_range_page_keeps_total(monkeypatch): ui_view_spend_logs, ) - # First query_raw call is the bounded count (7 matches), second is the - # out-of-range page (empty). - mock_prisma = MagicMock() - mock_prisma.db = MagicMock() - mock_prisma.db.query_raw = AsyncMock(side_effect=[[{"total_count": 7}], []]) - mock_prisma.db.litellm_spendlogs = MagicMock() - mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) + mock_prisma = _make_ui_spend_logs_mock(count_total=7, page_rows=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -634,10 +627,10 @@ async def test_spend_logs_ui_group_by_session_offset_pages_for_other_sorts(monke ) group_key = "COALESCE(NULLIF(session_id, ''), request_id), api_key" - count_sql = mock_prisma.db.query_raw.call_args_list[0][0][0] + count_sql = _query_raw_call_matching(mock_prisma, "COUNT(*) AS total_count")[0][0] assert f"GROUP BY {group_key}" in count_sql - page_call = mock_prisma.db.query_raw.call_args_list[1][0] + page_call = _query_raw_call_matching(mock_prisma, f"DISTINCT ON ({group_key})")[0] page_sql = page_call[0] assert f"DISTINCT ON ({group_key})" in page_sql assert "ORDER BY spend DESC" in page_sql @@ -682,7 +675,7 @@ async def test_spend_logs_ui_request_id_lookup_with_grouping_returns_exact_row(m group_by_session=True, ) - page_call = mock_prisma.db.query_raw.call_args_list[1] + page_call = _query_raw_call_matching(mock_prisma, "DISTINCT ON") assert "request_id = $" in page_call[0][0], "the request_id equality filter must survive grouping" assert "req-deep-link" in page_call[0] assert [row["request_id"] for row in response["data"]] == ["req-deep-link"]