mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
30e7fa6e6e
commit
f7b25d21f3
2 changed files with 242 additions and 0 deletions
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue