mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Mitigate deprecated spend logs Prisma OOM path
Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
9806e21871
commit
8529fc0b40
3 changed files with 205 additions and 56 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue