mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 1534adeb3a into 635085ac14
This commit is contained in:
commit
a90ee8e412
4 changed files with 74 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
|
|
|
|||
|
|
@ -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