Add unit tests for realtime spend log message storage

- Test user text input collection from conversation.item.create events
- Test system instructions collection from session.update events
- Test voice transcription collection from backend transcription events
- Test irrelevant events are ignored
- Test empty transcripts are not collected
- Test log_messages sets input_messages on logging object
- Test transcription captured during backend WebSocket session
- Test _get_messages_for_spend_logs_payload for realtime call types

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-25 03:48:42 +00:00
parent 30e7fa6e6e
commit f7b25d21f3
2 changed files with 242 additions and 0 deletions

View file

@ -66,6 +66,178 @@ def test_realtime_streaming_store_message():
assert len(streaming.messages) == 2 # Should not store the new message
def test_collect_user_input_from_text_conversation_item():
"""
Test that conversation.item.create with input_text content is collected as user input.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
msg = json.dumps({
"type": "conversation.item.create",
"item": {
"role": "user",
"content": [
{"type": "input_text", "text": "Hello, how are you?"}
]
}
})
streaming.store_input(msg)
assert len(streaming.input_messages) == 1
assert streaming.input_messages[0]["role"] == "user"
assert streaming.input_messages[0]["content"] == "Hello, how are you?"
def test_collect_user_input_from_session_update_instructions():
"""
Test that session.update with instructions is collected as system input.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
msg = json.dumps({
"type": "session.update",
"session": {
"instructions": "You are a helpful assistant."
}
})
streaming.store_input(msg)
assert len(streaming.input_messages) == 1
assert streaming.input_messages[0]["role"] == "system"
assert streaming.input_messages[0]["content"] == "You are a helpful assistant."
def test_collect_user_input_from_transcription_event():
"""
Test that conversation.item.input_audio_transcription.completed events
are collected as user input from backend events.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
event_obj = {
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "What is the weather today?",
"item_id": "item_123",
}
streaming._collect_user_input_from_backend_event(event_obj)
assert len(streaming.input_messages) == 1
assert streaming.input_messages[0]["role"] == "user"
assert streaming.input_messages[0]["content"] == "What is the weather today?"
def test_collect_user_input_ignores_irrelevant_events():
"""
Test that irrelevant client events don't get collected as user input.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
# input_audio_buffer.append should not be collected
msg = json.dumps({"type": "input_audio_buffer.append", "audio": "base64data"})
streaming.store_input(msg)
assert len(streaming.input_messages) == 0
# response.create should not be collected
msg = json.dumps({"type": "response.create"})
streaming.store_input(msg)
assert len(streaming.input_messages) == 0
def test_collect_user_input_empty_transcript_not_collected():
"""
Test that transcription events with empty transcripts are not collected.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
event_obj = {
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "",
"item_id": "item_123",
}
streaming._collect_user_input_from_backend_event(event_obj)
assert len(streaming.input_messages) == 0
@pytest.mark.asyncio
async def test_log_messages_sets_input_messages_on_logging_obj():
"""
Test that log_messages() sets input_messages on the logging object's model_call_details.
"""
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
logging_obj.model_call_details = {"messages": "default-message-value"}
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(websocket, backend_ws, logging_obj)
streaming.input_messages = [
{"role": "user", "content": "Hello from voice"},
{"role": "user", "content": "Tell me a joke"},
]
await streaming.log_messages()
assert logging_obj.model_call_details["messages"] == [
{"role": "user", "content": "Hello from voice"},
{"role": "user", "content": "Tell me a joke"},
]
@pytest.mark.asyncio
async def test_transcription_captured_in_backend_to_client():
"""
Test that conversation.item.input_audio_transcription.completed events
from the backend are captured as user input during the WebSocket session.
"""
import litellm
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
transcript_event = json.dumps({
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "What are the opening hours?",
"item_id": "item_789",
}).encode()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
transcript_event,
ConnectionClosed(None, None),
]
)
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.model_call_details = {"messages": "default-message-value"}
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
assert len(streaming.input_messages) == 1
assert streaming.input_messages[0]["role"] == "user"
assert streaming.input_messages[0]["content"] == "What are the opening hours?"
assert logging_obj.model_call_details["messages"] == streaming.input_messages
@pytest.mark.asyncio
async def test_realtime_guardrail_blocks_prompt_injection():
"""

View file

@ -19,6 +19,7 @@ import litellm
from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.spend_tracking.spend_tracking_utils import (
_get_messages_for_spend_logs_payload,
_get_proxy_server_request_for_spend_logs_payload,
_get_response_for_spend_logs_payload,
_get_spend_logs_metadata,
@ -245,6 +246,75 @@ def test_get_vector_store_request_for_spend_logs_payload_null_input(mock_should_
assert result is None
@patch(
"litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs"
)
def test_get_messages_for_spend_logs_realtime_returns_messages(mock_should_store):
"""
Test that _get_messages_for_spend_logs_payload returns messages
for realtime calls when store_prompts_in_spend_logs is True.
"""
mock_should_store.return_value = True
realtime_messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "What is the weather today?"},
]
payload = cast(
StandardLoggingPayload,
{
"call_type": "_arealtime",
"messages": realtime_messages,
},
)
result = _get_messages_for_spend_logs_payload(payload)
parsed = json.loads(result)
assert len(parsed) == 2
assert parsed[0]["role"] == "system"
assert parsed[0]["content"] == "You are a helpful assistant."
assert parsed[1]["role"] == "user"
assert parsed[1]["content"] == "What is the weather today?"
@patch(
"litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs"
)
def test_get_messages_for_spend_logs_realtime_empty_when_disabled(mock_should_store):
"""
Test that _get_messages_for_spend_logs_payload returns '{}' for realtime calls
when store_prompts_in_spend_logs is False.
"""
mock_should_store.return_value = False
payload = cast(
StandardLoggingPayload,
{
"call_type": "_arealtime",
"messages": [{"role": "user", "content": "Hello"}],
},
)
result = _get_messages_for_spend_logs_payload(payload)
assert result == "{}"
@patch(
"litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs"
)
def test_get_messages_for_spend_logs_non_realtime_returns_empty(mock_should_store):
"""
Test that _get_messages_for_spend_logs_payload returns '{}' for non-realtime
calls even when store_prompts_in_spend_logs is True.
"""
mock_should_store.return_value = True
payload = cast(
StandardLoggingPayload,
{
"call_type": "acompletion",
"messages": [{"role": "user", "content": "Hello"}],
},
)
result = _get_messages_for_spend_logs_payload(payload)
assert result == "{}"
@patch(
"litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs"
)