mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Fix: populate tools section in logs UI for realtime API tool calls
Backend: - Collect tool definitions from session.update events in RealTimeStreaming - Extract function_call items from response.done events as tool_calls - Enrich proxy_server_request with tools so UI extractToolsFromRequest finds them - Add tool_calls to response field so UI extractToolCallsFromResponse finds them UI: - Update extractToolCallsFromResponse to handle realtime response formats: - response.tool_calls (enriched by backend) - response.results[].response.output[].type === 'function_call' (raw format) Tests: - 4 new backend tests for tool collection and logging - 2 new UI tests for realtime tool call extraction Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
9fbaddf72d
commit
b5b96ad701
5 changed files with 276 additions and 2 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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<LogEntry> = {
|
||||
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<LogEntry> = {
|
||||
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", () => {
|
||||
|
|
|
|||
|
|
@ -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 [];
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue