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.
326 lines
11 KiB
Rust
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);
|
|
}
|
|
}
|