mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(responses): bound previous_response_id session reconstruction query
This commit is contained in:
parent
4d33964898
commit
6e6a9fd313
3 changed files with 50 additions and 6 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue