Fix provider tool history handling
This commit is contained in:
@@ -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![
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
|
||||
@@ -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, ¶ms.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,
|
||||
|
||||
Reference in New Issue
Block a user