This commit is contained in:
Deepanshu Pal 2026-10-04 18:29:59 -04:00 • committed by GitHub
commit a90ee8e412
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 74 additions and 7 deletions

View file

@ -896,9 +896,8 @@ def _get_session_id_for_spend_log(
omit_when_missing: bool,
batch_trace_session_id: str | None = None,
) -> str | None:
"""Under `omit` only `metadata.session_id`, the key Langfuse reads, counts as a session; `litellm_session_id` may
be a copied trace id. Batch call types carry a deterministic session derived from the batch id, which outranks
the per-request trace ids because those differ between the create call and the cost poller's row."""
"""Resolve the session id for the spend log row: `omit` honors only metadata.session_id,
batch sessions outrank everything, then the caller's litellm_session_id, then trace ids."""
if omit_when_missing:
session_id: Final = metadata.get("session_id") if metadata else None
return str(session_id) if session_id else None
@ -907,6 +906,9 @@ def _get_session_id_for_spend_log(
if batch_trace_session_id is not None:
return batch_trace_session_id
caller_session_id: Final = kwargs.get("litellm_session_id")
if caller_session_id:
return str(caller_session_id)
if standard_logging_payload is not None and standard_logging_payload.get("trace_id") is not None:
return str(standard_logging_payload.get("trace_id"))
if kwargs.get("litellm_trace_id") is not None:

View file

@ -272,7 +272,10 @@ class ResponsesSessionHandler:
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, api_key FROM spend_logs WHERE response_id = previous_response_id, SELECT * FROM spend_logs WHERE session_id = session_id AND api_key = api_key
Only rows written under the same API key as the referenced response are returned: the session id can be
caller supplied, so matching on it alone would let a caller replay another key's history.
A just-finished turn gets a short second chance: the worker that served it may
still be writing its spend log when the follow-up arrives, and an empty result
@ -294,13 +297,18 @@ class ResponsesSessionHandler:
query: Final = """
WITH matching_session AS (
SELECT session_id
SELECT session_id, api_key
FROM "LiteLLM_SpendLogs"
WHERE request_id = $1
)
SELECT *
FROM "LiteLLM_SpendLogs"
WHERE session_id IN (SELECT session_id FROM matching_session)
FROM "LiteLLM_SpendLogs" AS logs
WHERE EXISTS (
SELECT 1
FROM matching_session
WHERE matching_session.session_id = logs.session_id
AND matching_session.api_key = logs.api_key
)
ORDER BY "endTime" ASC;
"""

View file

@ -237,6 +237,27 @@ def test_batch_session_outranks_the_per_request_trace_id():
assert session_id == "batch-uid-1"
def test_caller_litellm_session_id_wins_over_the_per_request_trace_id():
session_id: Final = _get_session_id_for_spend_log(
kwargs={"litellm_session_id": "sess-1", "litellm_trace_id": "trace-abc"},
metadata={"trace_id": "trace-abc"},
standard_logging_payload=_TRACE_ONLY_STANDARD_LOGGING,
omit_when_missing=False,
)
assert session_id == "sess-1"
def test_batch_session_outranks_a_caller_litellm_session_id():
session_id: Final = _get_session_id_for_spend_log(
kwargs={"litellm_session_id": "sess-1", "litellm_trace_id": "trace-abc"},
metadata={"trace_id": "trace-abc"},
standard_logging_payload=_TRACE_ONLY_STANDARD_LOGGING,
omit_when_missing=False,
batch_trace_session_id="batch-uid-1",
)
assert session_id == "batch-uid-1"
def test_omit_policy_still_suppresses_batch_sessions():
session_id: Final = _get_session_id_for_spend_log(
kwargs={},

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