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.
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
use chrono::{DateTime, Utc};
|
||||
use mxlink::matrix_sdk::ruma::OwnedUserId;
|
||||
use mxlink::matrix_sdk::ruma::events::room::message::{
|
||||
FileMessageEventContent, ImageMessageEventContent,
|
||||
};
|
||||
@@ -16,6 +17,7 @@ pub enum Author {
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Message {
|
||||
pub author: Author,
|
||||
pub sender_id: Option<OwnedUserId>,
|
||||
pub timestamp: DateTime<Utc>,
|
||||
pub content: MessageContent,
|
||||
}
|
||||
@@ -104,6 +106,11 @@ impl Conversation {
|
||||
///
|
||||
/// 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.
|
||||
@@ -134,6 +141,10 @@ impl Conversation {
|
||||
text.push('\n');
|
||||
text.push_str(message_text_content);
|
||||
}
|
||||
|
||||
if last_message.sender_id != message.sender_id {
|
||||
last_message.sender_id = None;
|
||||
}
|
||||
}
|
||||
|
||||
Conversation {
|
||||
@@ -150,7 +161,7 @@ impl Conversation {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use chrono::{TimeZone, Utc};
|
||||
use mxlink::matrix_sdk::ruma::OwnedMxcUri;
|
||||
use mxlink::matrix_sdk::ruma::{OwnedMxcUri, OwnedUserId};
|
||||
use mxlink::mime;
|
||||
|
||||
#[test]
|
||||
@@ -173,21 +184,25 @@ mod tests {
|
||||
// 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,
|
||||
@@ -197,28 +212,33 @@ mod tests {
|
||||
},
|
||||
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,
|
||||
},
|
||||
@@ -267,4 +287,39 @@ mod tests {
|
||||
);
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,6 +20,7 @@ fn test_messages_by_the_bot_are_identified_correctly() {
|
||||
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id).unwrap();
|
||||
|
||||
assert_eq!(llm_message.author, Author::Assistant);
|
||||
assert_eq!(llm_message.sender_id, Some(bot_user_id.clone()));
|
||||
assert_eq!(
|
||||
llm_message.content,
|
||||
MessageContent::Text("Hello!".to_string())
|
||||
@@ -47,6 +48,7 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con
|
||||
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id).unwrap();
|
||||
|
||||
assert_eq!(llm_message.author, Author::User);
|
||||
assert_eq!(llm_message.sender_id, None);
|
||||
assert_eq!(
|
||||
llm_message.content,
|
||||
MessageContent::Text(source_message_text.to_string())
|
||||
@@ -75,6 +77,30 @@ fn test_notice_error_messages_by_bot_are_ignored() {
|
||||
assert!(llm_message.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_user_messages_preserve_sender_id() {
|
||||
let bot_user_id =
|
||||
OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID");
|
||||
|
||||
let user_id = OwnedUserId::try_from("@alice:example.com").expect("Failed to parse user ID");
|
||||
|
||||
let matrix_message = super::super::matrix::MatrixMessage {
|
||||
sender_id: user_id.clone(),
|
||||
content: super::super::matrix::MatrixMessageContent::Text("Hello!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
|
||||
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id).unwrap();
|
||||
|
||||
assert_eq!(llm_message.author, Author::User);
|
||||
assert_eq!(llm_message.sender_id, Some(user_id));
|
||||
assert_eq!(
|
||||
llm_message.content,
|
||||
MessageContent::Text("Hello!".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_other_notice_messages_by_the_bot_are_ignored() {
|
||||
// Also see `test_notice_error_messages_by_bot_are_ignored()`.
|
||||
|
||||
@@ -89,6 +89,7 @@ pub mod test {
|
||||
|
||||
let message = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("Hello there!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -108,6 +109,7 @@ pub mod test {
|
||||
|
||||
let prompt = super::Message {
|
||||
author: super::Author::Prompt,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("You are a bot!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -122,6 +124,7 @@ pub mod test {
|
||||
|
||||
let first = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("Hello there!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -136,6 +139,7 @@ pub mod test {
|
||||
|
||||
let second = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("Hello!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -150,6 +154,7 @@ pub mod test {
|
||||
|
||||
let third = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text(
|
||||
"This is the 3rd message in this conversation. It shall be preserved.".to_owned(),
|
||||
),
|
||||
@@ -166,6 +171,7 @@ pub mod test {
|
||||
|
||||
let forth = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text(
|
||||
"This is yet another message that shall be preserved.".to_owned(),
|
||||
),
|
||||
@@ -213,6 +219,7 @@ pub mod test {
|
||||
|
||||
let prompt = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("あなたはボットです。".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -227,6 +234,7 @@ pub mod test {
|
||||
|
||||
let first = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("こんにちは!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -241,6 +249,7 @@ pub mod test {
|
||||
|
||||
let second = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text("こんにちは。今日は元気ですか。".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -255,6 +264,7 @@ pub mod test {
|
||||
|
||||
let third = super::Message {
|
||||
author: super::Author::User,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text(
|
||||
"これは第3のメッセージなので、保存されます。".to_string(),
|
||||
),
|
||||
@@ -271,6 +281,7 @@ pub mod test {
|
||||
|
||||
let forth = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
sender_id: None,
|
||||
content: super::MessageContent::Text(
|
||||
"これはもう一つの保存されますメッセージです。".to_string(),
|
||||
),
|
||||
|
||||
@@ -17,14 +17,17 @@ pub fn convert_matrix_message_to_llm_message(
|
||||
|
||||
fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
match &matrix_message.content {
|
||||
MatrixMessageContent::Text(text) => {
|
||||
convert_bot_text_message(text, &matrix_message.timestamp)
|
||||
}
|
||||
MatrixMessageContent::Text(text) => convert_bot_text_message(
|
||||
text,
|
||||
&matrix_message.timestamp,
|
||||
matrix_message.sender_id.clone(),
|
||||
),
|
||||
MatrixMessageContent::Notice(text) => {
|
||||
convert_bot_notice_message(text, &matrix_message.timestamp)
|
||||
}
|
||||
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => Some(Message {
|
||||
author: Author::Assistant,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::Image(ImageDetails::new(
|
||||
image_content.clone(),
|
||||
mime_type.clone(),
|
||||
@@ -34,6 +37,7 @@ fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
}),
|
||||
MatrixMessageContent::File(file_content, mime_type, media_bytes) => Some(Message {
|
||||
author: Author::Assistant,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::File(FileDetails::new(
|
||||
file_content.clone(),
|
||||
mime_type.clone(),
|
||||
@@ -47,9 +51,11 @@ fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
fn convert_bot_text_message(
|
||||
text: &str,
|
||||
timestamp: &chrono::DateTime<chrono::Utc>,
|
||||
sender_id: OwnedUserId,
|
||||
) -> Option<Message> {
|
||||
Some(Message {
|
||||
author: Author::Assistant,
|
||||
sender_id: Some(sender_id),
|
||||
content: MessageContent::Text(text.to_owned()),
|
||||
timestamp: timestamp.to_owned(),
|
||||
})
|
||||
@@ -68,8 +74,10 @@ fn convert_bot_notice_message(
|
||||
|
||||
if let Some(text) = text_to_speech_utils::parse_transcribed_message_text(text) {
|
||||
// This is a transcription message. We remove the prefix and consider it as a message sent by the user.
|
||||
// sender_id is None because the original speaker is unknown.
|
||||
return Some(Message {
|
||||
author: Author::User,
|
||||
sender_id: None,
|
||||
content: MessageContent::Text(text.to_owned()),
|
||||
timestamp: timestamp.to_owned(),
|
||||
});
|
||||
@@ -82,16 +90,19 @@ fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
match &matrix_message.content {
|
||||
MatrixMessageContent::Text(text) => Some(Message {
|
||||
author: Author::User,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::Text(text.clone()),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
}),
|
||||
MatrixMessageContent::Notice(text) => Some(Message {
|
||||
author: Author::User,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::Text(text.clone()),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
}),
|
||||
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => Some(Message {
|
||||
author: Author::User,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::Image(ImageDetails::new(
|
||||
image_content.clone(),
|
||||
mime_type.clone(),
|
||||
@@ -101,6 +112,7 @@ fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
}),
|
||||
MatrixMessageContent::File(file_content, mime_type, media_bytes) => Some(Message {
|
||||
author: Author::User,
|
||||
sender_id: Some(matrix_message.sender_id.clone()),
|
||||
content: MessageContent::File(FileDetails::new(
|
||||
file_content.clone(),
|
||||
mime_type.clone(),
|
||||
|
||||
Reference in New Issue
Block a user