feat: stream LLM responses (#2677)
Co-authored-by: Michael Neale <michael.neale@gmail.com>
This commit is contained in:
+276
-237
@@ -14,7 +14,7 @@ use crate::agents::sub_recipe_execution_tool::sub_recipe_execute_task_tool::{
|
||||
};
|
||||
use crate::agents::sub_recipe_manager::SubRecipeManager;
|
||||
use crate::config::{Config, ExtensionConfigManager, PermissionManager};
|
||||
use crate::message::Message;
|
||||
use crate::message::{push_message, Message};
|
||||
use crate::permission::permission_judge::check_tool_permissions;
|
||||
use crate::permission::PermissionConfirmation;
|
||||
use crate::providers::base::Provider;
|
||||
@@ -722,6 +722,16 @@ impl Agent {
|
||||
});
|
||||
|
||||
loop {
|
||||
// Check for final output before incrementing turns or checking max_turns
|
||||
// This ensures that if we have a final output ready, we return it immediately
|
||||
// without being blocked by the max_turns limit - this is needed for streaming cases
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_some() {
|
||||
yield AgentEvent::Message(Message::assistant().with_text(final_output_tool.final_output.clone().unwrap()));
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
turns_taken += 1;
|
||||
if turns_taken > max_turns {
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
@@ -752,262 +762,291 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
match Self::generate_response_from_provider(
|
||||
let mut stream = Self::stream_response_from_provider(
|
||||
self.provider().await?,
|
||||
&system_prompt,
|
||||
&messages,
|
||||
&tools,
|
||||
&toolshim_tools,
|
||||
).await {
|
||||
Ok((response, usage)) => {
|
||||
// Emit model change event if provider is lead-worker
|
||||
let provider = self.provider().await?;
|
||||
if let Some(lead_worker) = provider.as_lead_worker() {
|
||||
// The actual model used is in the usage
|
||||
let active_model = usage.model.clone();
|
||||
let (lead_model, worker_model) = lead_worker.get_model_info();
|
||||
let mode = if active_model == lead_model {
|
||||
"lead"
|
||||
} else if active_model == worker_model {
|
||||
"worker"
|
||||
} else {
|
||||
"unknown"
|
||||
};
|
||||
).await?;
|
||||
|
||||
yield AgentEvent::ModelChange {
|
||||
model: active_model,
|
||||
mode: mode.to_string(),
|
||||
};
|
||||
}
|
||||
let mut added_message = false;
|
||||
while let Some(next) = stream.next().await {
|
||||
match next {
|
||||
Ok((response, usage)) => {
|
||||
// Emit model change event if provider is lead-worker
|
||||
let provider = self.provider().await?;
|
||||
if let Some(lead_worker) = provider.as_lead_worker() {
|
||||
if let Some(ref usage) = usage {
|
||||
// The actual model used is in the usage
|
||||
let active_model = usage.model.clone();
|
||||
let (lead_model, worker_model) = lead_worker.get_model_info();
|
||||
let mode = if active_model == lead_model {
|
||||
"lead"
|
||||
} else if active_model == worker_model {
|
||||
"worker"
|
||||
} else {
|
||||
"unknown"
|
||||
};
|
||||
|
||||
// record usage for the session in the session file
|
||||
if let Some(session_config) = session.clone() {
|
||||
Self::update_session_metrics(session_config, &usage, messages.len()).await?;
|
||||
}
|
||||
|
||||
// categorize the type of requests we need to handle
|
||||
let (frontend_requests,
|
||||
remaining_requests,
|
||||
filtered_response) =
|
||||
self.categorize_tool_requests(&response).await;
|
||||
|
||||
// Record tool calls in the router selector
|
||||
let selector = self.router_tool_selector.lock().await.clone();
|
||||
if let Some(selector) = selector {
|
||||
// Record frontend tool calls
|
||||
for request in &frontend_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if let Err(e) = selector.record_tool_call(&tool_call.name).await {
|
||||
tracing::error!("Failed to record frontend tool call: {}", e);
|
||||
}
|
||||
yield AgentEvent::ModelChange {
|
||||
model: active_model,
|
||||
mode: mode.to_string(),
|
||||
};
|
||||
}
|
||||
}
|
||||
// Record remaining tool calls
|
||||
for request in &remaining_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if let Err(e) = selector.record_tool_call(&tool_call.name).await {
|
||||
tracing::error!("Failed to record tool call: {}", e);
|
||||
}
|
||||
|
||||
// record usage for the session in the session file
|
||||
if let Some(session_config) = session.clone() {
|
||||
if let Some(ref usage) = usage {
|
||||
Self::update_session_metrics(session_config, usage, messages.len()).await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
// Yield the assistant's response with frontend tool requests filtered out
|
||||
yield AgentEvent::Message(filtered_response.clone());
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
if let Some(response) = response {
|
||||
// categorize the type of requests we need to handle
|
||||
let (frontend_requests,
|
||||
remaining_requests,
|
||||
filtered_response) =
|
||||
self.categorize_tool_requests(&response).await;
|
||||
|
||||
let num_tool_requests = frontend_requests.len() + remaining_requests.len();
|
||||
if num_tool_requests == 0 {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_none() {
|
||||
tracing::warn!("Final output tool has not been called yet. Continuing agent loop.");
|
||||
let message = Message::assistant().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE);
|
||||
messages.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
// Record tool calls in the router selector
|
||||
let selector = self.router_tool_selector.lock().await.clone();
|
||||
if let Some(selector) = selector {
|
||||
// Record frontend tool calls
|
||||
for request in &frontend_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if let Err(e) = selector.record_tool_call(&tool_call.name).await {
|
||||
tracing::error!("Failed to record frontend tool call: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
// Record remaining tool calls
|
||||
for request in &remaining_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if let Err(e) = selector.record_tool_call(&tool_call.name).await {
|
||||
tracing::error!("Failed to record tool call: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Yield the assistant's response with frontend tool requests filtered out
|
||||
yield AgentEvent::Message(filtered_response.clone());
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
let num_tool_requests = frontend_requests.len() + remaining_requests.len();
|
||||
if num_tool_requests == 0 {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_none() {
|
||||
tracing::warn!("Final output tool has not been called yet. Continuing agent loop.");
|
||||
let message = Message::assistant().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE);
|
||||
messages.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
continue;
|
||||
} else {
|
||||
let message = Message::assistant().with_text(final_output_tool.final_output.clone().unwrap());
|
||||
messages.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
// Set added_message to true and continue to end the current iteration
|
||||
added_message = true;
|
||||
push_message(&mut messages, response);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
// If there's no final output tool and no tool requests, continue the loop
|
||||
continue;
|
||||
} else {
|
||||
let message = Message::assistant().with_text(final_output_tool.final_output.clone().unwrap());
|
||||
messages.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
}
|
||||
|
||||
// Process tool requests depending on frontend tools and then goose_mode
|
||||
let message_tool_response = Arc::new(Mutex::new(Message::user()));
|
||||
|
||||
// First handle any frontend tool requests
|
||||
let mut frontend_tool_stream = self.handle_frontend_tool_requests(
|
||||
&frontend_requests,
|
||||
message_tool_response.clone()
|
||||
);
|
||||
|
||||
// we have a stream of frontend tools to handle, inside the stream
|
||||
// execution is yeield back to this reply loop, and is of the same Message
|
||||
// type, so we can yield that back up to be handled
|
||||
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
// Clone goose_mode once before the match to avoid move issues
|
||||
let mode = goose_mode.clone();
|
||||
if mode.as_str() == "chat" {
|
||||
// Skip all tool calls in chat mode
|
||||
for request in remaining_requests {
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(CHAT_MODE_TOOL_SKIPPED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
} else {
|
||||
// At this point, we have handled the frontend tool requests and know goose_mode != "chat"
|
||||
// What remains is handling the remaining tool requests (enable extension,
|
||||
// regular tool calls) in goose_mode == ["auto", "approve" or "smart_approve"]
|
||||
let mut permission_manager = PermissionManager::default();
|
||||
let (permission_check_result, enable_extension_request_ids) = check_tool_permissions(
|
||||
&remaining_requests,
|
||||
&mode,
|
||||
tools_with_readonly_annotation.clone(),
|
||||
tools_without_annotation.clone(),
|
||||
&mut permission_manager,
|
||||
self.provider().await?).await;
|
||||
|
||||
// Handle pre-approved and read-only tools in parallel
|
||||
let mut tool_futures: Vec<(String, ToolStream)> = Vec::new();
|
||||
|
||||
// Skip the confirmation for approved tools
|
||||
for request in &permission_check_result.approved {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
let (req_id, tool_result) = self.dispatch_tool_call(tool_call, request.id.clone()).await;
|
||||
|
||||
tool_futures.push((req_id, match tool_result {
|
||||
Ok(result) => tool_stream(
|
||||
result.notification_stream.unwrap_or_else(|| Box::new(stream::empty())),
|
||||
result.result,
|
||||
),
|
||||
Err(e) => tool_stream(
|
||||
Box::new(stream::empty()),
|
||||
futures::future::ready(Err(e)),
|
||||
),
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
for request in &permission_check_result.denied {
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(DECLINED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
|
||||
// We need interior mutability in handle_approval_tool_requests
|
||||
let tool_futures_arc = Arc::new(Mutex::new(tool_futures));
|
||||
|
||||
// Process tools requiring approval (enable extension, regular tool calls)
|
||||
let mut tool_approval_stream = self.handle_approval_tool_requests(
|
||||
&permission_check_result.needs_approval,
|
||||
tool_futures_arc.clone(),
|
||||
&mut permission_manager,
|
||||
message_tool_response.clone()
|
||||
);
|
||||
|
||||
// We have a stream of tool_approval_requests to handle
|
||||
// Execution is yielded back to this reply loop, and is of the same Message
|
||||
// type, so we can yield the Message back up to be handled and grab any
|
||||
// confirmations or denials
|
||||
while let Some(msg) = tool_approval_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
tool_futures = {
|
||||
// Lock the mutex asynchronously
|
||||
let mut futures_lock = tool_futures_arc.lock().await;
|
||||
// Drain the vector and collect into a new Vec
|
||||
futures_lock.drain(..).collect::<Vec<_>>()
|
||||
};
|
||||
|
||||
let with_id = tool_futures
|
||||
.into_iter()
|
||||
.map(|(request_id, stream)| {
|
||||
stream.map(move |item| (request_id.clone(), item))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let mut combined = stream::select_all(with_id);
|
||||
|
||||
let mut all_install_successful = true;
|
||||
|
||||
while let Some((request_id, item)) = combined.next().await {
|
||||
match item {
|
||||
ToolStreamItem::Result(output) => {
|
||||
if enable_extension_request_ids.contains(&request_id) && output.is_err(){
|
||||
all_install_successful = false;
|
||||
}
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(request_id, output);
|
||||
},
|
||||
ToolStreamItem::Message(msg) => {
|
||||
yield AgentEvent::McpNotification((request_id, msg))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Update system prompt and tools if installations were successful
|
||||
if all_install_successful {
|
||||
(tools, toolshim_tools, system_prompt) = self.prepare_tools_and_prompt().await?;
|
||||
}
|
||||
}
|
||||
|
||||
let final_message_tool_resp = message_tool_response.lock().await.clone();
|
||||
yield AgentEvent::Message(final_message_tool_resp.clone());
|
||||
|
||||
added_message = true;
|
||||
push_message(&mut messages, response);
|
||||
push_message(&mut messages, final_message_tool_resp);
|
||||
|
||||
// Check for MCP notifications from subagents again before next iteration
|
||||
// Note: These are already handled as McpNotification events above,
|
||||
// so we don't need to convert them to assistant messages here.
|
||||
// This was causing duplicate plain-text notifications.
|
||||
// let mcp_notifications = self.get_mcp_notifications().await;
|
||||
// for notification in mcp_notifications {
|
||||
// // Extract subagent info from the notification data for assistant messages
|
||||
// if let JsonRpcMessage::Notification(ref notif) = notification {
|
||||
// if let Some(params) = ¬if.params {
|
||||
// if let Some(data) = params.get("data") {
|
||||
// if let (Some(subagent_id), Some(message)) = (
|
||||
// data.get("subagent_id").and_then(|v| v.as_str()),
|
||||
// data.get("message").and_then(|v| v.as_str())
|
||||
// ) {
|
||||
// yield AgentEvent::Message(
|
||||
// Message::assistant().with_text(
|
||||
// format!("Subagent {}: {}", subagent_id, message)
|
||||
// )
|
||||
// );
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
}
|
||||
},
|
||||
Err(ProviderError::ContextLengthExceeded(_)) => {
|
||||
// At this point, the last message should be a user message
|
||||
// because call to provider led to context length exceeded error
|
||||
// Immediately yield a special message and break
|
||||
yield AgentEvent::Message(Message::assistant().with_context_length_exceeded(
|
||||
"The context length of the model has been exceeded. Please start a new session and try again.",
|
||||
));
|
||||
break;
|
||||
},
|
||||
Err(e) => {
|
||||
// Create an error message & terminate the stream
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(Message::assistant().with_text(format!("Ran into this error: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")));
|
||||
break;
|
||||
}
|
||||
|
||||
// Process tool requests depending on frontend tools and then goose_mode
|
||||
let message_tool_response = Arc::new(Mutex::new(Message::user()));
|
||||
|
||||
// First handle any frontend tool requests
|
||||
let mut frontend_tool_stream = self.handle_frontend_tool_requests(
|
||||
&frontend_requests,
|
||||
message_tool_response.clone()
|
||||
);
|
||||
|
||||
// we have a stream of frontend tools to handle, inside the stream
|
||||
// execution is yeield back to this reply loop, and is of the same Message
|
||||
// type, so we can yield that back up to be handled
|
||||
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
// Clone goose_mode once before the match to avoid move issues
|
||||
let mode = goose_mode.clone();
|
||||
if mode.as_str() == "chat" {
|
||||
// Skip all tool calls in chat mode
|
||||
for request in remaining_requests {
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(CHAT_MODE_TOOL_SKIPPED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
} else {
|
||||
// At this point, we have handled the frontend tool requests and know goose_mode != "chat"
|
||||
// What remains is handling the remaining tool requests (enable extension,
|
||||
// regular tool calls) in goose_mode == ["auto", "approve" or "smart_approve"]
|
||||
let mut permission_manager = PermissionManager::default();
|
||||
let (permission_check_result, enable_extension_request_ids) = check_tool_permissions(
|
||||
&remaining_requests,
|
||||
&mode,
|
||||
tools_with_readonly_annotation.clone(),
|
||||
tools_without_annotation.clone(),
|
||||
&mut permission_manager,
|
||||
self.provider().await?).await;
|
||||
|
||||
// Handle pre-approved and read-only tools in parallel
|
||||
let mut tool_futures: Vec<(String, ToolStream)> = Vec::new();
|
||||
|
||||
// Skip the confirmation for approved tools
|
||||
for request in &permission_check_result.approved {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
let (req_id, tool_result) = self.dispatch_tool_call(tool_call, request.id.clone()).await;
|
||||
|
||||
tool_futures.push((req_id, match tool_result {
|
||||
Ok(result) => tool_stream(
|
||||
result.notification_stream.unwrap_or_else(|| Box::new(stream::empty())),
|
||||
result.result,
|
||||
),
|
||||
Err(e) => tool_stream(
|
||||
Box::new(stream::empty()),
|
||||
futures::future::ready(Err(e)),
|
||||
),
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
for request in &permission_check_result.denied {
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(DECLINED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
|
||||
// We need interior mutability in handle_approval_tool_requests
|
||||
let tool_futures_arc = Arc::new(Mutex::new(tool_futures));
|
||||
|
||||
// Process tools requiring approval (enable extension, regular tool calls)
|
||||
let mut tool_approval_stream = self.handle_approval_tool_requests(
|
||||
&permission_check_result.needs_approval,
|
||||
tool_futures_arc.clone(),
|
||||
&mut permission_manager,
|
||||
message_tool_response.clone()
|
||||
);
|
||||
|
||||
// We have a stream of tool_approval_requests to handle
|
||||
// Execution is yielded back to this reply loop, and is of the same Message
|
||||
// type, so we can yield the Message back up to be handled and grab any
|
||||
// confirmations or denials
|
||||
while let Some(msg) = tool_approval_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
tool_futures = {
|
||||
// Lock the mutex asynchronously
|
||||
let mut futures_lock = tool_futures_arc.lock().await;
|
||||
// Drain the vector and collect into a new Vec
|
||||
futures_lock.drain(..).collect::<Vec<_>>()
|
||||
};
|
||||
|
||||
let with_id = tool_futures
|
||||
.into_iter()
|
||||
.map(|(request_id, stream)| {
|
||||
stream.map(move |item| (request_id.clone(), item))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let mut combined = stream::select_all(with_id);
|
||||
|
||||
let mut all_install_successful = true;
|
||||
|
||||
while let Some((request_id, item)) = combined.next().await {
|
||||
match item {
|
||||
ToolStreamItem::Result(output) => {
|
||||
if enable_extension_request_ids.contains(&request_id) && output.is_err(){
|
||||
all_install_successful = false;
|
||||
}
|
||||
let mut response = message_tool_response.lock().await;
|
||||
*response = response.clone().with_tool_response(request_id, output);
|
||||
},
|
||||
ToolStreamItem::Message(msg) => {
|
||||
yield AgentEvent::McpNotification((request_id, msg))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Update system prompt and tools if installations were successful
|
||||
if all_install_successful {
|
||||
(tools, toolshim_tools, system_prompt) = self.prepare_tools_and_prompt().await?;
|
||||
}
|
||||
}
|
||||
|
||||
let final_message_tool_resp = message_tool_response.lock().await.clone();
|
||||
yield AgentEvent::Message(final_message_tool_resp.clone());
|
||||
|
||||
messages.push(response);
|
||||
messages.push(final_message_tool_resp);
|
||||
|
||||
// Check for MCP notifications from subagents again before next iteration
|
||||
// Note: These are already handled as McpNotification events above,
|
||||
// so we don't need to convert them to assistant messages here.
|
||||
// This was causing duplicate plain-text notifications.
|
||||
// let mcp_notifications = self.get_mcp_notifications().await;
|
||||
// for notification in mcp_notifications {
|
||||
// // Extract subagent info from the notification data for assistant messages
|
||||
// if let JsonRpcMessage::Notification(ref notif) = notification {
|
||||
// if let Some(params) = ¬if.params {
|
||||
// if let Some(data) = params.get("data") {
|
||||
// if let (Some(subagent_id), Some(message)) = (
|
||||
// data.get("subagent_id").and_then(|v| v.as_str()),
|
||||
// data.get("message").and_then(|v| v.as_str())
|
||||
// ) {
|
||||
// yield AgentEvent::Message(
|
||||
// Message::assistant().with_text(
|
||||
// format!("Subagent {}: {}", subagent_id, message)
|
||||
// )
|
||||
// );
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
},
|
||||
Err(ProviderError::ContextLengthExceeded(_)) => {
|
||||
// At this point, the last message should be a user message
|
||||
// because call to provider led to context length exceeded error
|
||||
// Immediately yield a special message and break
|
||||
yield AgentEvent::Message(Message::assistant().with_context_length_exceeded(
|
||||
"The context length of the model has been exceeded. Please start a new session and try again.",
|
||||
));
|
||||
break;
|
||||
},
|
||||
Err(e) => {
|
||||
// Create an error message & terminate the stream
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(Message::assistant().with_text(format!("Ran into this error: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")));
|
||||
break;
|
||||
}
|
||||
}
|
||||
if !added_message {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_none() {
|
||||
tracing::warn!("Final output tool has not been called yet. Continuing agent loop.");
|
||||
yield AgentEvent::Message(Message::user().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE));
|
||||
continue;
|
||||
} else {
|
||||
yield AgentEvent::Message(Message::assistant().with_text(final_output_tool.final_output.clone().unwrap()));
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
// Yield control back to the scheduler to prevent blocking
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
@@ -2,10 +2,13 @@ use anyhow::Result;
|
||||
use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_stream::try_stream;
|
||||
use futures::stream::StreamExt;
|
||||
|
||||
use crate::agents::router_tool_selector::RouterToolSelectionStrategy;
|
||||
use crate::config::Config;
|
||||
use crate::message::{Message, MessageContent, ToolRequest};
|
||||
use crate::providers::base::{Provider, ProviderUsage};
|
||||
use crate::providers::base::{stream_from_single_message, MessageStream, Provider, ProviderUsage};
|
||||
use crate::providers::errors::ProviderError;
|
||||
use crate::providers::toolshim::{
|
||||
augment_message_with_tool_calls, convert_tool_messages_to_text,
|
||||
@@ -16,6 +19,19 @@ use mcp_core::tool::Tool;
|
||||
|
||||
use super::super::agents::Agent;
|
||||
|
||||
async fn toolshim_postprocess(
|
||||
response: Message,
|
||||
toolshim_tools: &[Tool],
|
||||
) -> Result<Message, ProviderError> {
|
||||
let interpreter = OllamaInterpreter::new().map_err(|e| {
|
||||
ProviderError::ExecutionError(format!("Failed to create OllamaInterpreter: {}", e))
|
||||
})?;
|
||||
|
||||
augment_message_with_tool_calls(&interpreter, response, toolshim_tools)
|
||||
.await
|
||||
.map_err(|e| ProviderError::ExecutionError(format!("Failed to augment message: {}", e)))
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
/// Prepares tools and system prompt for a provider request
|
||||
pub(crate) async fn prepare_tools_and_prompt(
|
||||
@@ -128,25 +144,67 @@ impl Agent {
|
||||
.complete(system_prompt, &messages_for_provider, tools)
|
||||
.await?;
|
||||
|
||||
// Store the model information in the global store
|
||||
crate::providers::base::set_current_model(&usage.model);
|
||||
|
||||
// Post-process / structure the response only if tool interpretation is enabled
|
||||
if config.toolshim {
|
||||
let interpreter = OllamaInterpreter::new().map_err(|e| {
|
||||
ProviderError::ExecutionError(format!("Failed to create OllamaInterpreter: {}", e))
|
||||
})?;
|
||||
|
||||
response = augment_message_with_tool_calls(&interpreter, response, toolshim_tools)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ProviderError::ExecutionError(format!("Failed to augment message: {}", e))
|
||||
})?;
|
||||
response = toolshim_postprocess(response, toolshim_tools).await?;
|
||||
}
|
||||
|
||||
Ok((response, usage))
|
||||
}
|
||||
|
||||
/// Stream a response from the LLM provider.
|
||||
/// Handles toolshim transformations if needed
|
||||
pub(crate) async fn stream_response_from_provider(
|
||||
provider: Arc<dyn Provider>,
|
||||
system_prompt: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
toolshim_tools: &[Tool],
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let config = provider.get_model_config();
|
||||
|
||||
// Convert tool messages to text if toolshim is enabled
|
||||
let messages_for_provider = if config.toolshim {
|
||||
convert_tool_messages_to_text(messages)
|
||||
} else {
|
||||
messages.to_vec()
|
||||
};
|
||||
|
||||
// Clone owned data to move into the async stream
|
||||
let system_prompt = system_prompt.to_owned();
|
||||
let tools = tools.to_owned();
|
||||
let toolshim_tools = toolshim_tools.to_owned();
|
||||
let provider = provider.clone();
|
||||
|
||||
let mut stream = if provider.supports_streaming() {
|
||||
provider
|
||||
.stream(system_prompt.as_str(), &messages_for_provider, &tools)
|
||||
.await?
|
||||
} else {
|
||||
let (message, usage) = provider
|
||||
.complete(system_prompt.as_str(), &messages_for_provider, &tools)
|
||||
.await?;
|
||||
stream_from_single_message(message, usage)
|
||||
};
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
while let Some(Ok((mut message, usage))) = stream.next().await {
|
||||
// Store the model information in the global store
|
||||
if let Some(usage) = usage.as_ref() {
|
||||
crate::providers::base::set_current_model(&usage.model);
|
||||
}
|
||||
|
||||
// Post-process / structure the response only if tool interpretation is enabled
|
||||
if message.is_some() && config.toolshim {
|
||||
message = Some(toolshim_postprocess(message.unwrap(), &toolshim_tools).await?);
|
||||
}
|
||||
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
/// Categorize tool requests from the response into different types
|
||||
/// Returns:
|
||||
/// - frontend_requests: Tool requests that should be handled by the frontend
|
||||
@@ -191,6 +249,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
let filtered_message = Message {
|
||||
id: response.id.clone(),
|
||||
role: response.role.clone(),
|
||||
created: response.created,
|
||||
content: filtered_content,
|
||||
|
||||
Reference in New Issue
Block a user