Fix multi tool calling (#5855)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
+329
-306
@@ -299,7 +299,7 @@ impl Agent {
|
||||
async fn handle_approved_and_denied_tools(
|
||||
&self,
|
||||
permission_check_result: &PermissionCheckResult,
|
||||
message_tool_response: Arc<Mutex<Message>>,
|
||||
request_to_response_map: &HashMap<String, Arc<Mutex<Message>>>,
|
||||
cancel_token: Option<tokio_util::sync::CancellationToken>,
|
||||
session: &Session,
|
||||
) -> Result<Vec<(String, ToolStream)>> {
|
||||
@@ -334,18 +334,25 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Handle denied tools
|
||||
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![rmcp::model::Content::text(DECLINED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
|
||||
Self::handle_denied_tools(permission_check_result, request_to_response_map).await;
|
||||
Ok(tool_futures)
|
||||
}
|
||||
|
||||
async fn handle_denied_tools(
|
||||
permission_check_result: &PermissionCheckResult,
|
||||
request_to_response_map: &HashMap<String, Arc<Mutex<Message>>>,
|
||||
) {
|
||||
for request in &permission_check_result.denied {
|
||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
||||
let mut response = response_msg.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![rmcp::model::Content::text(DECLINED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn set_scheduler(&self, scheduler: Arc<dyn SchedulerTrait>) {
|
||||
let mut scheduler_service = self.scheduler_service.lock().await;
|
||||
*scheduler_service = Some(scheduler);
|
||||
@@ -920,336 +927,352 @@ impl Agent {
|
||||
});
|
||||
|
||||
Ok(Box::pin(async_stream::try_stream! {
|
||||
let _ = reply_span.enter();
|
||||
let mut turns_taken = 0u32;
|
||||
let max_turns = session_config.max_turns.unwrap_or(DEFAULT_MAX_TURNS);
|
||||
let _ = reply_span.enter();
|
||||
let mut turns_taken = 0u32;
|
||||
let max_turns = session_config.max_turns.unwrap_or(DEFAULT_MAX_TURNS);
|
||||
|
||||
loop {
|
||||
if is_token_cancelled(&cancel_token) {
|
||||
break;
|
||||
loop {
|
||||
if is_token_cancelled(&cancel_token) {
|
||||
break;
|
||||
}
|
||||
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_some() {
|
||||
let final_event = AgentEvent::Message(
|
||||
Message::assistant().with_text(final_output_tool.final_output.clone().unwrap())
|
||||
);
|
||||
yield final_event;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
turns_taken += 1;
|
||||
if turns_taken > max_turns {
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
"I've reached the maximum number of actions I can do without user input. Would you like me to continue?"
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
let conversation_with_moim = super::moim::inject_moim(
|
||||
conversation.clone(),
|
||||
&self.extension_manager,
|
||||
).await;
|
||||
|
||||
let mut stream = Self::stream_response_from_provider(
|
||||
self.provider().await?,
|
||||
&system_prompt,
|
||||
conversation_with_moim.messages(),
|
||||
&tools,
|
||||
&toolshim_tools,
|
||||
).await?;
|
||||
|
||||
let mut no_tools_called = true;
|
||||
let mut messages_to_add = Conversation::default();
|
||||
let mut tools_updated = false;
|
||||
let mut did_recovery_compact_this_iteration = false;
|
||||
|
||||
while let Some(next) = stream.next().await {
|
||||
if is_token_cancelled(&cancel_token) {
|
||||
break;
|
||||
}
|
||||
|
||||
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 {
|
||||
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"
|
||||
};
|
||||
|
||||
yield AgentEvent::ModelChange {
|
||||
model: active_model,
|
||||
mode: mode.to_string(),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(ref usage) = usage {
|
||||
Self::update_session_metrics(&session_config, usage, false).await?;
|
||||
}
|
||||
|
||||
if let Some(response) = response {
|
||||
let ToolCategorizeResult {
|
||||
frontend_requests,
|
||||
remaining_requests,
|
||||
filtered_response,
|
||||
} = self.categorize_tools(&response, &tools).await;
|
||||
let requests_to_record: Vec<ToolRequest> = frontend_requests.iter().chain(remaining_requests.iter()).cloned().collect();
|
||||
self.tool_route_manager
|
||||
.record_tool_requests(&requests_to_record)
|
||||
.await;
|
||||
|
||||
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 {
|
||||
messages_to_add.push(response.clone());
|
||||
continue;
|
||||
}
|
||||
|
||||
let tool_response_messages: Vec<Arc<Mutex<Message>>> = (0..num_tool_requests)
|
||||
.map(|_| Arc::new(Mutex::new(Message::user().with_id(
|
||||
format!("msg_{}", Uuid::new_v4())
|
||||
))))
|
||||
.collect();
|
||||
|
||||
let mut request_to_response_map = HashMap::new();
|
||||
for (idx, request) in frontend_requests.iter().chain(remaining_requests.iter()).enumerate() {
|
||||
request_to_response_map.insert(request.id.clone(), tool_response_messages[idx].clone());
|
||||
}
|
||||
|
||||
for (idx, request) in frontend_requests.iter().enumerate() {
|
||||
let mut frontend_tool_stream = self.handle_frontend_tool_request(
|
||||
request,
|
||||
tool_response_messages[idx].clone(),
|
||||
);
|
||||
|
||||
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_some() {
|
||||
let final_event = AgentEvent::Message(
|
||||
Message::assistant().with_text(final_output_tool.final_output.clone().unwrap())
|
||||
}
|
||||
if goose_mode == GooseMode::Chat {
|
||||
// Skip all remaining tool calls in chat mode
|
||||
for request in remaining_requests.iter() {
|
||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
||||
let mut response = response_msg.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(CHAT_MODE_TOOL_SKIPPED_RESPONSE)]),
|
||||
);
|
||||
yield final_event;
|
||||
break;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Run all tool inspectors
|
||||
let inspection_results = self.tool_inspection_manager
|
||||
.inspect_tools(
|
||||
&remaining_requests,
|
||||
conversation.messages(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
let permission_check_result = self.tool_inspection_manager
|
||||
.process_inspection_results_with_permission_inspector(
|
||||
&remaining_requests,
|
||||
&inspection_results,
|
||||
)
|
||||
.unwrap_or_else(|| {
|
||||
let mut result = PermissionCheckResult {
|
||||
approved: vec![],
|
||||
needs_approval: vec![],
|
||||
denied: vec![],
|
||||
};
|
||||
result.needs_approval.extend(remaining_requests.iter().cloned());
|
||||
result
|
||||
});
|
||||
|
||||
// Track extension requests
|
||||
let mut enable_extension_request_ids = vec![];
|
||||
for request in &remaining_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if tool_call.name == MANAGE_EXTENSIONS_TOOL_NAME_COMPLETE {
|
||||
enable_extension_request_ids.push(request.id.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
turns_taken += 1;
|
||||
if turns_taken > max_turns {
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
"I've reached the maximum number of actions I can do without user input. Would you like me to continue?"
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
let conversation_with_moim = super::moim::inject_moim(
|
||||
conversation.clone(),
|
||||
&self.extension_manager,
|
||||
).await;
|
||||
|
||||
let mut stream = Self::stream_response_from_provider(
|
||||
self.provider().await?,
|
||||
&system_prompt,
|
||||
conversation_with_moim.messages(),
|
||||
&tools,
|
||||
&toolshim_tools,
|
||||
let mut tool_futures = self.handle_approved_and_denied_tools(
|
||||
&permission_check_result,
|
||||
&request_to_response_map,
|
||||
cancel_token.clone(),
|
||||
&session,
|
||||
).await?;
|
||||
|
||||
let mut no_tools_called = true;
|
||||
let mut messages_to_add = Conversation::default();
|
||||
let mut tools_updated = false;
|
||||
let mut did_recovery_compact_this_iteration = false;
|
||||
let tool_futures_arc = Arc::new(Mutex::new(tool_futures));
|
||||
|
||||
while let Some(next) = stream.next().await {
|
||||
let mut tool_approval_stream = self.handle_approval_tool_requests(
|
||||
&permission_check_result.needs_approval,
|
||||
tool_futures_arc.clone(),
|
||||
&request_to_response_map,
|
||||
cancel_token.clone(),
|
||||
&session,
|
||||
&inspection_results,
|
||||
);
|
||||
|
||||
while let Some(msg) = tool_approval_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
tool_futures = {
|
||||
let mut futures_lock = tool_futures_arc.lock().await;
|
||||
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 {
|
||||
if is_token_cancelled(&cancel_token) {
|
||||
break;
|
||||
}
|
||||
|
||||
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 {
|
||||
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"
|
||||
};
|
||||
|
||||
yield AgentEvent::ModelChange {
|
||||
model: active_model,
|
||||
mode: mode.to_string(),
|
||||
};
|
||||
}
|
||||
match item {
|
||||
ToolStreamItem::Result(output) => {
|
||||
if enable_extension_request_ids.contains(&request_id)
|
||||
&& output.is_err()
|
||||
{
|
||||
all_install_successful = false;
|
||||
}
|
||||
|
||||
if let Some(ref usage) = usage {
|
||||
Self::update_session_metrics(&session_config, usage, false).await?;
|
||||
}
|
||||
|
||||
if let Some(response) = response {
|
||||
messages_to_add.push(response.clone());
|
||||
let ToolCategorizeResult {
|
||||
frontend_requests,
|
||||
remaining_requests,
|
||||
filtered_response,
|
||||
} = self.categorize_tools(&response, &tools).await;
|
||||
let requests_to_record: Vec<ToolRequest> = frontend_requests.iter().chain(remaining_requests.iter()).cloned().collect();
|
||||
self.tool_route_manager
|
||||
.record_tool_requests(&requests_to_record)
|
||||
.await;
|
||||
|
||||
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 {
|
||||
continue;
|
||||
}
|
||||
|
||||
let message_tool_response = Arc::new(Mutex::new(Message::user().with_id(
|
||||
format!("msg_{}", Uuid::new_v4())
|
||||
)));
|
||||
|
||||
let mut frontend_tool_stream = self.handle_frontend_tool_requests(
|
||||
&frontend_requests,
|
||||
message_tool_response.clone(),
|
||||
);
|
||||
|
||||
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
if goose_mode == GooseMode::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 {
|
||||
// Run all tool inspectors (security, repetition, permission, etc.)
|
||||
let inspection_results = self.tool_inspection_manager
|
||||
.inspect_tools(
|
||||
&remaining_requests,
|
||||
conversation.messages(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Process inspection results into permission decisions using the permission inspector
|
||||
let permission_check_result = self.tool_inspection_manager
|
||||
.process_inspection_results_with_permission_inspector(
|
||||
&remaining_requests,
|
||||
&inspection_results,
|
||||
)
|
||||
.unwrap_or_else(|| {
|
||||
// Fallback if permission inspector not found - default to needs approval
|
||||
let mut result = PermissionCheckResult {
|
||||
approved: vec![],
|
||||
needs_approval: vec![],
|
||||
denied: vec![],
|
||||
};
|
||||
result.needs_approval.extend(remaining_requests.iter().cloned());
|
||||
result
|
||||
});
|
||||
|
||||
// Track extension requests for special handling
|
||||
let mut enable_extension_request_ids = vec![];
|
||||
for request in &remaining_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if tool_call.name == MANAGE_EXTENSIONS_TOOL_NAME_COMPLETE {
|
||||
enable_extension_request_ids.push(request.id.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut tool_futures = self.handle_approved_and_denied_tools(
|
||||
&permission_check_result,
|
||||
message_tool_response.clone(),
|
||||
cancel_token.clone(),
|
||||
&session,
|
||||
).await?;
|
||||
|
||||
let tool_futures_arc = Arc::new(Mutex::new(tool_futures));
|
||||
|
||||
let mut tool_approval_stream = self.handle_approval_tool_requests(
|
||||
&permission_check_result.needs_approval,
|
||||
tool_futures_arc.clone(),
|
||||
message_tool_response.clone(),
|
||||
cancel_token.clone(),
|
||||
&session,
|
||||
&inspection_results,
|
||||
);
|
||||
|
||||
while let Some(msg) = tool_approval_stream.try_next().await? {
|
||||
yield AgentEvent::Message(msg);
|
||||
}
|
||||
|
||||
tool_futures = {
|
||||
let mut futures_lock = tool_futures_arc.lock().await;
|
||||
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 {
|
||||
if is_token_cancelled(&cancel_token) {
|
||||
break;
|
||||
}
|
||||
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,
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if all_install_successful && !enable_extension_request_ids.is_empty() {
|
||||
if let Err(e) = self.save_extension_state(&session_config).await {
|
||||
warn!("Failed to save extension state after runtime changes: {}", e);
|
||||
}
|
||||
tools_updated = true;
|
||||
}
|
||||
}
|
||||
|
||||
let final_message_tool_resp = message_tool_response.lock().await.clone();
|
||||
yield AgentEvent::Message(final_message_tool_resp.clone());
|
||||
|
||||
no_tools_called = false;
|
||||
messages_to_add.push(final_message_tool_resp);
|
||||
if let Some(response_msg) = request_to_response_map.get(&request_id) {
|
||||
let mut response = response_msg.lock().await;
|
||||
*response = response.clone().with_tool_response(request_id, output);
|
||||
}
|
||||
}
|
||||
Err(ProviderError::ContextLengthExceeded(_error_msg)) => {
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_system_notification(
|
||||
SystemNotificationType::InlineMessage,
|
||||
"Context limit reached. Compacting to continue conversation...",
|
||||
)
|
||||
);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_system_notification(
|
||||
SystemNotificationType::ThinkingMessage,
|
||||
COMPACTION_THINKING_TEXT,
|
||||
)
|
||||
);
|
||||
ToolStreamItem::Message(msg) => {
|
||||
yield AgentEvent::McpNotification((request_id, msg));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match compact_messages(self.provider().await?.as_ref(), &conversation, false).await {
|
||||
Ok((compacted_conversation, usage)) => {
|
||||
SessionManager::replace_conversation(&session_config.id, &compacted_conversation).await?;
|
||||
Self::update_session_metrics(&session_config, &usage, true).await?;
|
||||
conversation = compacted_conversation;
|
||||
did_recovery_compact_this_iteration = true;
|
||||
yield AgentEvent::HistoryReplaced(conversation.clone());
|
||||
continue;
|
||||
if all_install_successful && !enable_extension_request_ids.is_empty() {
|
||||
if let Err(e) = self.save_extension_state(&session_config).await {
|
||||
warn!("Failed to save extension state after runtime changes: {}", e);
|
||||
}
|
||||
tools_updated = true;
|
||||
}
|
||||
}
|
||||
|
||||
for (idx, request) in frontend_requests.iter()
|
||||
.chain(remaining_requests.iter()).enumerate() {
|
||||
if request.tool_call.is_ok() {
|
||||
let request_msg = Message::assistant()
|
||||
.with_id(format!("msg_{}", Uuid::new_v4()))
|
||||
.with_tool_request(request.id.clone(), request.tool_call.clone());
|
||||
messages_to_add.push(request_msg);
|
||||
let final_response = tool_response_messages[idx]
|
||||
.lock().await.clone();
|
||||
yield AgentEvent::Message(final_response.clone());
|
||||
messages_to_add.push(final_response);
|
||||
}
|
||||
}
|
||||
no_tools_called = false;
|
||||
}
|
||||
}
|
||||
Err(ProviderError::ContextLengthExceeded(_error_msg)) => {
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_system_notification(
|
||||
SystemNotificationType::InlineMessage,
|
||||
"Context limit reached. Compacting to continue conversation...",
|
||||
)
|
||||
);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_system_notification(
|
||||
SystemNotificationType::ThinkingMessage,
|
||||
COMPACTION_THINKING_TEXT,
|
||||
)
|
||||
);
|
||||
|
||||
match compact_messages(self.provider().await?.as_ref(), &conversation, false).await {
|
||||
Ok((compacted_conversation, usage)) => {
|
||||
SessionManager::replace_conversation(&session_config.id, &compacted_conversation).await?;
|
||||
Self::update_session_metrics(&session_config, &usage, true).await?;
|
||||
conversation = compacted_conversation;
|
||||
did_recovery_compact_this_iteration = true;
|
||||
yield AgentEvent::HistoryReplaced(conversation.clone());
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Ran into this error trying to compact: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Ran into this error trying to compact: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||
format!("Ran into this error: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
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 tools_updated {
|
||||
(tools, toolshim_tools, system_prompt) =
|
||||
self.prepare_tools_and_prompt(&working_dir).await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
if tools_updated {
|
||||
(tools, toolshim_tools, system_prompt) =
|
||||
self.prepare_tools_and_prompt(&working_dir).await?;
|
||||
}
|
||||
let mut exit_chat = false;
|
||||
if no_tools_called {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_none() {
|
||||
warn!("Final output tool has not been called yet. Continuing agent loop.");
|
||||
let message = Message::user().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE);
|
||||
messages_to_add.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
} else {
|
||||
let message = Message::assistant().with_text(final_output_tool.final_output.clone().unwrap());
|
||||
messages_to_add.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
exit_chat = true;
|
||||
}
|
||||
} else if did_recovery_compact_this_iteration {
|
||||
// Avoid setting exit_chat; continue from last user message in the conversation
|
||||
} else {
|
||||
match self.handle_retry_logic(&mut conversation, &session_config, &initial_messages).await {
|
||||
Ok(should_retry) => {
|
||||
if should_retry {
|
||||
info!("Retry logic triggered, restarting agent loop");
|
||||
let mut exit_chat = false;
|
||||
if no_tools_called {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_none() {
|
||||
warn!("Final output tool has not been called yet. Continuing agent loop.");
|
||||
let message = Message::user().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE);
|
||||
messages_to_add.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
} else {
|
||||
let message = Message::assistant().with_text(final_output_tool.final_output.clone().unwrap());
|
||||
messages_to_add.push(message.clone());
|
||||
yield AgentEvent::Message(message);
|
||||
exit_chat = true;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Retry logic failed: {}", e);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Retry logic encountered an error: {}", e)
|
||||
)
|
||||
);
|
||||
exit_chat = true;
|
||||
} else if did_recovery_compact_this_iteration {
|
||||
// Avoid setting exit_chat; continue from last user message in the conversation
|
||||
} else {
|
||||
match self.handle_retry_logic(&mut conversation, &session_config, &initial_messages).await {
|
||||
Ok(should_retry) => {
|
||||
if should_retry {
|
||||
info!("Retry logic triggered, restarting agent loop");
|
||||
} else {
|
||||
exit_chat = true;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Retry logic failed: {}", e);
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Retry logic encountered an error: {}", e)
|
||||
)
|
||||
);
|
||||
exit_chat = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for msg in &messages_to_add {
|
||||
SessionManager::add_message(&session_config.id, msg).await?;
|
||||
}
|
||||
conversation.extend(messages_to_add);
|
||||
if exit_chat {
|
||||
break;
|
||||
}
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
}
|
||||
|
||||
for msg in &messages_to_add {
|
||||
SessionManager::add_message(&session_config.id, msg).await?;
|
||||
}
|
||||
conversation.extend(messages_to_add);
|
||||
if exit_chat {
|
||||
break;
|
||||
}
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
}))
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn extend_system_prompt(&self, instruction: String) {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -52,97 +53,98 @@ impl Agent {
|
||||
&'a self,
|
||||
tool_requests: &'a [ToolRequest],
|
||||
tool_futures: Arc<Mutex<Vec<(String, ToolStream)>>>,
|
||||
message_tool_response: Arc<Mutex<Message>>,
|
||||
request_to_response_map: &'a HashMap<String, Arc<Mutex<Message>>>,
|
||||
cancellation_token: Option<CancellationToken>,
|
||||
session: &'a Session,
|
||||
inspection_results: &'a [crate::tool_inspection::InspectionResult],
|
||||
) -> BoxStream<'a, anyhow::Result<Message>> {
|
||||
try_stream! {
|
||||
for request in tool_requests.iter() {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
// Find the corresponding inspection result for this tool request
|
||||
let security_message = inspection_results.iter()
|
||||
.find(|result| result.tool_request_id == request.id)
|
||||
.and_then(|result| {
|
||||
if let crate::tool_inspection::InspectionAction::RequireApproval(Some(message)) = &result.action {
|
||||
Some(message.clone())
|
||||
} else {
|
||||
None
|
||||
for request in tool_requests.iter() {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
// Find the corresponding inspection result for this tool request
|
||||
let security_message = inspection_results.iter()
|
||||
.find(|result| result.tool_request_id == request.id)
|
||||
.and_then(|result| {
|
||||
if let crate::tool_inspection::InspectionAction::RequireApproval(Some(message)) = &result.action {
|
||||
Some(message.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
});
|
||||
|
||||
let confirmation = Message::assistant()
|
||||
.with_tool_confirmation_request(
|
||||
request.id.clone(),
|
||||
tool_call.name.to_string().clone(),
|
||||
tool_call.arguments.clone().unwrap_or_default(),
|
||||
security_message,
|
||||
)
|
||||
.user_only();
|
||||
yield confirmation;
|
||||
|
||||
let mut rx = self.confirmation_rx.lock().await;
|
||||
while let Some((req_id, confirmation)) = rx.recv().await {
|
||||
if req_id == request.id {
|
||||
// Log user decision if this was a security alert
|
||||
if let Some(finding_id) = get_security_finding_id_from_results(&request.id, inspection_results) {
|
||||
tracing::info!(
|
||||
counter.goose.prompt_injection_user_decisions = 1,
|
||||
decision = ?confirmation.permission,
|
||||
finding_id = %finding_id,
|
||||
"User security decision"
|
||||
);
|
||||
}
|
||||
|
||||
if confirmation.permission == Permission::AllowOnce || confirmation.permission == Permission::AlwaysAllow {
|
||||
let (req_id, tool_result) = self.dispatch_tool_call(tool_call.clone(), request.id.clone(), cancellation_token.clone(), session).await;
|
||||
let mut futures = tool_futures.lock().await;
|
||||
|
||||
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)),
|
||||
),
|
||||
}));
|
||||
|
||||
// Update the shared permission manager when user selects "Always Allow"
|
||||
if confirmation.permission == Permission::AlwaysAllow {
|
||||
self.tool_inspection_manager
|
||||
.update_permission_manager(&tool_call.name, PermissionLevel::AlwaysAllow)
|
||||
.await;
|
||||
}
|
||||
});
|
||||
|
||||
let confirmation = Message::assistant()
|
||||
.with_tool_confirmation_request(
|
||||
request.id.clone(),
|
||||
tool_call.name.to_string().clone(),
|
||||
tool_call.arguments.clone().unwrap_or_default(),
|
||||
security_message,
|
||||
)
|
||||
.user_only();
|
||||
yield confirmation;
|
||||
|
||||
let mut rx = self.confirmation_rx.lock().await;
|
||||
while let Some((req_id, confirmation)) = rx.recv().await {
|
||||
if req_id == request.id {
|
||||
// Log user decision if this was a security alert
|
||||
if let Some(finding_id) = get_security_finding_id_from_results(&request.id, inspection_results) {
|
||||
tracing::info!(
|
||||
counter.goose.prompt_injection_user_decisions = 1,
|
||||
decision = ?confirmation.permission,
|
||||
finding_id = %finding_id,
|
||||
"User security decision"
|
||||
);
|
||||
}
|
||||
|
||||
if confirmation.permission == Permission::AllowOnce || confirmation.permission == Permission::AlwaysAllow {
|
||||
let (req_id, tool_result) = self.dispatch_tool_call(tool_call.clone(), request.id.clone(), cancellation_token.clone(), session).await;
|
||||
let mut futures = tool_futures.lock().await;
|
||||
|
||||
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)),
|
||||
),
|
||||
}));
|
||||
|
||||
// Update the shared permission manager when user selects "Always Allow"
|
||||
if confirmation.permission == Permission::AlwaysAllow {
|
||||
self.tool_inspection_manager
|
||||
.update_permission_manager(&tool_call.name, PermissionLevel::AlwaysAllow)
|
||||
.await;
|
||||
}
|
||||
} else {
|
||||
// User declined - add declined response
|
||||
let mut response = message_tool_response.lock().await;
|
||||
} else {
|
||||
// User declined - update the specific response message for this request
|
||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
||||
let mut response = response_msg.lock().await;
|
||||
*response = response.clone().with_tool_response(
|
||||
request.id.clone(),
|
||||
Ok(vec![Content::text(DECLINED_RESPONSE)]),
|
||||
);
|
||||
}
|
||||
break; // Exit the loop once the matching `req_id` is found
|
||||
}
|
||||
break; // Exit the loop once the matching `req_id` is found
|
||||
}
|
||||
}
|
||||
}
|
||||
}.boxed()
|
||||
}
|
||||
}.boxed()
|
||||
}
|
||||
|
||||
pub(crate) fn handle_frontend_tool_requests<'a>(
|
||||
pub(crate) fn handle_frontend_tool_request<'a>(
|
||||
&'a self,
|
||||
tool_requests: &'a [ToolRequest],
|
||||
tool_request: &'a ToolRequest,
|
||||
message_tool_response: Arc<Mutex<Message>>,
|
||||
) -> BoxStream<'a, anyhow::Result<Message>> {
|
||||
try_stream! {
|
||||
for request in tool_requests {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
if let Ok(tool_call) = tool_request.tool_call.clone() {
|
||||
if self.is_frontend_tool(&tool_call.name).await {
|
||||
// Send frontend tool request and wait for response
|
||||
yield Message::assistant().with_frontend_tool_request(
|
||||
request.id.clone(),
|
||||
tool_request.id.clone(),
|
||||
Ok(tool_call.clone())
|
||||
);
|
||||
|
||||
@@ -151,7 +153,6 @@ impl Agent {
|
||||
*response = response.clone().with_tool_response(id, result);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
.boxed()
|
||||
|
||||
Reference in New Issue
Block a user