fix(context): deduplicate parallel tool-pair summaries (#11195)

Co-authored-by: Lifei Zhou <lifei@squareup.com>
This commit is contained in:
marquezini
2026-08-25 07:30:43 +00:00
committed by GitHub
parent bcbf3c7723
commit b0d5ae6fe5
+120 -1
View File
@@ -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::<Vec<_>>();
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<usize>,
captured_system: std::sync::Mutex<Option<String>>,
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<MessageStream, ProviderError> {
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("<mock summary>");