Files
baibot-withmcp/src/conversation/llm/entity.rs
kschwank 2d659964a7 Add sender context mode for text generation (#104)
Add a per-room/global `sender_context_mode` setting that optionally
prefixes conversation messages with sender metadata before sending them
to the model provider.

This helps models distinguish between participants in multi-user rooms.
                                                                                                                                                                                                                                                             
Three modes are supported:
- `disabled` (default, no change)
- `matrix_user_id` (prefixes with `[sender=@user:server]`)
- `matrix_user_id_and_timestamp` (adds `send_at`; Example: `[sender=@user:server sent_at=<ISO 8601>]`)
                                                                                                                                                                                                                                                             
Sender context is applied to user and assistant text messages only,
skipping system prompts and non-text content.

Mixed-sender merged turns (something we intentionally do for Anthropic)
have their `sender_id` cleared to avoid misattribution.
2026-03-25 20:14:34 +02:00

326 lines
11 KiB
Rust

use chrono::{DateTime, Utc};
use mxlink::matrix_sdk::ruma::OwnedUserId;
use mxlink::matrix_sdk::ruma::events::room::message::{
FileMessageEventContent, ImageMessageEventContent,
};
use mxlink::mime::Mime;
use crate::agent::provider::ImageSource;
#[derive(Debug, Clone, PartialEq)]
pub enum Author {
Prompt,
Assistant,
User,
}
#[derive(Debug, Clone)]
pub struct Message {
pub author: Author,
pub sender_id: Option<OwnedUserId>,
pub timestamp: DateTime<Utc>,
pub content: MessageContent,
}
#[derive(Debug, Clone)]
pub struct ImageDetails {
pub event_content: ImageMessageEventContent,
pub mime: Mime,
pub data: Vec<u8>,
}
impl ImageDetails {
pub fn new(event_content: ImageMessageEventContent, mime: Mime, data: Vec<u8>) -> Self {
Self {
event_content,
mime,
data,
}
}
pub fn filename(&self) -> String {
self.event_content
.filename
.clone()
.unwrap_or(self.event_content.body.clone())
}
}
impl From<ImageDetails> for ImageSource {
fn from(value: ImageDetails) -> Self {
ImageSource::new(value.filename(), value.data.clone(), value.mime.clone())
}
}
#[derive(Debug, Clone)]
pub struct FileDetails {
pub event_content: FileMessageEventContent,
pub mime: Mime,
pub data: Vec<u8>,
}
impl FileDetails {
pub fn new(event_content: FileMessageEventContent, mime: Mime, data: Vec<u8>) -> Self {
Self {
event_content,
mime,
data,
}
}
pub fn filename(&self) -> String {
self.event_content
.filename
.clone()
.unwrap_or(self.event_content.body.clone())
}
}
#[derive(Debug, Clone)]
pub enum MessageContent {
Text(String),
Image(ImageDetails),
File(FileDetails),
}
impl PartialEq for MessageContent {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(MessageContent::Text(a), MessageContent::Text(b)) => a == b,
(MessageContent::Image(a), MessageContent::Image(b)) => {
// We can probably do better than this by inspecting `.event_conten1t.source`, but for now this is good enough.
a.filename() == b.filename()
}
(MessageContent::File(a), MessageContent::File(b)) => a.filename() == b.filename(),
_ => false,
}
}
}
#[derive(Debug)]
pub struct Conversation {
pub messages: Vec<Message>,
}
impl Conversation {
/// Combine consecutive messages by the same author into a single message.
///
/// Certain models (like Anthropic) cannot tolerate consecutive messages by the same author,
/// so combining them helps avoid issues.
///
/// When multiple text messages by the same author are merged, the resulting message keeps a
/// `sender_id` only if all merged messages came from the same sender. Mixed-sender merges are
/// possible for user turns in multi-user rooms, so `sender_id` is cleared in that case to
/// avoid incorrectly attributing the whole merged turn to the first sender.
/// See: https://github.com/etkecc/baibot/issues/13
pub fn combine_consecutive_messages(&self) -> Conversation {
// We'll likely get fewer messages, but let's reserve the maximum we expect.
let mut new_messages = Vec::with_capacity(self.messages.len());
let mut last_seen_text_from_author: Option<Author> = None;
for message in &self.messages {
let MessageContent::Text(message_text_content) = &message.content else {
last_seen_text_from_author = None;
new_messages.push(message.clone());
continue;
};
let Some(last_seen_author_clone) = last_seen_text_from_author.clone() else {
last_seen_text_from_author = Some(message.author.clone());
new_messages.push(message.clone());
continue;
};
if message.author != last_seen_author_clone {
last_seen_text_from_author = Some(message.author.clone());
new_messages.push(message.clone());
continue;
}
let last_message = new_messages.last_mut().unwrap();
if let MessageContent::Text(ref mut text) = last_message.content {
text.push('\n');
text.push_str(message_text_content);
}
if last_message.sender_id != message.sender_id {
last_message.sender_id = None;
}
}
Conversation {
messages: new_messages,
}
}
pub fn start_time(&self) -> Option<DateTime<Utc>> {
self.messages.first().map(|message| message.timestamp)
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{TimeZone, Utc};
use mxlink::matrix_sdk::ruma::{OwnedMxcUri, OwnedUserId};
use mxlink::mime;
#[test]
fn combine_consecutive_messages() {
let timestamp_1 = Utc.with_ymd_and_hms(2024, 9, 20, 18, 34, 15).unwrap();
let timestamp_2 = Utc.with_ymd_and_hms(2024, 9, 21, 18, 34, 16).unwrap();
let timestamp_3 = Utc.with_ymd_and_hms(2024, 9, 22, 18, 34, 17).unwrap();
let timestamp_4 = Utc.with_ymd_and_hms(2024, 9, 23, 18, 34, 18).unwrap();
let image_event_content = ImageMessageEventContent::plain(
"image.png".to_string(),
OwnedMxcUri::from("mxc://example.com/1234567890"),
);
let conversation = Conversation {
messages: vec![
// User's turn
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("Hello".to_string()),
timestamp: timestamp_1,
},
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("How are you?".to_string()),
timestamp: timestamp_2,
},
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("I'm OK, btw.".to_string()),
timestamp: timestamp_3,
},
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Image(ImageDetails::new(
image_event_content.clone(),
mime::IMAGE_PNG,
vec![],
)),
timestamp: timestamp_4,
},
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("Above is an image.".to_string()),
timestamp: timestamp_4,
},
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("Would you take a look at it?".to_string()),
timestamp: timestamp_4,
},
// Assistant's turn
Message {
author: Author::Assistant,
sender_id: None,
content: MessageContent::Text("Hi there!".to_string()),
timestamp: timestamp_2,
},
Message {
author: Author::Assistant,
sender_id: None,
content: MessageContent::Text("I'm doing well, thank you.".to_string()),
timestamp: timestamp_3,
},
// User's turn
Message {
author: Author::User,
sender_id: None,
content: MessageContent::Text("That's great!".to_string()),
timestamp: timestamp_3,
},
],
};
let conversation = conversation.combine_consecutive_messages();
assert_eq!(conversation.messages.len(), 5);
assert_eq!(conversation.messages[0].author, Author::User);
assert_eq!(
conversation.messages[0].content,
MessageContent::Text("Hello\nHow are you?\nI'm OK, btw.".to_string())
);
assert_eq!(conversation.messages[0].timestamp, timestamp_1);
assert_eq!(conversation.messages[1].author, Author::User);
assert_eq!(
conversation.messages[1].content,
MessageContent::Image(ImageDetails::new(
image_event_content.clone(),
mime::IMAGE_PNG,
vec![],
))
);
assert_eq!(conversation.messages[2].author, Author::User);
assert_eq!(
conversation.messages[2].content,
MessageContent::Text("Above is an image.\nWould you take a look at it?".to_string())
);
assert_eq!(conversation.messages[2].timestamp, timestamp_4);
assert_eq!(conversation.messages[3].author, Author::Assistant);
assert_eq!(
conversation.messages[3].content,
MessageContent::Text("Hi there!\nI'm doing well, thank you.".to_string())
);
assert_eq!(conversation.messages[3].timestamp, timestamp_2);
assert_eq!(conversation.messages[4].author, Author::User);
assert_eq!(
conversation.messages[4].content,
MessageContent::Text("That's great!".to_string())
);
assert_eq!(conversation.messages[4].timestamp, timestamp_3);
}
#[test]
fn combine_consecutive_messages_clears_sender_id_for_mixed_sender_turns() {
let timestamp_1 = Utc.with_ymd_and_hms(2024, 9, 20, 18, 34, 15).unwrap();
let timestamp_2 = Utc.with_ymd_and_hms(2024, 9, 20, 18, 34, 16).unwrap();
let sender_1 = OwnedUserId::try_from("@alice:example.com").unwrap();
let sender_2 = OwnedUserId::try_from("@bob:example.com").unwrap();
let conversation = Conversation {
messages: vec![
Message {
author: Author::User,
sender_id: Some(sender_1),
content: MessageContent::Text("Hello".to_string()),
timestamp: timestamp_1,
},
Message {
author: Author::User,
sender_id: Some(sender_2),
content: MessageContent::Text("Hi there".to_string()),
timestamp: timestamp_2,
},
],
};
let conversation = conversation.combine_consecutive_messages();
assert_eq!(conversation.messages.len(), 1);
assert_eq!(conversation.messages[0].sender_id, None);
assert_eq!(
conversation.messages[0].content,
MessageContent::Text("Hello\nHi there".to_string())
);
assert_eq!(conversation.messages[0].timestamp, timestamp_1);
}
}