From 8529fc0b40b8c99111fd529651b00a1a4144510a Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 26 Feb 2026 02:02:10 +0000 Subject: [PATCH] Mitigate deprecated spend logs Prisma OOM path Co-authored-by: Ishaan Jaff --- .../spend_management_endpoints.py | 161 ++++++++++++++---- litellm/proxy/utils.py | 10 +- .../test_spend_management_endpoints.py | 90 +++++++--- 3 files changed, 205 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 5b58fbe70a0..829f167365b 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f6613b5548f..a02b9fb0de0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: 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 e439dfd693c..808499215d9 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 @@ -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