diff --git a/tests/unit/responses/litellm_completion_transformation/test_session_handler.py b/tests/unit/responses/litellm_completion_transformation/test_session_handler.py index 002a595a5e7..ace4101719c 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_session_handler.py +++ b/tests/unit/responses/litellm_completion_transformation/test_session_handler.py @@ -1,4 +1,5 @@ import json +import sqlite3 from typing import Final from unittest.mock import AsyncMock, patch @@ -787,3 +788,38 @@ async def test_message_history_replays_real_key_named_tool_payloads() -> None: tool_message: Final = result["messages"][2] assert json.loads(assistant_message["tool_calls"][0]["function"]["arguments"]) == function_arguments assert json.loads(tool_message["content"]) == function_output + + +class _SqliteBackedPrismaDB: + def __init__(self, rows: list[tuple[str, str, str, str]]): + self._connection = sqlite3.connect(":memory:") + self._connection.row_factory = sqlite3.Row + self._connection.execute( + 'CREATE TABLE "LiteLLM_SpendLogs" (request_id TEXT, api_key TEXT, session_id TEXT, "endTime" TEXT)' + ) + self._connection.executemany('INSERT INTO "LiteLLM_SpendLogs" VALUES (?, ?, ?, ?)', rows) + + async def query_raw(self, query, *args): + return [dict(row) for row in self._connection.execute(query.replace("$1", "?"), args)] + + +@pytest.mark.asyncio +async def test_session_lookup_does_not_return_another_keys_rows_for_a_shared_session_id(): + """A caller-supplied litellm_session_id can collide with another key's session; history must stay per key.""" + sqlite_db: Final = _SqliteBackedPrismaDB( + [ + ("victim-1", "victim-key", "shared-session", "2026-01-01T00:00:01"), + ("victim-2", "victim-key", "shared-session", "2026-01-01T00:00:02"), + ("attacker-1", "attacker-key", "shared-session", "2026-01-01T00:00:03"), + ("attacker-2", "attacker-key", "other-session", "2026-01-01T00:00:04"), + ] + ) + fake_prisma_client = _FakePrismaClient(results=[]) + fake_prisma_client.db = sqlite_db + + with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma_client): + spend_logs = await ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id("attacker-1") + victim_logs = await ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id("victim-1") + + assert [row["request_id"] for row in spend_logs] == ["attacker-1"] + assert [row["request_id"] for row in victim_logs] == ["victim-1", "victim-2"]