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());
|
||||
|
||||
Reference in New Issue
Block a user