221 lines
8.0 KiB
Rust
221 lines
8.0 KiB
Rust
use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageContent, MessageRole};
|
|
|
|
/// Sanitizes messages for OpenAI API compatibility.
|
|
///
|
|
/// OpenAI is more lenient than Bedrock — it doesn't require strict user/assistant
|
|
/// alternation and allows system messages anywhere. The main constraints are:
|
|
/// - Tool results must reference a valid tool_call_id from a preceding assistant message
|
|
/// - Tool calls in assistant messages must eventually have matching tool results
|
|
pub fn sanitize_messages_for_openai(messages: &mut Vec<ConversationMessage>) {
|
|
remove_orphaned_tool_results(messages);
|
|
synthesize_missing_tool_results(messages);
|
|
}
|
|
|
|
/// 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();
|
|
|
|
// First pass: collect all tool_use_ids from assistant messages
|
|
for msg in messages.iter() {
|
|
if msg.role != MessageRole::Assistant {
|
|
continue;
|
|
}
|
|
collect_tool_use_ids(&msg.content, &mut known_tool_use_ids);
|
|
}
|
|
|
|
// Second pass: remove tool_results that reference unknown IDs
|
|
messages.retain(|msg| {
|
|
if msg.role != MessageRole::User {
|
|
return true;
|
|
}
|
|
match &msg.content {
|
|
MessageContent::ToolResult { tool_use_id, .. } => {
|
|
known_tool_use_ids.contains(tool_use_id)
|
|
}
|
|
MessageContent::MultiPart(parts) => {
|
|
// Keep the message if it has at least one non-orphaned part
|
|
parts.iter().any(|part| match part {
|
|
ContentPart::ToolResult { tool_use_id, .. } => {
|
|
known_tool_use_ids.contains(tool_use_id)
|
|
}
|
|
_ => true,
|
|
})
|
|
}
|
|
_ => true,
|
|
}
|
|
});
|
|
}
|
|
|
|
/// 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();
|
|
|
|
// First pass: collect all existing tool_result IDs
|
|
for msg in messages.iter() {
|
|
if msg.role == MessageRole::User {
|
|
collect_tool_result_ids(&msg.content, &mut answered_ids);
|
|
}
|
|
}
|
|
|
|
// Second pass: walk through messages and insert synthetic results after
|
|
// assistant tool_use messages that have unanswered IDs.
|
|
let mut i = 0;
|
|
while i < messages.len() {
|
|
if messages[i].role != MessageRole::Assistant {
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
let mut unanswered: Vec<String> = Vec::new();
|
|
collect_tool_use_ids_vec(&messages[i].content, &mut unanswered);
|
|
unanswered.retain(|id| !answered_ids.contains(id));
|
|
|
|
if unanswered.is_empty() {
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
log::warn!(
|
|
"[openai] Synthesizing {} missing tool_result(s) after message {} for IDs: {:?}",
|
|
unanswered.len(),
|
|
i,
|
|
unanswered
|
|
);
|
|
|
|
let synthetic_parts: Vec<ContentPart> = unanswered
|
|
.iter()
|
|
.map(|id| ContentPart::ToolResult {
|
|
tool_use_id: id.clone(),
|
|
content: "Tool call result unavailable (conversation was interrupted).".to_string(),
|
|
is_error: true,
|
|
})
|
|
.collect();
|
|
|
|
let insert_idx = i + 1;
|
|
|
|
// If next message is a user message, merge synthetic results into it
|
|
if insert_idx < messages.len() && messages[insert_idx].role == MessageRole::User {
|
|
match &mut messages[insert_idx].content {
|
|
MessageContent::MultiPart(parts) => {
|
|
let existing = std::mem::take(parts);
|
|
parts.extend(synthetic_parts);
|
|
parts.extend(existing);
|
|
}
|
|
existing => {
|
|
let existing_part =
|
|
match std::mem::replace(existing, MessageContent::Text(String::new())) {
|
|
MessageContent::Text(t) => ContentPart::Text(t),
|
|
MessageContent::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
} => ContentPart::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
},
|
|
MessageContent::ToolUse {
|
|
tool_use_id,
|
|
name,
|
|
input,
|
|
} => ContentPart::ToolUse {
|
|
tool_use_id,
|
|
name,
|
|
input,
|
|
},
|
|
MessageContent::MultiPart(_) => unreachable!(),
|
|
};
|
|
let mut parts = synthetic_parts;
|
|
parts.push(existing_part);
|
|
*existing = MessageContent::MultiPart(parts);
|
|
}
|
|
}
|
|
} else {
|
|
// No user message follows — insert a new one
|
|
let content = if synthetic_parts.len() == 1 {
|
|
match synthetic_parts.into_iter().next().unwrap() {
|
|
ContentPart::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
} => MessageContent::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
},
|
|
_ => unreachable!(),
|
|
}
|
|
} else {
|
|
MessageContent::MultiPart(synthetic_parts)
|
|
};
|
|
messages.insert(
|
|
insert_idx,
|
|
ConversationMessage {
|
|
role: MessageRole::User,
|
|
content,
|
|
},
|
|
);
|
|
}
|
|
|
|
// Mark these as answered so we don't double-synthesize
|
|
for id in unanswered {
|
|
answered_ids.insert(id);
|
|
}
|
|
|
|
i += 2; // Skip past the inserted/modified message
|
|
}
|
|
}
|
|
|
|
fn collect_tool_use_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
|
|
match content {
|
|
MessageContent::ToolUse { tool_use_id, .. } => {
|
|
ids.insert(tool_use_id.clone());
|
|
}
|
|
MessageContent::MultiPart(parts) => {
|
|
for part in parts {
|
|
if let ContentPart::ToolUse { tool_use_id, .. } = part {
|
|
ids.insert(tool_use_id.clone());
|
|
}
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|
|
|
|
fn collect_tool_use_ids_vec(content: &MessageContent, ids: &mut Vec<String>) {
|
|
match content {
|
|
MessageContent::ToolUse { tool_use_id, .. } => {
|
|
ids.push(tool_use_id.clone());
|
|
}
|
|
MessageContent::MultiPart(parts) => {
|
|
for part in parts {
|
|
if let ContentPart::ToolUse { tool_use_id, .. } = part {
|
|
ids.push(tool_use_id.clone());
|
|
}
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|
|
|
|
fn collect_tool_result_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
|
|
match content {
|
|
MessageContent::ToolResult { tool_use_id, .. } => {
|
|
ids.insert(tool_use_id.clone());
|
|
}
|
|
MessageContent::MultiPart(parts) => {
|
|
for part in parts {
|
|
if let ContentPart::ToolResult { tool_use_id, .. } = part {
|
|
ids.insert(tool_use_id.clone());
|
|
}
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|