Fix provider tool history handling

This commit is contained in:
2026-08-12 14:19:51 -05:00
parent c79634e76f
commit b3f3a72435
18 changed files with 838 additions and 34 deletions
+58 -5
View File
@@ -1,3 +1,5 @@
use std::collections::HashSet;
use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageContent, MessageRole};
/// Sanitizes messages for OpenAI API compatibility.
@@ -9,6 +11,7 @@ use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageConten
/// - When routed to Bedrock via LiteLLM, the conversation must end with a user message
pub fn sanitize_messages_for_openai(messages: &mut Vec<ConversationMessage>) {
remove_orphaned_tool_results(messages);
remove_misplaced_tool_results(messages);
synthesize_missing_tool_results(messages);
ensure_ends_with_user_message(messages);
}
@@ -16,8 +19,7 @@ pub fn sanitize_messages_for_openai(messages: &mut Vec<ConversationMessage>) {
/// Removes tool_result messages that reference tool_use_ids not found in any
/// preceding assistant message.
fn remove_orphaned_tool_results(messages: &mut Vec<ConversationMessage>) {
let mut known_tool_use_ids: std::collections::HashSet<String> =
std::collections::HashSet::new();
let mut known_tool_use_ids = HashSet::new();
// First pass: collect all tool_use_ids from assistant messages
for msg in messages.iter() {
@@ -50,12 +52,63 @@ fn remove_orphaned_tool_results(messages: &mut Vec<ConversationMessage>) {
});
}
/// LiteLLM may route OpenAI-compatible requests to Bedrock, which requires a
/// user turn containing tool_result blocks to directly answer the tool_use
/// blocks from the immediately previous assistant turn. Late results from
/// cancelled or superseded actions are valid history globally, but invalid in
/// that later user turn, so drop them before request conversion.
fn remove_misplaced_tool_results(messages: &mut Vec<ConversationMessage>) {
let mut i = 0;
while i < messages.len() {
if messages[i].role != MessageRole::User {
i += 1;
continue;
}
let mut allowed_tool_use_ids = if i > 0 && messages[i - 1].role == MessageRole::Assistant {
let mut ids = HashSet::new();
collect_tool_use_ids(&messages[i - 1].content, &mut ids);
ids
} else {
HashSet::new()
};
if retain_allowed_tool_results(&mut messages[i].content, &mut allowed_tool_use_ids) {
i += 1;
} else {
messages.remove(i);
}
}
}
fn retain_allowed_tool_results(
content: &mut MessageContent,
allowed_tool_use_ids: &mut HashSet<String>,
) -> bool {
match content {
MessageContent::Text(_) | MessageContent::ToolUse { .. } => true,
MessageContent::ToolResult { tool_use_id, .. } => allowed_tool_use_ids.remove(tool_use_id),
MessageContent::MultiPart(parts) => {
parts.retain(|part| match part {
ContentPart::ToolResult { tool_use_id, .. } => {
allowed_tool_use_ids.remove(tool_use_id)
}
ContentPart::Text(_)
| ContentPart::Reasoning { .. }
| ContentPart::Image { .. }
| ContentPart::ToolUse { .. } => true,
});
!parts.is_empty()
}
}
}
/// For any assistant tool_use that doesn't have a matching tool_result in a
/// subsequent user message, synthesize a result immediately after the tool_use.
/// This satisfies Bedrock's requirement (via LiteLLM) that tool_result blocks
/// appear immediately after the corresponding tool_use message.
fn synthesize_missing_tool_results(messages: &mut Vec<ConversationMessage>) {
let mut answered_ids: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut answered_ids = HashSet::new();
// First pass: collect all existing tool_result IDs
for msg in messages.iter() {
@@ -173,7 +226,7 @@ fn synthesize_missing_tool_results(messages: &mut Vec<ConversationMessage>) {
}
}
fn collect_tool_use_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
fn collect_tool_use_ids(content: &MessageContent, ids: &mut HashSet<String>) {
match content {
MessageContent::ToolUse { tool_use_id, .. } => {
ids.insert(tool_use_id.clone());
@@ -205,7 +258,7 @@ fn collect_tool_use_ids_vec(content: &MessageContent, ids: &mut Vec<String>) {
}
}
fn collect_tool_result_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
fn collect_tool_result_ids(content: &MessageContent, ids: &mut HashSet<String>) {
match content {
MessageContent::ToolResult { tool_use_id, .. } => {
ids.insert(tool_use_id.clone());
@@ -168,6 +168,106 @@ fn test_multipart_tool_uses_all_get_results() {
}
}
#[test]
fn sanitizer_drops_stale_tool_result_from_current_user_turn() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "old_tool".to_string(),
name: "read_shell_command_output".to_string(),
input: json!({"block_id": "block-1"}),
},
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "old_tool".to_string(),
content: "cancelled".to_string(),
is_error: true,
},
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "current_tool".to_string(),
name: "read_notebook".to_string(),
input: json!({"document_id": "doc-1"}),
},
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::MultiPart(vec![
ContentPart::ToolResult {
tool_use_id: "old_tool".to_string(),
content: "late cancellation".to_string(),
is_error: true,
},
ContentPart::ToolResult {
tool_use_id: "current_tool".to_string(),
content: "notebook contents".to_string(),
is_error: false,
},
]),
},
];
sanitize_messages_for_openai(&mut messages);
assert_eq!(messages.len(), 4);
let MessageContent::MultiPart(parts) = &messages[3].content else {
panic!("expected current user message to remain multipart");
};
assert_eq!(parts.len(), 1);
assert!(matches!(
&parts[0],
ContentPart::ToolResult { tool_use_id, content, is_error }
if tool_use_id == "current_tool" && content == "notebook contents" && !is_error
));
}
#[test]
fn sanitizer_drops_tool_result_message_not_following_its_tool_use() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "old_tool".to_string(),
name: "read_shell_command_output".to_string(),
input: json!({"block_id": "block-1"}),
},
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "old_tool".to_string(),
content: "cancelled".to_string(),
is_error: true,
},
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text("Continuing.".to_string()),
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "old_tool".to_string(),
content: "late cancellation".to_string(),
is_error: true,
},
},
];
sanitize_messages_for_openai(&mut messages);
assert_eq!(messages.len(), 4);
let MessageContent::Text(text) = &messages[3].content else {
panic!("expected appended continuation message after stale tool result was dropped");
};
assert_eq!(text, "Continue.");
}
#[test]
fn test_ensure_ends_with_user_message_no_op_when_already_user() {
let mut messages = vec![
+2
View File
@@ -679,6 +679,8 @@ const KNOWN_TOOLS: &[&str] = &[
"read_documents",
"create_documents",
"edit_documents",
"run_agents",
"wait_for_events",
"start_agent",
"ask_user_question",
"read_skill",
@@ -136,6 +136,8 @@ async fn recall_tool_history_does_not_emit_a_client_tool_call() {
fn direct_provider_known_tools_exclude_hosted_only_tools() {
assert!(!is_known_tool("send_message_to_agent"));
assert!(!is_known_tool("suggest_next_prompt"));
assert!(is_known_tool("run_agents"));
assert!(is_known_tool("wait_for_events"));
assert!(is_known_tool("recall_tool_history"));
assert!(is_known_tool("interrupt_shell_command"));
}
+12 -2
View File
@@ -8,7 +8,9 @@ use super::request_translator::sanitize_messages_for_openai;
use super::response_translator::{openai_stream_to_response_events, OpenAIStreamContext};
use crate::ai::agent::api::LegacyResponseStream;
use crate::ai::bedrock::request_translator;
use crate::ai::provider::types::{ConversationMessage, MessageContent, MessageRole};
use crate::ai::provider::types::{
flatten_tool_history_for_no_tools_turn, ConversationMessage, MessageContent, MessageRole,
};
const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 64_000;
@@ -92,6 +94,10 @@ pub(crate) fn prepare_turn(params: &TranslatorRequest, request: &mut api::Reques
message.truncate_tool_results_for_provider_request();
}
sanitize_messages_for_openai(&mut messages);
let tools = request_translator::extract_tools(request);
if tools_are_inline_only(&tools) {
flatten_tool_history_for_no_tools_turn(&mut messages);
}
PreparedTurn {
task_id,
@@ -99,12 +105,16 @@ pub(crate) fn prepare_turn(params: &TranslatorRequest, request: &mut api::Reques
user_query: request_translator::extract_user_query_text(request),
messages,
system_prompt: request_translator::extract_system_prompt(request, &params.global_rules),
tools: request_translator::extract_tools(request),
tools,
model_id,
persistent_message_count,
}
}
fn tools_are_inline_only(tools: &[crate::ai::provider::types::ToolDefinition]) -> bool {
tools.iter().all(|tool| tool.name == "recall_tool_history")
}
pub async fn execute(
params: TranslatorRequest,
request: &mut api::Request,