From 93ffd74c92a27dcf990156bbbca8227842100f48 Mon Sep 17 00:00:00 2001 From: yassin Date: Sun, 23 Aug 2026 10:12:10 +0000 Subject: [PATCH] fix(proxy): batch spend log reads on deprecated /spend/logs to bound Prisma query engine memory Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/workflows/test-unit.yml | 6 +- litellm/constants.py | 1 + .../spend_management_endpoints.py | 41 +++-- .../test_spend_management_endpoints.py | 145 ++++++++++++++++++ 4 files changed, 181 insertions(+), 12 deletions(-) diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index a7c67f2b35d..2dfca3d308f 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -211,7 +211,7 @@ jobs: workers: 2 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 55 + job-timeout-minutes: 60 - shard: proxy-extras artifact-name: proxy-extras @@ -219,7 +219,7 @@ jobs: workers: 2 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 55 + job-timeout-minutes: 60 - shard: enterprise-package artifact-name: enterprise-package @@ -227,7 +227,7 @@ jobs: workers: 4 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 55 + job-timeout-minutes: 60 - shard: responses-caching-types artifact-name: responses-caching-types diff --git a/litellm/constants.py b/litellm/constants.py index aaaddd063e7..eaca1f584b4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1539,6 +1539,7 @@ SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) SPEND_LOG_WRITE_BATCH_MAX_ROWS: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_ROWS", "100"))) +SPEND_LOG_READ_BATCH_MAX_ROWS: Final = 1000 SPEND_LOG_QUEUE_SIZE_THRESHOLD: Final = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) SPEND_LOG_QUEUE_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_QUEUE_MAX_BYTES", "64000000"))) SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index ed2ecd8325a..561a6f33819 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1,5 +1,6 @@ #### SPEND MANAGEMENT ##### import collections +import itertools import json import os from collections.abc import Mapping, Sequence @@ -11,6 +12,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import SPEND_LOG_READ_BATCH_MAX_ROWS from litellm.proxy._types import * from litellm.proxy._types import ProviderBudgetResponse, ProviderBudgetResponseObject from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -46,6 +48,10 @@ class _SupportsModelDump(Protocol): def model_dump(self) -> Mapping[str, object]: ... +class _SpendLogRow(_SupportsModelDump, Protocol): + request_id: str + + class _SpendLogOwnershipRow(Protocol): user: str | None team_id: str | None @@ -153,8 +159,14 @@ class _SpendLogsTable(Protocol): """The subset of the Prisma spend-logs table API this module uses.""" async def find_many( - self, *, where: Mapping[str, object], order: Mapping[str, str] - ) -> Sequence[_SupportsModelDump]: ... + self, + *, + where: Mapping[str, object], + order: Sequence[Mapping[str, str]], + take: int, + skip: int, + cursor: Mapping[str, str] | None, + ) -> Sequence[_SpendLogRow]: ... async def find_unique( self, *, where: Mapping[str, object], include: None = None @@ -199,9 +211,25 @@ async def _find_spend_logs( prisma_client: PrismaClient, where: Mapping[str, object], order: Mapping[str, str], + batch_size: int = SPEND_LOG_READ_BATCH_MAX_ROWS, ) -> Sequence[_SupportsModelDump]: - """Read spend log rows as Prisma model instances.""" - return await _spend_logs_table(prisma_client).find_many(where=where, order=order) + """Read spend log rows as Prisma model instances in bounded keyset-paginated batches.""" + table: Final = _spend_logs_table(prisma_client) + stable_order: Final = [dict(order), {"request_id": "desc"}] # mutable-ok: Prisma expects order as a list of dicts + batches: Final[list[Sequence[_SpendLogRow]]] = [] # mutable-ok: accumulator for paged async reads + cursor: Mapping[str, str] | None = None + while True: + batch = await table.find_many( + where=where, + order=stable_order, + take=batch_size, + skip=0 if cursor is None else 1, + cursor=cursor, + ) + batches.append(batch) + if len(batch) < batch_size: + return tuple(itertools.chain.from_iterable(batches)) + cursor = {"request_id": batch[-1].request_id} # rebind-ok: cursor advances to the last row of each page async def _find_spend_log_row(prisma_client: PrismaClient, request_id: str) -> _SpendLogOwnershipRow | None: @@ -2921,7 +2949,6 @@ async def view_spend_logs( raise Exception( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - spend_logs = [] if ( start_date is not None and isinstance(start_date, str) @@ -3029,10 +3056,6 @@ async def view_spend_logs( if user_id is not None and isinstance(user_id, str): scoped_filter["user"] = user_id - if not scoped_filter: - spend_logs = await prisma_client.get_data(table_name="spend", query_type="find_all") - return spend_logs - data = await _find_spend_logs( prisma_client, where=scoped_filter, 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 2b062d9020d..d21151a8962 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 @@ -3135,6 +3135,151 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch): app.dependency_overrides.pop(ps.user_api_key_auth, None) +class _FakeSpendLogRow(dict): + """Dict-backed row that also exposes request_id as an attribute like a Prisma model.""" + + @property + def request_id(self): + return self["request_id"] + + +class _BatchRecordingSpendLogsDB: + """Fake db whose litellm_spendlogs.find_many honors take/skip/cursor and records every call.""" + + def __init__(self, rows): + self.rows = rows + self.calls = [] + self.litellm_spendlogs = self + + async def find_many(self, **kwargs): + self.calls.append(kwargs) + take = kwargs["take"] + skip = kwargs["skip"] + cursor = kwargs["cursor"] + start = ( + 0 + if cursor is None + else next(i for i, r in enumerate(self.rows) if r["request_id"] == cursor["request_id"]) + ) + return self.rows[start + skip : start + skip + take] + + +@pytest.mark.asyncio +async def test_find_spend_logs_batches_reads(monkeypatch): + """Regression for LIT-4765: no single Prisma statement may carry an unbounded row payload.""" + from litellm.proxy.spend_tracking.spend_management_endpoints import _find_spend_logs + + rows = [_FakeSpendLogRow({"request_id": f"req-{i}", "spend": 0.01}) for i in range(5)] + db = _BatchRecordingSpendLogsDB(rows) + + class MockPrismaClient: + def __init__(self): + self.db = db + + result = await _find_spend_logs( + MockPrismaClient(), + where={}, + order={"startTime": "desc"}, + batch_size=2, + ) + + assert list(result) == rows + assert [c["skip"] for c in db.calls] == [0, 1, 1] + assert [c["cursor"] for c in db.calls] == [None, {"request_id": "req-1"}, {"request_id": "req-3"}] + assert all(c["take"] == 2 for c in db.calls) + assert all(c["order"] == [{"startTime": "desc"}, {"request_id": "desc"}] for c in db.calls) + + +@pytest.mark.asyncio +async def test_find_spend_logs_is_stable_under_concurrent_inserts(monkeypatch): + """Regression for LIT-4765: rows inserted mid-pagination must not duplicate or drop results.""" + from litellm.proxy.spend_tracking.spend_management_endpoints import _find_spend_logs + + rows = [_FakeSpendLogRow({"request_id": f"req-{i}", "spend": 0.01}) for i in range(5)] + db = _BatchRecordingSpendLogsDB(list(rows)) + + original_find_many = db.find_many + + async def find_many_with_insert(**kwargs): + batch = await original_find_many(**kwargs) + if len(db.calls) == 1: + db.rows.insert(0, _FakeSpendLogRow({"request_id": "req-new", "spend": 0.01})) + return batch + + db.find_many = find_many_with_insert + + class MockPrismaClient: + def __init__(self): + self.db = db + + result = await _find_spend_logs( + MockPrismaClient(), + where={}, + order={"startTime": "desc"}, + batch_size=2, + ) + + assert [r["request_id"] for r in result] == [f"req-{i}" for i in range(5)] + + +@pytest.mark.asyncio +async def test_view_spend_logs_no_filter_uses_bounded_reads(client, monkeypatch): + """Regression for LIT-4765: the no-filter /spend/logs path must not read the whole table in one statement.""" + from litellm.constants import SPEND_LOG_READ_BATCH_MAX_ROWS + + row_count = SPEND_LOG_READ_BATCH_MAX_ROWS * 2 + SPEND_LOG_READ_BATCH_MAX_ROWS // 2 + rows = [_FakeSpendLogRow({"request_id": f"req-{i}", "spend": 0.01}) for i in range(row_count)] + db = _BatchRecordingSpendLogsDB(rows) + + class MockPrismaClient: + def __init__(self): + self.db = db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + response = client.get("/spend/logs", headers={"Authorization": "Bearer sk-test"}) + assert response.status_code == 200 + assert len(response.json()) == row_count + assert len(db.calls) == 3 + assert all(c["take"] == SPEND_LOG_READ_BATCH_MAX_ROWS for c in db.calls) + assert all(c["where"] == {} for c in db.calls) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_view_spend_logs_user_filter_uses_bounded_reads(client, monkeypatch): + """Regression for LIT-4765: the user_id-filtered /spend/logs path must batch its reads too.""" + from litellm.constants import SPEND_LOG_READ_BATCH_MAX_ROWS + + row_count = SPEND_LOG_READ_BATCH_MAX_ROWS + 1 + rows = [_FakeSpendLogRow({"request_id": f"req-{i}", "user": "u1", "spend": 0.01}) for i in range(row_count)] + db = _BatchRecordingSpendLogsDB(rows) + + class MockPrismaClient: + def __init__(self): + self.db = db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + response = client.get( + "/spend/logs", params={"user_id": "u1"}, headers={"Authorization": "Bearer sk-test"} + ) + assert response.status_code == 200 + assert len(response.json()) == row_count + assert len(db.calls) == 2 + assert all(c["take"] == SPEND_LOG_READ_BATCH_MAX_ROWS for c in db.calls) + assert all(c["where"] == {"user": "u1"} for c in db.calls) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_view_spend_tags(client, monkeypatch): """Test the /spend/tags endpoint"""