From b0d5ae6fe5ac66f6b69f87e4b7d8d3fecc04b742 Mon Sep 17 00:00:00 2001 From: marquezini Date: Tue, 25 Aug 2026 07:30:43 +0000 Subject: [PATCH] fix(context): deduplicate parallel tool-pair summaries (#11195) Co-authored-by: Lifei Zhou --- crates/goose/src/context_mgmt/mod.rs | 121 ++++++++++++++++++++++++++- 1 file changed, 120 insertions(+), 1 deletion(-) diff --git a/crates/goose/src/context_mgmt/mod.rs b/crates/goose/src/context_mgmt/mod.rs index 7f1c1c2ea..4242e4805 100644 --- a/crates/goose/src/context_mgmt/mod.rs +++ b/crates/goose/src/context_mgmt/mod.rs @@ -509,9 +509,69 @@ pub fn maybe_summarize_tool_pairs( return None; } + // A request/response message can contain multiple parallel tool calls. + // Summarization formats whole messages, so sibling IDs in the same pair + // must be compacted as one group or we'd issue duplicate summary calls. + let mut seen_message_pairs = std::collections::HashSet::new(); + let mut grouped_tool_ids = Vec::new(); + for tool_id in tool_ids { + let pair = match agent_visible_tool_pair(&conversation, &tool_id) { + Ok(pair) => pair, + Err(error) => { + warn!("Failed to identify tool pair for summarization: {}", error); + continue; + } + }; + if pair.len() != 2 { + warn!( + "Expected a tool request/response pair for '{}', found {} messages", + tool_id, + pair.len() + ); + continue; + } + + let request_ids: std::collections::HashSet<&str> = pair + .iter() + .flat_map(|message| message.get_tool_request_ids()) + .collect(); + let response_ids: std::collections::HashSet<&str> = pair + .iter() + .flat_map(|message| message.get_tool_response_ids()) + .collect(); + if request_ids != response_ids { + warn!( + "Tool pair for '{}' has siblings answered elsewhere; skipping", + tool_id + ); + continue; + } + + let mut message_ids = pair + .iter() + .filter_map(|message| message.id.clone()) + .collect::>(); + if message_ids.len() != 2 { + warn!( + "Expected two persisted messages for tool pair '{}', found {}", + tool_id, + message_ids.len() + ); + continue; + } + message_ids.sort_unstable(); + if seen_message_pairs.insert(message_ids) { + grouped_tool_ids.push(tool_id); + } + } + + if grouped_tool_ids.is_empty() { + return None; + } + Some(tokio::spawn(async move { let mut results = Vec::new(); - for tool_id in tool_ids { + for tool_id in grouped_tool_ids { match summarize_tool_call( provider.as_ref(), &model_config, @@ -567,6 +627,7 @@ mod tests { config: ModelConfig, max_tool_responses: Option, captured_system: std::sync::Mutex>, + calls: std::sync::atomic::AtomicUsize, } impl MockProvider { @@ -587,6 +648,7 @@ mod tests { }, max_tool_responses: None, captured_system: std::sync::Mutex::new(None), + calls: std::sync::atomic::AtomicUsize::new(0), } } @@ -594,6 +656,10 @@ mod tests { self.max_tool_responses = Some(max); self } + + fn call_count(&self) -> usize { + self.calls.load(std::sync::atomic::Ordering::Relaxed) + } } #[async_trait] @@ -609,6 +675,8 @@ mod tests { messages: &[Message], _tools: &[Tool], ) -> Result { + self.calls + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); *self.captured_system.lock().unwrap() = Some(system.to_string()); // If max_tool_responses is set, fail if we have too many if let Some(max) = self.max_tool_responses { @@ -1110,6 +1178,57 @@ mod tests { assert!(error.to_string().contains("No agent-visible tool pair")); } + #[tokio::test] + async fn parallel_tool_calls_share_one_summary_request() { + let provider = Arc::new(MockProvider::new( + Message::assistant().with_text("summary"), + 1000, + )); + let mut messages = vec![Message::user().with_text("start")]; + for index in 0..7 { + let first_id = format!("tool_{index}a"); + let second_id = format!("tool_{index}b"); + messages.push( + Message::assistant() + .with_tool_request(&first_id, Ok(CallToolRequestParams::new("read_file"))) + .with_tool_request(&second_id, Ok(CallToolRequestParams::new("read_file"))) + .with_id(format!("request_{index}")), + ); + messages.push( + Message::user() + .with_tool_response( + &first_id, + Ok(rmcp::model::CallToolResult::success(vec![ + ContentBlock::text("first result"), + ])), + ) + .with_tool_response( + &second_id, + Ok(rmcp::model::CallToolResult::success(vec![ + ContentBlock::text("second result"), + ])), + ) + .with_id(format!("response_{index}")), + ); + } + + let conversation = Conversation::new_unvalidated(messages); + let summaries = maybe_summarize_tool_pairs( + provider.clone(), + provider.config.clone(), + "test-session-id".to_string(), + conversation, + 2, + 0, + ) + .unwrap() + .await + .unwrap(); + + assert_eq!(summaries.len(), 5); + assert_eq!(provider.call_count(), 5); + } + #[tokio::test] async fn test_progressive_removal_on_context_exceeded() { let response_message = Message::assistant().with_text("");