Implement SQLite-backed DB for session tests

Added a SQLite-backed database class for testing and a new test to ensure session isolation between different API keys.
This commit is contained in:
Deepanshu Pal 2026-10-02 21:42:36 +05:30 • committed by GitHub
parent 174cba307b
commit 1c9cffe39d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -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"]