refactor: remove unnecessary Arc<Mutex> from tool execution pipeline (#7979)
Signed-off-by: rabi <ramishra@redhat.com> Co-authored-by: jh-block <jhugo@block.xyz>
This commit is contained in:
@@ -395,7 +395,7 @@ impl Agent {
|
|||||||
async fn handle_approved_and_denied_tools(
|
async fn handle_approved_and_denied_tools(
|
||||||
&self,
|
&self,
|
||||||
permission_check_result: &PermissionCheckResult,
|
permission_check_result: &PermissionCheckResult,
|
||||||
request_to_response_map: &HashMap<String, Arc<Mutex<Message>>>,
|
request_to_response_map: &mut HashMap<String, Message>,
|
||||||
cancel_token: Option<tokio_util::sync::CancellationToken>,
|
cancel_token: Option<tokio_util::sync::CancellationToken>,
|
||||||
session: &Session,
|
session: &Session,
|
||||||
) -> Result<Vec<(String, ToolStream)>> {
|
) -> Result<Vec<(String, ToolStream)>> {
|
||||||
@@ -430,18 +430,17 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Self::handle_denied_tools(permission_check_result, request_to_response_map).await;
|
Self::handle_denied_tools(permission_check_result, request_to_response_map);
|
||||||
Ok(tool_futures)
|
Ok(tool_futures)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn handle_denied_tools(
|
fn handle_denied_tools(
|
||||||
permission_check_result: &PermissionCheckResult,
|
permission_check_result: &PermissionCheckResult,
|
||||||
request_to_response_map: &HashMap<String, Arc<Mutex<Message>>>,
|
request_to_response_map: &mut HashMap<String, Message>,
|
||||||
) {
|
) {
|
||||||
for request in &permission_check_result.denied {
|
for request in &permission_check_result.denied {
|
||||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
if let Some(response) = request_to_response_map.get_mut(&request.id) {
|
||||||
let mut response = response_msg.lock().await;
|
response.add_tool_response_with_metadata(
|
||||||
*response = response.clone().with_tool_response_with_metadata(
|
|
||||||
request.id.clone(),
|
request.id.clone(),
|
||||||
Ok(CallToolResult::error(vec![rmcp::model::Content::text(
|
Ok(CallToolResult::error(vec![rmcp::model::Content::text(
|
||||||
DECLINED_RESPONSE,
|
DECLINED_RESPONSE,
|
||||||
@@ -1246,21 +1245,19 @@ impl Agent {
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
let tool_response_messages: Vec<Arc<Mutex<Message>>> = (0..num_tool_requests)
|
|
||||||
.map(|_| Arc::new(Mutex::new(Message::user().with_generated_id())))
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
let mut request_to_response_map = HashMap::new();
|
let mut request_to_response_map = HashMap::new();
|
||||||
let mut request_metadata: HashMap<String, Option<ProviderMetadata>> = HashMap::new();
|
let mut request_metadata: HashMap<String, Option<ProviderMetadata>> = HashMap::new();
|
||||||
for (idx, request) in frontend_requests.iter().chain(remaining_requests.iter()).enumerate() {
|
for request in frontend_requests.iter().chain(remaining_requests.iter()) {
|
||||||
request_to_response_map.insert(request.id.clone(), tool_response_messages[idx].clone());
|
request_to_response_map.insert(request.id.clone(), Message::user().with_generated_id());
|
||||||
request_metadata.insert(request.id.clone(), request.metadata.clone());
|
request_metadata.insert(request.id.clone(), request.metadata.clone());
|
||||||
}
|
}
|
||||||
|
|
||||||
for (idx, request) in frontend_requests.iter().enumerate() {
|
for request in frontend_requests.iter() {
|
||||||
|
let response_msg = request_to_response_map.get_mut(&request.id)
|
||||||
|
.ok_or_else(|| anyhow::anyhow!("missing response entry for request {}", request.id))?;
|
||||||
let mut frontend_tool_stream = self.handle_frontend_tool_request(
|
let mut frontend_tool_stream = self.handle_frontend_tool_request(
|
||||||
request,
|
request,
|
||||||
tool_response_messages[idx].clone(),
|
response_msg,
|
||||||
);
|
);
|
||||||
|
|
||||||
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
while let Some(msg) = frontend_tool_stream.try_next().await? {
|
||||||
@@ -1268,11 +1265,9 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if goose_mode == GooseMode::Chat {
|
if goose_mode == GooseMode::Chat {
|
||||||
// Skip all remaining tool calls in chat mode
|
|
||||||
for request in remaining_requests.iter() {
|
for request in remaining_requests.iter() {
|
||||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
if let Some(response) = request_to_response_map.get_mut(&request.id) {
|
||||||
let mut response = response_msg.lock().await;
|
response.add_tool_response_with_metadata(
|
||||||
*response = response.clone().with_tool_response_with_metadata(
|
|
||||||
request.id.clone(),
|
request.id.clone(),
|
||||||
Ok(CallToolResult::success(vec![Content::text(CHAT_MODE_TOOL_SKIPPED_RESPONSE)])),
|
Ok(CallToolResult::success(vec![Content::text(CHAT_MODE_TOOL_SKIPPED_RESPONSE)])),
|
||||||
request.metadata.as_ref(),
|
request.metadata.as_ref(),
|
||||||
@@ -1317,31 +1312,26 @@ impl Agent {
|
|||||||
|
|
||||||
let mut tool_futures = self.handle_approved_and_denied_tools(
|
let mut tool_futures = self.handle_approved_and_denied_tools(
|
||||||
&permission_check_result,
|
&permission_check_result,
|
||||||
&request_to_response_map,
|
&mut request_to_response_map,
|
||||||
cancel_token.clone(),
|
cancel_token.clone(),
|
||||||
&session,
|
&session,
|
||||||
).await?;
|
).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,
|
||||||
|
&mut tool_futures,
|
||||||
|
&mut request_to_response_map,
|
||||||
|
cancel_token.clone(),
|
||||||
|
&session,
|
||||||
|
&inspection_results,
|
||||||
|
);
|
||||||
|
|
||||||
let mut tool_approval_stream = self.handle_approval_tool_requests(
|
while let Some(msg) = tool_approval_stream.try_next().await? {
|
||||||
&permission_check_result.needs_approval,
|
yield AgentEvent::Message(msg);
|
||||||
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
|
let with_id = tool_futures
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|(request_id, stream)| {
|
.map(|(request_id, stream)| {
|
||||||
@@ -1391,10 +1381,9 @@ impl Agent {
|
|||||||
{
|
{
|
||||||
all_install_successful = false;
|
all_install_successful = false;
|
||||||
}
|
}
|
||||||
if let Some(response_msg) = request_to_response_map.get(&request_id) {
|
if let Some(response) = request_to_response_map.get_mut(&request_id) {
|
||||||
let metadata = request_metadata.get(&request_id).and_then(|m| m.as_ref());
|
let metadata = request_metadata.get(&request_id).and_then(|m| m.as_ref());
|
||||||
let mut response = response_msg.lock().await;
|
response.add_tool_response_with_metadata(request_id, output, metadata);
|
||||||
*response = response.clone().with_tool_response_with_metadata(request_id, output, metadata);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
ToolStreamItem::Message(msg) => {
|
ToolStreamItem::Message(msg) => {
|
||||||
@@ -1447,12 +1436,11 @@ impl Agent {
|
|||||||
.cloned()
|
.cloned()
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
for (idx, request) in frontend_requests.iter().chain(remaining_requests.iter()).enumerate() {
|
for request in frontend_requests.iter().chain(remaining_requests.iter()) {
|
||||||
if request.tool_call.is_ok() {
|
if request.tool_call.is_ok() {
|
||||||
let mut request_msg = Message::assistant()
|
let mut request_msg = Message::assistant()
|
||||||
.with_id(format!("msg_{}", Uuid::new_v4()));
|
.with_id(format!("msg_{}", Uuid::new_v4()));
|
||||||
|
|
||||||
// Attach reasoning content to EVERY split tool request message.
|
|
||||||
// Providers like Kimi require reasoning_content on all assistant
|
// Providers like Kimi require reasoning_content on all assistant
|
||||||
// messages with tool_calls when thinking mode is enabled.
|
// messages with tool_calls when thinking mode is enabled.
|
||||||
for rc in &reasoning_content {
|
for rc in &reasoning_content {
|
||||||
@@ -1467,8 +1455,9 @@ impl Agent {
|
|||||||
request.tool_meta.clone(),
|
request.tool_meta.clone(),
|
||||||
);
|
);
|
||||||
messages_to_add.push(request_msg);
|
messages_to_add.push(request_msg);
|
||||||
let final_response = tool_response_messages[idx]
|
let final_response = request_to_response_map
|
||||||
.lock().await.clone();
|
.remove(&request.id)
|
||||||
|
.unwrap_or_else(|| Message::user().with_generated_id());
|
||||||
yield AgentEvent::Message(final_response.clone());
|
yield AgentEvent::Message(final_response.clone());
|
||||||
messages_to_add.push(final_response);
|
messages_to_add.push(final_response);
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -1,11 +1,9 @@
|
|||||||
use std::collections::HashMap;
|
|
||||||
use std::future::Future;
|
|
||||||
use std::sync::Arc;
|
|
||||||
|
|
||||||
use async_stream::try_stream;
|
use async_stream::try_stream;
|
||||||
use futures::stream::{self, BoxStream};
|
use futures::stream::{self, BoxStream};
|
||||||
use futures::{Stream, StreamExt};
|
use futures::{Stream, StreamExt};
|
||||||
use tokio::sync::Mutex;
|
use rmcp::model::CallToolResult;
|
||||||
|
use std::collections::HashMap;
|
||||||
|
use std::future::Future;
|
||||||
use tokio_util::sync::CancellationToken;
|
use tokio_util::sync::CancellationToken;
|
||||||
|
|
||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
@@ -79,8 +77,8 @@ impl Agent {
|
|||||||
pub(crate) fn handle_approval_tool_requests<'a>(
|
pub(crate) fn handle_approval_tool_requests<'a>(
|
||||||
&'a self,
|
&'a self,
|
||||||
tool_requests: &'a [ToolRequest],
|
tool_requests: &'a [ToolRequest],
|
||||||
tool_futures: Arc<Mutex<Vec<(String, ToolStream)>>>,
|
tool_futures: &'a mut Vec<(String, ToolStream)>,
|
||||||
request_to_response_map: &'a HashMap<String, Arc<Mutex<Message>>>,
|
request_to_response_map: &'a mut HashMap<String, Message>,
|
||||||
cancellation_token: Option<CancellationToken>,
|
cancellation_token: Option<CancellationToken>,
|
||||||
session: &'a Session,
|
session: &'a Session,
|
||||||
inspection_results: &'a [crate::tool_inspection::InspectionResult],
|
inspection_results: &'a [crate::tool_inspection::InspectionResult],
|
||||||
@@ -88,7 +86,6 @@ impl Agent {
|
|||||||
try_stream! {
|
try_stream! {
|
||||||
for request in tool_requests.iter() {
|
for request in tool_requests.iter() {
|
||||||
if let Ok(tool_call) = request.tool_call.clone() {
|
if let Ok(tool_call) = request.tool_call.clone() {
|
||||||
// Find the corresponding inspection result for this tool request
|
|
||||||
let security_message = inspection_results.iter()
|
let security_message = inspection_results.iter()
|
||||||
.find(|result| result.tool_request_id == request.id)
|
.find(|result| result.tool_request_id == request.id)
|
||||||
.and_then(|result| {
|
.and_then(|result| {
|
||||||
@@ -114,7 +111,6 @@ impl Agent {
|
|||||||
let confirmation = confirmation_rx.await
|
let confirmation = confirmation_rx.await
|
||||||
.map_err(|_| anyhow::anyhow!("Confirmation channel closed for request {}", request.id))?;
|
.map_err(|_| anyhow::anyhow!("Confirmation channel closed for request {}", 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) {
|
if let Some(finding_id) = get_security_finding_id_from_results(&request.id, inspection_results) {
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
monotonic_counter.goose.prompt_injection_user_decisions = 1,
|
monotonic_counter.goose.prompt_injection_user_decisions = 1,
|
||||||
@@ -127,9 +123,8 @@ impl Agent {
|
|||||||
|
|
||||||
if confirmation.permission == Permission::AllowOnce || confirmation.permission == Permission::AlwaysAllow {
|
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 (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 {
|
tool_futures.push((req_id, match tool_result {
|
||||||
Ok(result) => tool_stream(
|
Ok(result) => tool_stream(
|
||||||
result.notification_stream.unwrap_or_else(|| Box::new(stream::empty())),
|
result.notification_stream.unwrap_or_else(|| Box::new(stream::empty())),
|
||||||
result.result,
|
result.result,
|
||||||
@@ -140,19 +135,16 @@ impl Agent {
|
|||||||
),
|
),
|
||||||
}));
|
}));
|
||||||
|
|
||||||
// Update the shared permission manager when user selects "Always Allow"
|
|
||||||
if confirmation.permission == Permission::AlwaysAllow {
|
if confirmation.permission == Permission::AlwaysAllow {
|
||||||
self.tool_inspection_manager
|
self.tool_inspection_manager
|
||||||
.update_permission_manager(&tool_call.name, PermissionLevel::AlwaysAllow)
|
.update_permission_manager(&tool_call.name, PermissionLevel::AlwaysAllow)
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// User declined - update the specific response message for this request
|
if let Some(response) = request_to_response_map.get_mut(&request.id) {
|
||||||
if let Some(response_msg) = request_to_response_map.get(&request.id) {
|
response.add_tool_response_with_metadata(
|
||||||
let mut response = response_msg.lock().await;
|
|
||||||
*response = response.clone().with_tool_response_with_metadata(
|
|
||||||
request.id.clone(),
|
request.id.clone(),
|
||||||
Ok(rmcp::model::CallToolResult::error(vec![Content::text(DECLINED_RESPONSE)])),
|
Ok(CallToolResult::error(vec![Content::text(DECLINED_RESPONSE)])),
|
||||||
request.metadata.as_ref(),
|
request.metadata.as_ref(),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
@@ -171,20 +163,18 @@ impl Agent {
|
|||||||
pub(crate) fn handle_frontend_tool_request<'a>(
|
pub(crate) fn handle_frontend_tool_request<'a>(
|
||||||
&'a self,
|
&'a self,
|
||||||
tool_request: &'a ToolRequest,
|
tool_request: &'a ToolRequest,
|
||||||
message_tool_response: Arc<Mutex<Message>>,
|
message_tool_response: &'a mut Message,
|
||||||
) -> BoxStream<'a, anyhow::Result<Message>> {
|
) -> BoxStream<'a, anyhow::Result<Message>> {
|
||||||
try_stream! {
|
try_stream! {
|
||||||
if let Ok(tool_call) = tool_request.tool_call.clone() {
|
if let Ok(tool_call) = tool_request.tool_call.clone() {
|
||||||
if self.is_frontend_tool(&tool_call.name).await {
|
if self.is_frontend_tool(&tool_call.name).await {
|
||||||
// Send frontend tool request and wait for response
|
|
||||||
yield Message::assistant().with_frontend_tool_request(
|
yield Message::assistant().with_frontend_tool_request(
|
||||||
tool_request.id.clone(),
|
tool_request.id.clone(),
|
||||||
Ok(tool_call.clone())
|
Ok(tool_call.clone())
|
||||||
);
|
);
|
||||||
|
|
||||||
if let Some((id, result)) = self.tool_result_rx.lock().await.recv().await {
|
if let Some((id, result)) = self.tool_result_rx.lock().await.recv().await {
|
||||||
let mut response = message_tool_response.lock().await;
|
message_tool_response.add_tool_response_with_metadata(
|
||||||
*response = response.clone().with_tool_response_with_metadata(
|
|
||||||
id,
|
id,
|
||||||
result,
|
result,
|
||||||
tool_request.metadata.as_ref(),
|
tool_request.metadata.as_ref(),
|
||||||
|
|||||||
@@ -774,7 +774,6 @@ impl Message {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Add a tool response to the message
|
|
||||||
pub fn with_tool_response<S: Into<String>>(
|
pub fn with_tool_response<S: Into<String>>(
|
||||||
self,
|
self,
|
||||||
id: S,
|
id: S,
|
||||||
@@ -783,15 +782,16 @@ impl Message {
|
|||||||
self.with_content(MessageContent::tool_response(id, result))
|
self.with_content(MessageContent::tool_response(id, result))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn with_tool_response_with_metadata<S: Into<String>>(
|
pub fn add_tool_response_with_metadata<S: Into<String>>(
|
||||||
self,
|
&mut self,
|
||||||
id: S,
|
id: S,
|
||||||
result: ToolResult<CallToolResult>,
|
result: ToolResult<CallToolResult>,
|
||||||
metadata: Option<&ProviderMetadata>,
|
metadata: Option<&ProviderMetadata>,
|
||||||
) -> Self {
|
) {
|
||||||
self.with_content(MessageContent::tool_response_with_metadata(
|
self.content
|
||||||
id, result, metadata,
|
.push(MessageContent::tool_response_with_metadata(
|
||||||
))
|
id, result, metadata,
|
||||||
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Add an action required message for tool confirmation
|
/// Add an action required message for tool confirmation
|
||||||
|
|||||||
@@ -1012,7 +1012,8 @@ mod tests {
|
|||||||
"Should inherit"
|
"Should inherit"
|
||||||
);
|
);
|
||||||
|
|
||||||
let tool_response = Message::user().with_tool_response_with_metadata(
|
let mut tool_response = Message::user();
|
||||||
|
tool_response.add_tool_response_with_metadata(
|
||||||
req1.id.clone(),
|
req1.id.clone(),
|
||||||
Ok(tool_result("output")),
|
Ok(tool_result("output")),
|
||||||
req1.metadata.as_ref(),
|
req1.metadata.as_ref(),
|
||||||
|
|||||||
Reference in New Issue
Block a user