fix: conversation fixer merges assistant text blocks and drops empty text messages (#4898)

This commit is contained in:
Jack Amadeo
2025-10-01 15:44:59 -04:00
committed by GitHub
parent b22b35a3f4
commit ffe7e26640
2 changed files with 191 additions and 46 deletions
+2 -15
View File
@@ -379,7 +379,7 @@ impl From<PromptMessage> for Message {
} }
} }
#[derive(ToSchema, Clone, Copy, PartialEq, Serialize, Deserialize)] #[derive(ToSchema, Clone, Copy, PartialEq, Serialize, Deserialize, Debug)]
/// Metadata for message visibility /// Metadata for message visibility
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct MessageMetadata { pub struct MessageMetadata {
@@ -462,7 +462,7 @@ fn default_true() -> bool {
true true
} }
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize)] #[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
/// A message to or from an LLM /// A message to or from an LLM
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct Message { pub struct Message {
@@ -476,19 +476,6 @@ pub struct Message {
pub metadata: MessageMetadata, pub metadata: MessageMetadata,
} }
impl fmt::Debug for Message {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let joined_content: String = self
.content
.iter()
.map(|c| format!("{c}"))
.collect::<Vec<_>>()
.join(" ");
write!(f, "{:?}: {}", self.role, joined_content)
}
}
fn default_created() -> i64 { fn default_created() -> i64 {
0 // old messages do not have timestamps. 0 // old messages do not have timestamps.
} }
+189 -31
View File
@@ -168,20 +168,61 @@ pub fn fix_conversation(conversation: Conversation) -> (Conversation, Vec<String
} }
fn fix_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) { fn fix_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
let (messages_1, empty_removed) = remove_empty_messages(messages); [
let (messages_2, tool_calling_fixed) = fix_tool_calling(messages_1); merge_text_content_items,
let (messages_3, messages_merged) = merge_consecutive_messages(messages_2); remove_empty_messages,
let (messages_4, lead_trail_fixed) = fix_lead_trail(messages_3); fix_tool_calling,
let (messages_5, populated_if_empty) = populate_if_empty(messages_4); merge_consecutive_messages,
fix_lead_trail,
populate_if_empty,
]
.into_iter()
.fold(
(messages, Vec::new()),
|(msgs, mut all_issues), processor| {
let (new_msgs, issues) = processor(msgs);
all_issues.extend(issues);
(new_msgs, all_issues)
},
)
}
let mut issues = Vec::new(); fn merge_text_content_in_message(mut msg: Message) -> Message {
issues.extend(empty_removed); if msg.role != Role::Assistant {
issues.extend(tool_calling_fixed); return msg;
issues.extend(messages_merged); }
issues.extend(lead_trail_fixed); msg.content = msg
issues.extend(populated_if_empty); .content
.into_iter()
.fold(Vec::new(), |mut content, item| {
match item {
MessageContent::Text(text) => {
if let Some(MessageContent::Text(ref mut last)) = content.last_mut() {
last.text.push_str(&text.text);
} else {
content.push(MessageContent::Text(text));
}
}
other => content.push(other),
}
content
});
msg
}
(messages_5, issues) fn merge_text_content_items(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
messages.into_iter().fold(
(Vec::new(), Vec::new()),
|(mut messages, mut issues), message| {
let content_len = message.content.len();
let message = merge_text_content_in_message(message);
if content_len != message.content.len() {
issues.push(String::from("Merged text content"))
}
messages.push(message);
(messages, issues)
},
)
} }
fn remove_empty_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) { fn remove_empty_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
@@ -189,7 +230,11 @@ fn remove_empty_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>)
let filtered_messages = messages let filtered_messages = messages
.into_iter() .into_iter()
.filter(|msg| { .filter(|msg| {
if msg.content.is_empty() { if msg
.content
.iter()
.all(|c| c.as_text().is_some_and(str::is_empty))
{
issues.push("Removed empty message".to_string()); issues.push("Removed empty message".to_string());
false false
} else { } else {
@@ -402,6 +447,24 @@ mod tests {
use rmcp::model::{CallToolRequestParam, Role}; use rmcp::model::{CallToolRequestParam, Role};
use rmcp::object; use rmcp::object;
macro_rules! assert_has_issues_unordered {
($fixed:expr, $issues:expr, $($expected:expr),+ $(,)?) => {
{
let mut expected: Vec<&str> = vec![$($expected),+];
let mut actual: Vec<&str> = $issues.iter().map(|s| s.as_str()).collect();
expected.sort();
actual.sort();
if actual != expected {
panic!(
"assertion failed: issues don't match\nexpected: {:?}\n actual: {:?}. Fixed conversation is:\n{:#?}",
expected, $issues, $fixed,
);
}
}
};
}
fn run_verify(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) { fn run_verify(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
let (fixed, issues) = fix_conversation(Conversation::new_unvalidated(messages.clone())); let (fixed, issues) = fix_conversation(Conversation::new_unvalidated(messages.clone()));
@@ -486,17 +549,15 @@ mod tests {
let (fixed, issues) = run_verify(messages); let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3); assert_eq!(fixed.len(), 3);
assert_eq!(issues.len(), 4);
assert!(issues assert_has_issues_unordered!(
.iter() fixed,
.any(|i| i.contains("Merged consecutive user messages"))); issues,
assert!(issues "Merged consecutive assistant messages",
.iter() "Merged consecutive user messages",
.any(|i| i.contains("Removed tool response 'orphan_1' from assistant message"))); "Removed tool response 'orphan_1' from assistant message",
assert!(issues "Removed tool request 'bad_req' from user message",
.iter() );
.any(|i| i.contains("Removed tool request 'bad_req' from user message")));
assert_eq!(fixed[0].role, Role::User); assert_eq!(fixed[0].role, Role::User);
assert_eq!(fixed[1].role, Role::Assistant); assert_eq!(fixed[1].role, Role::Assistant);
@@ -536,10 +597,18 @@ mod tests {
assert_eq!(fixed.len(), 1); assert_eq!(fixed.len(), 1);
assert!(issues.iter().any(|i| i.contains("Removed empty message"))); assert_has_issues_unordered!(
assert!(issues fixed,
.iter() issues,
.any(|i| i.contains("Removed orphaned tool response 'wrong_id'"))); "Removed empty message",
"Removed orphaned tool response 'wrong_id'",
"Removed orphaned tool request 'search_1'",
"Removed orphaned tool request 'search_2'",
"Removed empty message",
"Removed empty message",
"Removed leading assistant message",
"Added placeholder user message to empty conversation",
);
assert_eq!(fixed[0].role, Role::User); assert_eq!(fixed[0].role, Role::User);
assert_eq!(fixed[0].as_concat_text(), "Hello"); assert_eq!(fixed[0].as_concat_text(), "Hello");
@@ -569,9 +638,12 @@ mod tests {
let (fixed, issues) = fix_conversation(conversation); let (fixed, issues) = fix_conversation(conversation);
assert_eq!(fixed.len(), 5); assert_eq!(fixed.len(), 5);
assert_eq!(issues.len(), 2); assert_has_issues_unordered!(
assert!(issues[0].contains("Removed orphaned tool request")); fixed,
assert!(issues[1].contains("Merged consecutive assistant messages")); issues,
"Removed orphaned tool request 'toolu_bdrk_018adWbP4X26CfoJU5hkhu3i'",
"Merged consecutive assistant messages"
)
} }
#[test] #[test]
@@ -592,6 +664,92 @@ mod tests {
]; ];
let (_fixed, issues) = run_verify(messages); let (_fixed, issues) = run_verify(messages);
assert_eq!(issues.len(), 0); assert!(issues.is_empty());
}
#[test]
fn test_merge_text_content_items() {
use crate::conversation::message::MessageContent;
use rmcp::model::{AnnotateAble, RawTextContent};
let mut message = Message::assistant().with_text("Hello");
message.content.push(MessageContent::Text(
RawTextContent {
text: " world".to_string(),
meta: None,
}
.no_annotation(),
));
message.content.push(MessageContent::Text(
RawTextContent {
text: "!".to_string(),
meta: None,
}
.no_annotation(),
));
let messages = vec![
Message::user().with_text("hello"),
message,
Message::user().with_text("thanks"),
];
let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3);
assert_has_issues_unordered!(fixed, issues, "Merged text content");
let fixed_msg = &fixed[1];
assert_eq!(fixed_msg.content.len(), 1);
if let MessageContent::Text(text_content) = &fixed_msg.content[0] {
assert_eq!(text_content.text, "Hello world!");
} else {
panic!("Expected text content");
}
}
#[test]
fn test_merge_text_content_items_with_mixed_content() {
use crate::conversation::message::MessageContent;
use rmcp::model::{AnnotateAble, RawTextContent};
let mut image_message = Message::assistant().with_text("Look at");
image_message.content.push(MessageContent::Text(
RawTextContent {
text: " this image:".to_string(),
meta: None,
}
.no_annotation(),
));
image_message = image_message.with_image("", "");
let messages = vec![
Message::user().with_text("hello"),
image_message,
Message::user().with_text("thanks"),
];
let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3);
assert_has_issues_unordered!(fixed, issues, "Merged text content");
let fixed_msg = &fixed[1];
assert_eq!(fixed_msg.content.len(), 2);
if let MessageContent::Text(text_content) = &fixed_msg.content[0] {
assert_eq!(text_content.text, "Look at this image:");
} else {
panic!("Expected first item to be text content");
}
if let MessageContent::Image(_) = &fixed_msg.content[1] {
// Good
} else {
panic!("Expected second item to be an image");
}
} }
} }