mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
f005afa146
commit
93ffd74c92
4 changed files with 181 additions and 12 deletions
6
.github/workflows/test-unit.yml
vendored
6
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue