Mitigate deprecated spend logs Prisma OOM path

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-26 02:02:10 +00:00
parent 9806e21871
commit 8529fc0b40
3 changed files with 205 additions and 56 deletions

View file

@ -2,7 +2,7 @@
import collections
import json
import os
from datetime import datetime, timedelta, timezone
from datetime import date, datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
import fastapi
@ -2074,6 +2074,16 @@ async def view_spend_logs( # noqa: PLR0915
default=True,
description="When start_date and end_date are provided, summarize=true returns aggregated data by date (legacy behavior), summarize=false returns filtered individual logs",
),
limit: Optional[int] = fastapi.Query(
default=None,
ge=1,
description="Deprecated endpoint safety cap for number of rows returned. Capped by SPEND_LOGS_DEPRECATED_MAX_ROWS (default: 1000).",
),
offset: int = fastapi.Query(
default=0,
ge=0,
description="Deprecated endpoint row offset for basic pagination.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
@ -2124,6 +2134,20 @@ async def view_spend_logs( # noqa: PLR0915
):
user_id = user_api_key_dict.user_id
max_deprecated_rows = max(
1, int(os.getenv("SPEND_LOGS_DEPRECATED_MAX_ROWS", "1000"))
)
requested_limit = limit if limit is not None else max_deprecated_rows
safe_limit = min(requested_limit, max_deprecated_rows)
if requested_limit > max_deprecated_rows:
verbose_proxy_logger.warning(
"/spend/logs requested limit=%s exceeds SPEND_LOGS_DEPRECATED_MAX_ROWS=%s. "
"Clamping to %s.",
requested_limit,
max_deprecated_rows,
safe_limit,
)
try:
verbose_proxy_logger.debug("inside view_spend_logs")
if prisma_client is None:
@ -2156,8 +2180,16 @@ async def view_spend_logs( # noqa: PLR0915
}
}
if api_key is not None and isinstance(api_key, str):
filter_query["api_key"] = api_key # type: ignore
effective_api_key = api_key
if (
effective_api_key is not None
and isinstance(effective_api_key, str)
and effective_api_key.startswith("sk-")
):
effective_api_key = prisma_client.hash_token(token=effective_api_key)
if effective_api_key is not None and isinstance(effective_api_key, str):
filter_query["api_key"] = effective_api_key # type: ignore
elif request_id is not None and isinstance(request_id, str):
filter_query["request_id"] = request_id # type: ignore
elif user_id is not None and isinstance(user_id, str):
@ -2171,18 +2203,48 @@ async def view_spend_logs( # noqa: PLR0915
order={
"startTime": "desc",
},
take=safe_limit,
skip=offset,
)
if len(data) >= safe_limit:
verbose_proxy_logger.warning(
"/spend/logs summarize=false returned %s rows (safety cap reached). "
"Use /spend/logs/v2 for full pagination.",
safe_limit,
)
return data
# Legacy behavior: return summarized data (when summarize=true)
# SQL query
response = await prisma_client.db.litellm_spendlogs.group_by(
by=["api_key", "user", "model", "startTime"],
where=filter_query, # type: ignore
sum={
"spend": True,
},
)
# Aggregate in SQL by DATE(startTime) to avoid near row-level group cardinality.
sql_conditions: List[str] = ['"startTime" >= $1', '"startTime" <= $2']
sql_params: List[Any] = [start_date_obj, end_date_obj]
param_idx = 3
if effective_api_key is not None and isinstance(effective_api_key, str):
sql_conditions.append(f"api_key = ${param_idx}")
sql_params.append(effective_api_key)
param_idx += 1
elif request_id is not None and isinstance(request_id, str):
sql_conditions.append(f"request_id = ${param_idx}")
sql_params.append(request_id)
param_idx += 1
elif user_id is not None and isinstance(user_id, str):
sql_conditions.append(f'"user" = ${param_idx}')
sql_params.append(user_id)
param_idx += 1
sql_query = f"""
SELECT
DATE("startTime") AS spend_date,
api_key,
"user",
model,
SUM(spend)::double precision AS spend
FROM "LiteLLM_SpendLogs"
WHERE {' AND '.join(sql_conditions)}
GROUP BY DATE("startTime"), api_key, "user", model
ORDER BY DATE("startTime")
"""
response = await prisma_client.db.query_raw(sql_query, *sql_params)
if (
isinstance(response, list)
@ -2191,25 +2253,37 @@ async def view_spend_logs( # noqa: PLR0915
):
result: dict = {}
for record in response:
dt_object = datetime.strptime(str(record["startTime"]), "%Y-%m-%dT%H:%M:%S.%fZ") # type: ignore
date = dt_object.date()
if date not in result:
result[date] = {"users": {}, "models": {}}
api_key = record["api_key"] # type: ignore
user_id = record["user"] # type: ignore
model = record["model"] # type: ignore
result[date]["spend"] = result[date].get("spend", 0) + record.get(
"_sum", {}
).get("spend", 0)
result[date][api_key] = result[date].get(api_key, 0) + record.get(
"_sum", {}
).get("spend", 0)
result[date]["users"][user_id] = result[date]["users"].get(
user_id, 0
) + record.get("_sum", {}).get("spend", 0)
result[date]["models"][model] = result[date]["models"].get(
model, 0
) + record.get("_sum", {}).get("spend", 0)
spend_date_value = record.get("spend_date")
date_key: Optional[date] = None
if isinstance(spend_date_value, datetime):
date_key = spend_date_value.date()
elif isinstance(spend_date_value, date):
date_key = spend_date_value
elif isinstance(spend_date_value, str):
try:
date_key = datetime.fromisoformat(spend_date_value).date()
except ValueError:
date_key = datetime.strptime(
spend_date_value, "%Y-%m-%d"
).date()
if date_key is None:
continue
if date_key not in result:
result[date_key] = {"users": {}, "models": {}}
_api_key = record.get("api_key") or ""
_user_id = record.get("user") or ""
_model = record.get("model") or ""
spend = float(record.get("spend") or 0)
result[date_key]["spend"] = result[date_key].get("spend", 0) + spend
result[date_key][_api_key] = result[date_key].get(_api_key, 0) + spend
result[date_key]["users"][_user_id] = result[date_key]["users"].get(
_user_id, 0
) + spend
result[date_key]["models"][_model] = result[date_key]["models"].get(
_model, 0
) + spend
return_list = []
final_date = None
for k, v in sorted(result.items()):
@ -2244,10 +2318,18 @@ async def view_spend_logs( # noqa: PLR0915
table_name="spend",
query_type="find_all",
key_val={"key": "api_key", "value": hashed_token},
limit=safe_limit,
offset=offset,
)
if spend_log is None:
return []
if isinstance(spend_log, list):
if len(spend_log) >= safe_limit:
verbose_proxy_logger.warning(
"/spend/logs api_key filter returned %s rows (safety cap reached). "
"Use /spend/logs/v2 for full pagination.",
safe_limit,
)
return spend_log
else:
return [spend_log]
@ -2265,17 +2347,34 @@ async def view_spend_logs( # noqa: PLR0915
table_name="spend",
query_type="find_all",
key_val={"key": "user", "value": user_id},
limit=safe_limit,
offset=offset,
)
if spend_log is None:
return []
if isinstance(spend_log, list):
if len(spend_log) >= safe_limit:
verbose_proxy_logger.warning(
"/spend/logs user_id filter returned %s rows (safety cap reached). "
"Use /spend/logs/v2 for full pagination.",
safe_limit,
)
return spend_log
else:
return [spend_log]
else:
spend_logs = await prisma_client.get_data(
table_name="spend", query_type="find_all"
table_name="spend",
query_type="find_all",
limit=safe_limit,
offset=offset,
)
if isinstance(spend_logs, list) and len(spend_logs) >= safe_limit:
verbose_proxy_logger.warning(
"/spend/logs returned %s rows (safety cap reached). "
"Use /spend/logs/v2 for full pagination.",
safe_limit,
)
return spend_logs

View file

@ -2731,6 +2731,11 @@ class PrismaClient:
verbose_proxy_logger.debug(
"PrismaClient: get_data: table_name == 'spend'"
)
spend_find_many_kwargs: Dict[str, Any] = {"order": {"startTime": "desc"}}
if isinstance(limit, int):
spend_find_many_kwargs["take"] = limit
if isinstance(offset, int) and offset >= 0:
spend_find_many_kwargs["skip"] = offset
if key_val is not None:
if query_type == "find_unique":
response = await self.db.litellm_spendlogs.find_unique( # type: ignore
@ -2742,12 +2747,13 @@ class PrismaClient:
response = await self.db.litellm_spendlogs.find_many( # type: ignore
where={
key_val["key"]: key_val["value"], # type: ignore
}
},
**spend_find_many_kwargs,
)
return response
else:
response = await self.db.litellm_spendlogs.find_many( # type: ignore
order={"startTime": "desc"},
**spend_find_many_kwargs,
)
return response
elif table_name == "budget" and reset_at is not None:

View file

@ -1753,24 +1753,27 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch):
# Return individual log entries when summarize=false
return mock_spend_logs
async def group_by(self, *args, **kwargs):
async def query_raw(self, sql_query, *params):
# Return grouped data when summarize=true
# Simplified mock response for grouped data
# Simplified mock response for aggregated SQL query
assert 'DATE("startTime") AS spend_date' in sql_query
assert 'GROUP BY DATE("startTime"), api_key, "user", model' in sql_query
assert len(params) == 2 # start_date, end_date
yesterday = datetime.datetime.now(timezone.utc) - timedelta(days=1)
return [
{
"spend_date": yesterday.strftime("%Y-%m-%d"),
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-3.5-turbo",
"startTime": yesterday.strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
"_sum": {"spend": 0.05},
"spend": 0.05,
},
{
"spend_date": yesterday.strftime("%Y-%m-%d"),
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-4",
"startTime": yesterday.strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
"_sum": {"spend": 0.10},
"spend": 0.10,
},
]
@ -1859,6 +1862,52 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch):
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_view_spend_logs_summarize_false_applies_safety_limit(client, monkeypatch):
"""Deprecated /spend/logs should clamp limit to SPEND_LOGS_DEPRECATED_MAX_ROWS."""
captured_kwargs = {}
class MockDB:
def __init__(self):
self.litellm_spendlogs = self
async def find_many(self, *args, **kwargs):
captured_kwargs.clear()
captured_kwargs.update(kwargs)
return [{"request_id": "req-1"}]
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
def hash_token(self, token: str) -> str:
return token
monkeypatch.setenv("SPEND_LOGS_DEPRECATED_MAX_ROWS", "5")
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={
"start_date": "2025-01-01",
"end_date": "2025-01-02",
"summarize": "false",
"limit": 100,
"offset": 2,
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert captured_kwargs.get("take") == 5
assert captured_kwargs.get("skip") == 2
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"""
@ -2006,20 +2055,20 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
"""
Tests the /spend/logs endpoint with both start_date and end_date,
ensuring it returns summarized data and not an empty list.
This test specifically validates the fix for dates being passed as ISO strings.
This test validates SQL date aggregation for summarize=true.
"""
from datetime import datetime, timedelta, timezone
# This simulates the summarized data that Prisma's `group_by` would return.
# This simulates the summarized rows returned by query_raw.
mock_summarized_response = [
{
"spend_date": (datetime.now(timezone.utc) - timedelta(days=1)).strftime(
"%Y-%m-%d"
),
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-4",
"startTime": (datetime.now(timezone.utc) - timedelta(days=1)).strftime(
"%Y-%m-%dT%H:%M:%S.%fZ"
),
"_sum": {"spend": 0.15},
"spend": 0.15,
}
]
@ -2028,17 +2077,12 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
def __init__(self):
self.litellm_spendlogs = self
async def group_by(self, *args, **kwargs):
# We assert that the `gte` and `lte` values are strings in ISO format.
# If they were datetime objects, this test would fail.
where_clause = kwargs.get("where", {})
start_time_filter = where_clause.get("startTime", {})
assert "gte" in start_time_filter
assert "lte" in start_time_filter
assert isinstance(start_time_filter["gte"], str)
assert isinstance(start_time_filter["lte"], str)
assert "T" in start_time_filter["gte"] # Check for ISO format 'T' separator
async def query_raw(self, sql_query, *params):
assert 'DATE("startTime") AS spend_date' in sql_query
assert 'GROUP BY DATE("startTime"), api_key, "user", model' in sql_query
assert len(params) == 2
assert isinstance(params[0], datetime)
assert isinstance(params[1], datetime)
# If the assertions pass, return the mock response.
return mock_summarized_response