Fix: store realtime API user input and response content in spend logs

- Collect user text input (conversation.item.create) and voice transcriptions
  (conversation.item.input_audio_transcription.completed) during WebSocket sessions
- Pass collected input_messages to logging object's model_call_details before logging
- Update _get_messages_for_spend_logs_payload to return messages for realtime calls
  when store_prompts_in_spend_logs is enabled
- Capture transcriptions in both raw and provider_config backend message handlers

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

View file

@ -49,6 +49,7 @@ class RealTimeStreaming:
self.logging_obj = logging_obj
self.messages: List[OpenAIRealtimeEvents] = []
self.input_message: Dict = {}
self.input_messages: List[Dict[str, str]] = []
_logged_real_time_event_types = litellm.logged_real_time_event_types
@ -100,15 +101,74 @@ class RealTimeStreaming:
if self._should_store_message(message_obj):
self.messages.append(message_obj)
def store_input(self, message: dict):
def _collect_user_input_from_client_event(
self, message: Union[str, dict]
) -> None:
"""Extract user text content from client WebSocket events for spend logging."""
try:
if isinstance(message, str):
msg_obj = json.loads(message)
elif isinstance(message, dict):
msg_obj = message
else:
return
msg_type = msg_obj.get("type", "")
if msg_type == "conversation.item.create":
item = msg_obj.get("item", {})
if item.get("role") == "user":
content_list = item.get("content", [])
for content in content_list:
if (
isinstance(content, dict)
and content.get("type") == "input_text"
):
text = content.get("text", "")
if text:
self.input_messages.append(
{"role": "user", "content": text}
)
elif msg_type == "session.update":
session = msg_obj.get("session", {})
instructions = session.get("instructions", "")
if instructions:
self.input_messages.append(
{"role": "system", "content": instructions}
)
except (json.JSONDecodeError, AttributeError, TypeError):
pass
def _collect_user_input_from_backend_event(self, event_obj: dict) -> None:
"""Extract user voice transcription from backend events for spend logging."""
try:
event_type = event_obj.get("type", "")
if (
event_type
== "conversation.item.input_audio_transcription.completed"
):
transcript = event_obj.get("transcript", "")
if transcript:
self.input_messages.append(
{"role": "user", "content": transcript}
)
except (AttributeError, TypeError):
pass
def store_input(self, message: Union[str, dict]):
"""Store input message"""
self.input_message = message
self.input_message = message if isinstance(message, dict) else {}
self._collect_user_input_from_client_event(message)
if self.logging_obj:
self.logging_obj.pre_call(input=message, api_key="")
async def log_messages(self):
"""Log messages in list"""
if self.logging_obj:
if self.input_messages:
self.logging_obj.model_call_details["messages"] = (
self.input_messages
)
## ASYNC LOGGING
# Create an event loop for the new thread
asyncio.create_task(self.logging_obj.async_success_handler(self.messages))
@ -259,6 +319,7 @@ class RealTimeStreaming:
== "conversation.item.input_audio_transcription.completed"
):
transcript = event.get("transcript", "")
self._collect_user_input_from_backend_event(event)
self.store_message(event_str)
await self.websocket.send_text(event_str)
blocked = await self.run_realtime_guardrails(
@ -308,6 +369,7 @@ class RealTimeStreaming:
== "conversation.item.input_audio_transcription.completed"
):
transcript = event_obj.get("transcript", "")
self._collect_user_input_from_backend_event(event_obj)
## LOGGING — must happen before continue below
self.store_message(raw_response)
# Forward transcript to client so user sees what they said

View file

@ -582,6 +582,16 @@ def _get_messages_for_spend_logs_payload(
standard_logging_payload: Optional[StandardLoggingPayload],
metadata: Optional[dict] = None,
) -> str:
if _should_store_prompts_and_responses_in_spend_logs():
if standard_logging_payload is not None:
call_type = standard_logging_payload.get("call_type", "")
if call_type == "_arealtime":
messages = standard_logging_payload.get("messages")
if messages is not None:
try:
return json.dumps(messages, default=str)
except Exception:
return "{}"
return "{}"