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>
This commit is contained in:
yassin 2026-08-23 10:12:10 +00:00
parent f005afa146
commit 93ffd74c92
4 changed files with 181 additions and 12 deletions

View file

@ -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

View file

@ -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))

View file

@ -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,

View file

@ -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"""