diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 60d921669ac..d15d23f8eea 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -50,6 +50,8 @@ class RealTimeStreaming: self.messages: List[OpenAIRealtimeEvents] = [] self.input_message: Dict = {} self.input_messages: List[Dict[str, str]] = [] + self.session_tools: List[Dict] = [] + self.tool_calls: List[Dict] = [] _logged_real_time_event_types = litellm.logged_real_time_event_types @@ -86,6 +88,7 @@ class RealTimeStreaming: message_obj = message else: message_obj = json.loads(message) + self._collect_tool_calls_from_response_done(message_obj) try: if ( not isinstance(message, dict) @@ -136,6 +139,9 @@ class RealTimeStreaming: self.input_messages.append( {"role": "system", "content": instructions} ) + tools = session.get("tools") + if tools and isinstance(tools, list): + self.session_tools = tools except (json.JSONDecodeError, AttributeError, TypeError): pass @@ -155,6 +161,29 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass + def _collect_tool_calls_from_response_done( + self, event_obj: dict + ) -> None: + """Extract function_call items from response.done events for spend logging.""" + try: + if event_obj.get("type") != "response.done": + return + response = event_obj.get("response", {}) + for item in response.get("output", []): + if item.get("type") == "function_call": + self.tool_calls.append( + { + "id": item.get("call_id", ""), + "type": "function", + "function": { + "name": item.get("name", ""), + "arguments": item.get("arguments", "{}"), + }, + } + ) + except (AttributeError, TypeError): + pass + def store_input(self, message: Union[str, dict]): """Store input message""" self.input_message = message if isinstance(message, dict) else {} @@ -169,6 +198,13 @@ class RealTimeStreaming: self.logging_obj.model_call_details["messages"] = ( self.input_messages ) + if self.session_tools or self.tool_calls: + self.logging_obj.model_call_details[ + "realtime_tools" + ] = self.session_tools + self.logging_obj.model_call_details[ + "realtime_tool_calls" + ] = self.tool_calls ## ASYNC LOGGING # Create an event loop for the new thread asyncio.create_task(self.logging_obj.async_success_handler(self.messages)) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 766837ebbb5..bbffe33a089 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -733,7 +733,13 @@ def _get_proxy_server_request_for_spend_logs_payload( ) if _proxy_server_request is not None: _request_body = _proxy_server_request.get("body", {}) or {} - + + if kwargs is not None: + realtime_tools = kwargs.get("realtime_tools") + if realtime_tools: + _request_body = dict(_request_body) + _request_body["tools"] = realtime_tools + # Apply message redaction if turn_off_message_logging is enabled if kwargs is not None: from litellm.litellm_core_utils.redact_messages import ( @@ -795,7 +801,13 @@ def _get_response_for_spend_logs_payload( response_obj: Any = payload.get("response") if response_obj is None: return "{}" - + + if kwargs is not None: + realtime_tool_calls = kwargs.get("realtime_tool_calls") + if realtime_tool_calls and isinstance(response_obj, dict): + response_obj = dict(response_obj) + response_obj["tool_calls"] = realtime_tool_calls + # Apply message redaction if turn_off_message_logging is enabled if kwargs is not None: from litellm.litellm_core_utils.redact_messages import ( diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index d32803b5701..aaaab95ce6f 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -238,6 +238,128 @@ async def test_transcription_captured_in_backend_to_client(): assert logging_obj.model_call_details["messages"] == streaming.input_messages +def test_collect_session_tools_from_session_update(): + """ + Test that tools from session.update events are collected. + """ + websocket = MagicMock() + backend_ws = MagicMock() + logging_obj = MagicMock() + streaming = RealTimeStreaming(websocket, backend_ws, logging_obj) + + msg = json.dumps({ + "type": "session.update", + "session": { + "tools": [ + { + "type": "function", + "name": "get_weather", + "description": "Get the current weather", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + } + ], + "instructions": "You are a weather assistant." + } + }) + streaming.store_input(msg) + + assert len(streaming.session_tools) == 1 + assert streaming.session_tools[0]["name"] == "get_weather" + assert len(streaming.input_messages) == 1 + assert streaming.input_messages[0]["role"] == "system" + + +def test_collect_tool_calls_from_response_done(): + """ + Test that function_call items are extracted from response.done events. + """ + websocket = MagicMock() + backend_ws = MagicMock() + logging_obj = MagicMock() + streaming = RealTimeStreaming(websocket, backend_ws, logging_obj) + streaming.logged_real_time_event_types = "*" + + response_done = json.dumps({ + "type": "response.done", + "event_id": "evt_123", + "response": { + "output": [ + { + "type": "function_call", + "call_id": "call_abc123", + "name": "get_weather", + "arguments": '{"location": "Paris"}', + } + ] + } + }) + streaming.store_message(response_done) + + assert len(streaming.tool_calls) == 1 + assert streaming.tool_calls[0]["id"] == "call_abc123" + assert streaming.tool_calls[0]["type"] == "function" + assert streaming.tool_calls[0]["function"]["name"] == "get_weather" + assert streaming.tool_calls[0]["function"]["arguments"] == '{"location": "Paris"}' + + +def test_tool_calls_not_collected_from_non_function_call_output(): + """ + Test that non-function_call output items in response.done are not collected. + """ + websocket = MagicMock() + backend_ws = MagicMock() + logging_obj = MagicMock() + streaming = RealTimeStreaming(websocket, backend_ws, logging_obj) + streaming.logged_real_time_event_types = "*" + + response_done = json.dumps({ + "type": "response.done", + "event_id": "evt_456", + "response": { + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello!"}] + } + ] + } + }) + streaming.store_message(response_done) + + assert len(streaming.tool_calls) == 0 + + +@pytest.mark.asyncio +async def test_log_messages_includes_tools_in_model_call_details(): + """ + Test that log_messages() sets session_tools and tool_calls on the logging object. + """ + 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.session_tools = [ + {"type": "function", "name": "get_weather", "description": "Get weather"} + ] + streaming.tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{"location": "Paris"}'}} + ] + + await streaming.log_messages() + + assert logging_obj.model_call_details["realtime_tools"] == streaming.session_tools + assert logging_obj.model_call_details["realtime_tool_calls"] == streaming.tool_calls + + @pytest.mark.asyncio async def test_realtime_guardrail_blocks_prompt_injection(): """ diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts b/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts index 75f975e9a13..b3c118ef368 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts +++ b/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts @@ -258,6 +258,83 @@ describe("ToolsSection utils", () => { expect(result[1].index).toBe(2); expect(result[2].index).toBe(3); }); + + it("should parse tool calls from realtime API response with tool_calls field", () => { + const log: Partial = { + request_id: "test-realtime-1", + proxy_server_request: { + tools: [ + { + type: "function", + name: "get_weather", + description: "Get current weather", + }, + ], + }, + response: { + results: [], + usage: {}, + tool_calls: [ + { + id: "call_abc", + type: "function", + function: { + name: "get_weather", + arguments: '{"location": "Paris"}', + }, + }, + ], + }, + } as any; + + const result = parseToolsFromLog(log as LogEntry); + + expect(result).toHaveLength(1); + expect(result[0].name).toBe("get_weather"); + expect(result[0].called).toBe(true); + expect(result[0].callData?.arguments).toEqual({ location: "Paris" }); + }); + + it("should parse tool calls from realtime response.done events", () => { + const log: Partial = { + request_id: "test-realtime-2", + proxy_server_request: { + tools: [ + { + type: "function", + name: "get_weather", + description: "Get current weather", + }, + ], + }, + response: { + results: [ + { type: "session.created", session: {} }, + { + type: "response.done", + response: { + output: [ + { + type: "function_call", + call_id: "call_xyz", + name: "get_weather", + arguments: '{"location": "Tokyo"}', + }, + ], + }, + }, + ], + usage: {}, + }, + } as any; + + const result = parseToolsFromLog(log as LogEntry); + + expect(result).toHaveLength(1); + expect(result[0].name).toBe("get_weather"); + expect(result[0].called).toBe(true); + expect(result[0].callData?.arguments).toEqual({ location: "Tokyo" }); + }); }); describe("hasTools", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts b/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts index 351bbf51169..d0fd662efec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts +++ b/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts @@ -77,6 +77,33 @@ function extractToolCallsFromResponse(log: LogEntry): ToolCall[] { } } + // Realtime API format: response.tool_calls (added by spend tracking for realtime calls) + if (Array.isArray(responseData.tool_calls)) { + return responseData.tool_calls; + } + + // Realtime API format: response.results[].response.output[].type === "function_call" + if (Array.isArray(responseData.results)) { + const toolCalls: ToolCall[] = []; + for (const result of responseData.results) { + if (result.type === "response.done" && result.response?.output) { + for (const item of result.response.output) { + if (item.type === "function_call") { + toolCalls.push({ + id: item.call_id || "", + type: "function", + function: { + name: item.name || "", + arguments: item.arguments || "{}", + }, + }); + } + } + } + } + if (toolCalls.length > 0) return toolCalls; + } + return []; }