diff --git a/litellm/constants.py b/litellm/constants.py index e104c937a9b..891d3efa7bd 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -6,6 +6,7 @@ from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_non DEFAULT_HEALTH_CHECK_PROMPT = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) AZURE_DEFAULT_RESPONSES_API_VERSION = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) +MAX_SPEND_LOGS_PER_RESPONSES_SESSION = get_env_int("MAX_SPEND_LOGS_PER_RESPONSES_SESSION", 1000) ROUTER_MAX_FALLBACKS = int(os.getenv("ROUTER_MAX_FALLBACKS", 5)) DEFAULT_BATCH_SIZE = int(os.getenv("DEFAULT_BATCH_SIZE", 512)) DEFAULT_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5)) diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index 68637ea97b3..7913d356fe3 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -252,13 +252,16 @@ class ResponsesSessionHandler: previous_response_id: str, ) -> List[SpendLogsPayload]: """ - Get all spend logs for a previous response id - + Get the spend logs for a previous response id, bounded to the most recent + ``MAX_SPEND_LOGS_PER_RESPONSES_SESSION`` rows so a single ``/v1/responses`` request + cannot make the query engine buffer an unbounded amount of ``LiteLLM_SpendLogs`` data. SQL query - SELECT session_id FROM spend_logs WHERE response_id = previous_response_id, SELECT * FROM spend_logs WHERE session_id = session_id + SELECT session_id FROM spend_logs WHERE request_id = previous_response_id, then the + most recent N rows for that session, returned in ascending endTime order. """ + from litellm.constants import MAX_SPEND_LOGS_PER_RESPONSES_SESSION from litellm.proxy.proxy_server import prisma_client verbose_proxy_logger.debug("decoding response id=%s", previous_response_id) @@ -273,14 +276,20 @@ class ResponsesSessionHandler: SELECT session_id FROM "LiteLLM_SpendLogs" WHERE request_id = $1 + ), + recent_logs AS ( + SELECT * + FROM "LiteLLM_SpendLogs" + WHERE session_id IN (SELECT session_id FROM matching_session) + ORDER BY "endTime" DESC + LIMIT $2 ) SELECT * - FROM "LiteLLM_SpendLogs" - WHERE session_id IN (SELECT session_id FROM matching_session) + FROM recent_logs ORDER BY "endTime" ASC; """ - spend_logs = await prisma_client.db.query_raw(query, previous_response_id) + spend_logs = await prisma_client.db.query_raw(query, previous_response_id, MAX_SPEND_LOGS_PER_RESPONSES_SESSION) verbose_proxy_logger.debug( "Found the following spend logs for previous response id %s: %s", diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py index 9c354101e22..8c9abc1d763 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py @@ -388,6 +388,40 @@ async def test_should_check_cold_storage_for_full_payload(): ), "Should return False when cold storage is not configured, even with truncated content" +@pytest.mark.asyncio +async def test_get_all_spend_logs_for_previous_response_id_is_bounded(monkeypatch): + """ + Regression test for the unbounded SELECT * over LiteLLM_SpendLogs during + Responses API session reconstruction (query-engine OOM). + + The reconstruction query must be capped to the most recent + MAX_SPEND_LOGS_PER_RESPONSES_SESSION rows via a LIMIT, and that bound must be + passed to the query engine so it cannot buffer an entire session into memory. + """ + from litellm.proxy import proxy_server + + monkeypatch.setattr(litellm.constants, "MAX_SPEND_LOGS_PER_RESPONSES_SESSION", 7) + + mock_prisma = AsyncMock() + mock_prisma.db.query_raw = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma) + + await ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id( + "chatcmpl-bounded-test" + ) + + mock_prisma.db.query_raw.assert_awaited_once() + call_args = mock_prisma.db.query_raw.await_args + query = call_args.args[0] + + # The query must bind a LIMIT so the query engine never buffers the whole session + assert "LIMIT $2" in query + assert "ORDER BY \"endTime\" DESC" in query + + # The configured cap must be the value passed to the query engine + assert call_args.args[2] == 7 + + @pytest.mark.asyncio async def test_get_chat_completion_message_history_empty_response_dict(): """