[router][grpc] Replace Vec<(String, String, String)> with ExtractedToolCall (#16598)

This commit is contained in:
Chang Su
2026-01-06 16:32:59 -08:00
committed by GitHub
parent 4f443f445a
commit 05b54b6d7b
3 changed files with 73 additions and 59 deletions
@@ -114,10 +114,18 @@ pub(super) fn prepare_chat_tools_and_choice(
};
}
/// Tool call extracted from a ChatCompletionResponse
#[derive(Debug, Clone)]
pub(super) struct ExtractedToolCall {
pub call_id: String,
pub name: String,
pub arguments: String,
}
/// Extract all tool calls from chat response (for parallel tool call support)
pub(super) fn extract_all_tool_calls_from_chat(
response: &crate::protocols::chat::ChatCompletionResponse,
) -> Vec<(String, String, String)> {
) -> Vec<ExtractedToolCall> {
// Check if response has choices with tool calls
let Some(choice) = response.choices.first() else {
return Vec::new();
@@ -128,16 +136,14 @@ pub(super) fn extract_all_tool_calls_from_chat(
if let Some(tool_calls) = &message.tool_calls {
tool_calls
.iter()
.map(|tool_call| {
(
tool_call.id.clone(),
tool_call.function.name.clone(),
tool_call
.function
.arguments
.clone()
.unwrap_or_else(|| "{}".to_string()),
)
.map(|tool_call| ExtractedToolCall {
call_id: tool_call.id.clone(),
name: tool_call.function.name.clone(),
arguments: tool_call
.function
.arguments
.clone()
.unwrap_or_else(|| "{}".to_string()),
})
.collect()
} else {
@@ -15,7 +15,7 @@ use super::{
common::{
build_mcp_list_tools_item, build_next_request, convert_mcp_tools_to_chat_tools,
extract_all_tool_calls_from_chat, load_conversation_history, prepare_chat_tools_and_choice,
ToolLoopState,
ExtractedToolCall, ToolLoopState,
},
conversions,
};
@@ -232,9 +232,10 @@ pub(super) async fn execute_tool_loop(
// Separate MCP and function tool calls
let mcp_tool_names: std::collections::HashSet<&str> =
mcp_tools.iter().map(|t| t.name.as_ref()).collect();
let (mcp_tool_calls, function_tool_calls): (Vec<_>, Vec<_>) = tool_calls
.into_iter()
.partition(|(_, tool_name, _)| mcp_tool_names.contains(tool_name.as_str()));
let (mcp_tool_calls, function_tool_calls): (Vec<ExtractedToolCall>, Vec<_>) =
tool_calls
.into_iter()
.partition(|tc| mcp_tool_names.contains(tc.name.as_str()));
trace!(
"Separated tool calls: {} MCP, {} function",
@@ -312,18 +313,18 @@ pub(super) async fn execute_tool_loop(
}
// Execute all MCP tools
for (call_id, tool_name, args_json_str) in mcp_tool_calls {
for tool_call in mcp_tool_calls {
trace!(
"Calling MCP tool '{}' (call_id: {}) with args: {}",
tool_name,
call_id,
args_json_str
tool_call.name,
tool_call.call_id,
tool_call.arguments
);
let tool_start = Instant::now();
let (output_str, success, error) = match ctx
.mcp_manager
.call_tool(tool_name.as_str(), args_json_str.as_str())
.call_tool(tool_call.name.as_str(), tool_call.arguments.as_str())
.await
{
Ok(result) => match serde_json::to_string(&result) {
@@ -348,12 +349,12 @@ pub(super) async fn execute_tool_loop(
// Record MCP tool metrics
Metrics::record_mcp_tool_duration(
&current_request.model,
&tool_name,
&tool_call.name,
tool_duration,
);
Metrics::record_mcp_tool_call(
&current_request.model,
&tool_name,
&tool_call.name,
if success {
metrics_labels::RESULT_SUCCESS
} else {
@@ -363,9 +364,9 @@ pub(super) async fn execute_tool_loop(
// Record the call in state
state.record_call(
call_id,
tool_name,
args_json_str,
tool_call.call_id,
tool_call.name,
tool_call.arguments,
output_str,
success,
error,
@@ -27,7 +27,7 @@ use uuid::Uuid;
use super::{
common::{
build_next_request, convert_mcp_tools_to_chat_tools, extract_all_tool_calls_from_chat,
prepare_chat_tools_and_choice, ToolLoopState,
prepare_chat_tools_and_choice, ExtractedToolCall, ToolLoopState,
},
conversions,
};
@@ -630,9 +630,10 @@ async fn execute_tool_loop_streaming_internal(
// Separate MCP and function tool calls
let mcp_tool_names: std::collections::HashSet<&str> =
mcp_tools.iter().map(|t| t.name.as_ref()).collect();
let (mcp_tool_calls, function_tool_calls): (Vec<_>, Vec<_>) = tool_calls
.into_iter()
.partition(|(_, tool_name, _)| mcp_tool_names.contains(tool_name.as_str()));
let (mcp_tool_calls, function_tool_calls): (Vec<ExtractedToolCall>, Vec<_>) =
tool_calls
.into_iter()
.partition(|tc| mcp_tool_names.contains(tc.name.as_str()));
trace!(
"Separated tool calls: {} MCP, {} function",
@@ -659,15 +660,15 @@ async fn execute_tool_loop_streaming_internal(
}
// Process each MCP tool call
for (call_id, tool_name, args_json_str) in mcp_tool_calls {
for tool_call in mcp_tool_calls {
state.total_calls += 1;
trace!(
"Executing tool call {}/{}: {} (call_id: {})",
state.total_calls,
state.total_calls,
tool_name,
call_id
tool_call.name,
tool_call.call_id
);
// Allocate output_index for this mcp_call item
@@ -678,7 +679,7 @@ async fn execute_tool_loop_streaming_internal(
let item = json!({
"id": item_id,
"type": "mcp_call",
"name": tool_name,
"name": tool_call.name,
"server_label": state.server_label,
"status": "in_progress",
"arguments": ""
@@ -693,25 +694,31 @@ async fn execute_tool_loop_streaming_internal(
emitter.send_event(&event, &tx)?;
// Emit mcp_call_arguments.delta (simulate streaming by sending full arguments)
let event =
emitter.emit_mcp_call_arguments_delta(output_index, &item_id, &args_json_str);
let event = emitter.emit_mcp_call_arguments_delta(
output_index,
&item_id,
&tool_call.arguments,
);
emitter.send_event(&event, &tx)?;
// Emit mcp_call_arguments.done
let event =
emitter.emit_mcp_call_arguments_done(output_index, &item_id, &args_json_str);
let event = emitter.emit_mcp_call_arguments_done(
output_index,
&item_id,
&tool_call.arguments,
);
emitter.send_event(&event, &tx)?;
// Execute the MCP tool - manager handles parsing and type coercion
trace!(
"Calling MCP tool '{}' with args: {}",
tool_name,
args_json_str
tool_call.name,
tool_call.arguments
);
let tool_start = Instant::now();
let (output_str, success, error) = match ctx
.mcp_manager
.call_tool(tool_name.as_str(), args_json_str.as_str())
.call_tool(tool_call.name.as_str(), tool_call.arguments.as_str())
.await
{
Ok(result) => match serde_json::to_string(&result) {
@@ -724,10 +731,10 @@ async fn execute_tool_loop_streaming_internal(
let item_done = json!({
"id": item_id,
"type": "mcp_call",
"name": tool_name,
"name": tool_call.name,
"server_label": state.server_label,
"status": "completed",
"arguments": args_json_str,
"arguments": tool_call.arguments,
"output": output
});
@@ -749,10 +756,10 @@ async fn execute_tool_loop_streaming_internal(
let item_done = json!({
"id": item_id,
"type": "mcp_call",
"name": tool_name,
"name": tool_call.name,
"server_label": state.server_label,
"status": "failed",
"arguments": args_json_str,
"arguments": tool_call.arguments,
"error": &err
});
@@ -776,10 +783,10 @@ async fn execute_tool_loop_streaming_internal(
let item_done = json!({
"id": item_id,
"type": "mcp_call",
"name": tool_name,
"name": tool_call.name,
"server_label": state.server_label,
"status": "failed",
"arguments": args_json_str,
"arguments": tool_call.arguments,
"error": &err_str
});
@@ -795,10 +802,10 @@ async fn execute_tool_loop_streaming_internal(
let tool_duration = tool_start.elapsed();
// Record MCP tool metrics
Metrics::record_mcp_tool_duration(&model, &tool_name, tool_duration);
Metrics::record_mcp_tool_duration(&model, &tool_call.name, tool_duration);
Metrics::record_mcp_tool_call(
&model,
&tool_name,
&tool_call.name,
if success {
metrics_labels::RESULT_SUCCESS
} else {
@@ -808,9 +815,9 @@ async fn execute_tool_loop_streaming_internal(
// Record the call in state
state.record_call(
call_id,
tool_name,
args_json_str,
tool_call.call_id,
tool_call.name,
tool_call.arguments,
output_str,
success,
error,
@@ -825,7 +832,7 @@ async fn execute_tool_loop_streaming_internal(
);
// Emit function_tool_call events for each function tool
for (call_id, tool_name, args_json_str) in function_tool_calls {
for tool_call in function_tool_calls {
// Allocate output_index for this function_tool_call item
let (output_index, item_id) =
emitter.allocate_output_index(OutputItemType::FunctionCall);
@@ -834,8 +841,8 @@ async fn execute_tool_loop_streaming_internal(
let item = json!({
"id": item_id,
"type": "function_tool_call",
"call_id": call_id,
"name": tool_name,
"call_id": tool_call.call_id,
"name": tool_call.name,
"status": "in_progress",
"arguments": ""
});
@@ -848,7 +855,7 @@ async fn execute_tool_loop_streaming_internal(
let event = emitter.emit_function_call_arguments_delta(
output_index,
&item_id,
&args_json_str,
&tool_call.arguments,
);
emitter.send_event(&event, &tx)?;
@@ -856,7 +863,7 @@ async fn execute_tool_loop_streaming_internal(
let event = emitter.emit_function_call_arguments_done(
output_index,
&item_id,
&args_json_str,
&tool_call.arguments,
);
emitter.send_event(&event, &tx)?;
@@ -864,10 +871,10 @@ async fn execute_tool_loop_streaming_internal(
let item_complete = json!({
"id": item_id,
"type": "function_tool_call",
"call_id": call_id,
"name": tool_name,
"call_id": tool_call.call_id,
"name": tool_call.name,
"status": "completed",
"arguments": args_json_str
"arguments": tool_call.arguments
});
// Emit output_item.done