mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
174cba307b
commit
1c9cffe39d
1 changed files with 36 additions and 0 deletions
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue