Initial work on Vision support in text conversations and Image Editing
This is a huge patch which does some major refactoring like: - renaming "Image Generation" to "Image Creation" in most places, to better match its new command (`!bai image create`) - relocating image creation command (`!bai image` -> `!bai image create`), so it wouldn't conflict with the new image editing command (`!bai image edit`) - introducing a new image editing command (`!bai image edit`), which is meant to work only with the OpenAI provider, but doesn't fully work yet due to https://github.com/64bit/async-openai/issues/364, though a next patch will fix it - adding support for reading images off of Matrix conversations and forwarding them to text conversations. Works for OpenAI, but not for Anthropic yet (requires custom patches) and not for OpenAI-Compat (no support for images there) - relocating some utils around (base64, mime)
This commit is contained in:
@@ -1,4 +1,8 @@
|
||||
use chrono::{DateTime, Utc};
|
||||
use mxlink::matrix_sdk::ruma::events::room::message::ImageMessageEventContent;
|
||||
use mxlink::mime::Mime;
|
||||
|
||||
use crate::agent::provider::ImageSource;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum Author {
|
||||
@@ -10,10 +14,51 @@ pub enum Author {
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Message {
|
||||
pub author: Author,
|
||||
pub message_text: String,
|
||||
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 Into<ImageSource> for &ImageDetails {
|
||||
fn into(self) -> ImageSource {
|
||||
ImageSource::new(self.filename(), self.data.clone(), self.mime.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum MessageContent {
|
||||
Text(String),
|
||||
Image(ImageDetails),
|
||||
}
|
||||
|
||||
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()
|
||||
},
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Debug)]
|
||||
pub struct Conversation {
|
||||
pub messages: Vec<Message>,
|
||||
@@ -28,27 +73,32 @@ impl Conversation {
|
||||
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_author: Option<Author> = None;
|
||||
let mut last_seen_text_from_author: Option<Author> = None;
|
||||
|
||||
for message in &self.messages {
|
||||
let Some(last_seen_author_clone) = last_seen_author.clone() else {
|
||||
last_seen_author = Some(message.author.clone());
|
||||
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_author = Some(message.author.clone());
|
||||
last_seen_text_from_author = Some(message.author.clone());
|
||||
new_messages.push(message.clone());
|
||||
continue;
|
||||
}
|
||||
|
||||
new_messages.last_mut().unwrap().message_text.push('\n');
|
||||
new_messages
|
||||
.last_mut()
|
||||
.unwrap()
|
||||
.message_text
|
||||
.push_str(&message.message_text);
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
Conversation {
|
||||
@@ -65,48 +115,76 @@ impl Conversation {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use chrono::{TimeZone, Utc};
|
||||
use mxlink::matrix_sdk::ruma::OwnedMxcUri;
|
||||
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, 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, 15).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,
|
||||
message_text: "Hello".to_string(),
|
||||
content: MessageContent::Text("Hello".to_string()),
|
||||
timestamp: timestamp_1,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "How are you?".to_string(),
|
||||
content: MessageContent::Text("How are you?".to_string()),
|
||||
timestamp: timestamp_2,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "I'm OK, btw.".to_string(),
|
||||
content: MessageContent::Text("I'm OK, btw.".to_string()),
|
||||
timestamp: timestamp_3,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Image(ImageDetails::new(
|
||||
image_event_content.clone(),
|
||||
mime::IMAGE_PNG,
|
||||
vec![],
|
||||
)),
|
||||
timestamp: timestamp_4,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Text("Above is an image.".to_string()),
|
||||
timestamp: timestamp_4,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Text("Would you take a look at it?".to_string()),
|
||||
timestamp: timestamp_4,
|
||||
},
|
||||
// Assistant's turn
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Hi there!".to_string(),
|
||||
content: MessageContent::Text("Hi there!".to_string()),
|
||||
timestamp: timestamp_2,
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "I'm doing well, thank you.".to_string(),
|
||||
content: MessageContent::Text("I'm doing well, thank you.".to_string()),
|
||||
timestamp: timestamp_3,
|
||||
},
|
||||
// User's turn
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "That's great!".to_string(),
|
||||
content: MessageContent::Text("That's great!".to_string()),
|
||||
timestamp: timestamp_3,
|
||||
},
|
||||
],
|
||||
@@ -114,23 +192,44 @@ mod tests {
|
||||
|
||||
let conversation = conversation.combine_consecutive_messages();
|
||||
|
||||
assert_eq!(conversation.messages.len(), 3);
|
||||
assert_eq!(conversation.messages.len(), 5);
|
||||
|
||||
assert_eq!(conversation.messages[0].author, Author::User);
|
||||
assert_eq!(
|
||||
conversation.messages[0].message_text,
|
||||
"Hello\nHow are you?\nI'm OK, btw."
|
||||
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::Assistant);
|
||||
assert_eq!(conversation.messages[1].author, Author::User);
|
||||
assert_eq!(
|
||||
conversation.messages[1].message_text,
|
||||
"Hi there!\nI'm doing well, thank you."
|
||||
conversation.messages[1].content,
|
||||
MessageContent::Image(ImageDetails::new(
|
||||
image_event_content.clone(),
|
||||
mime::IMAGE_PNG,
|
||||
vec![],
|
||||
))
|
||||
);
|
||||
assert_eq!(conversation.messages[1].timestamp, timestamp_2);
|
||||
|
||||
assert_eq!(conversation.messages[2].author, Author::User);
|
||||
assert_eq!(conversation.messages[2].message_text, "That's great!");
|
||||
assert_eq!(conversation.messages[2].timestamp, timestamp_3);
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,8 +12,7 @@ fn test_messages_by_the_bot_are_identified_correctly() {
|
||||
|
||||
let matrix_message = super::super::matrix::MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: super::super::matrix::MatrixMessageType::Text,
|
||||
message_text: "Hello!".to_owned(),
|
||||
content: super::super::matrix::MatrixMessageContent::Text("Hello!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -21,12 +20,11 @@ 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.message_text, "Hello!");
|
||||
assert_eq!(llm_message.content, MessageContent::Text("Hello!".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user()
|
||||
{
|
||||
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user() {
|
||||
let bot_user_id =
|
||||
OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID");
|
||||
|
||||
@@ -37,8 +35,7 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con
|
||||
|
||||
let matrix_message = super::super::matrix::MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: super::super::matrix::MatrixMessageType::Notice,
|
||||
message_text,
|
||||
content: super::super::matrix::MatrixMessageContent::Notice(message_text),
|
||||
mentioned_users: vec![],
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -46,7 +43,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.message_text, source_message_text);
|
||||
assert_eq!(llm_message.content, MessageContent::Text(source_message_text.to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -61,8 +58,7 @@ fn test_notice_error_messages_by_bot_are_ignored() {
|
||||
|
||||
let matrix_message = super::super::matrix::MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: super::super::matrix::MatrixMessageType::Notice,
|
||||
message_text,
|
||||
content: super::super::matrix::MatrixMessageContent::Notice(message_text),
|
||||
mentioned_users: vec![],
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
@@ -86,8 +82,7 @@ fn test_other_notice_messages_by_the_bot_are_ignored() {
|
||||
|
||||
let matrix_message = super::super::matrix::MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: super::super::matrix::MatrixMessageType::Notice,
|
||||
message_text: message_text.to_owned(),
|
||||
content: super::super::matrix::MatrixMessageContent::Notice(message_text.to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
|
||||
@@ -2,7 +2,7 @@ use tiktoken_rs::CoreBPE;
|
||||
use tiktoken_rs::get_bpe_from_tokenizer;
|
||||
use tiktoken_rs::tokenizer;
|
||||
|
||||
use super::{Author, Message};
|
||||
use super::{Author, Message, MessageContent};
|
||||
|
||||
fn get_bpe_for_model(model: &str) -> CoreBPE {
|
||||
let tokenizer = tokenizer::get_tokenizer(model)
|
||||
@@ -71,7 +71,10 @@ fn calculate_token_size_for_message(bpe: &CoreBPE, model: &str, message: &Messag
|
||||
Author::Prompt => bpe.encode_with_special_tokens("system").len() as i32,
|
||||
};
|
||||
|
||||
let text_length = bpe.encode_with_special_tokens(&message.message_text).len() as i32;
|
||||
let text_length = match &message.content {
|
||||
MessageContent::Text(text) => bpe.encode_with_special_tokens(text).len() as i32,
|
||||
MessageContent::Image(..) => 0,
|
||||
};
|
||||
|
||||
(text_length + role_length + tokens_per_message + tokens_per_name) as u32
|
||||
}
|
||||
@@ -85,7 +88,7 @@ pub mod test {
|
||||
|
||||
let message = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "Hello there!".to_owned(),
|
||||
content: super::MessageContent::Text("Hello there!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
|
||||
@@ -104,7 +107,7 @@ pub mod test {
|
||||
|
||||
let prompt = super::Message {
|
||||
author: super::Author::Prompt,
|
||||
message_text: "You are a bot!".to_owned(),
|
||||
content: super::MessageContent::Text("You are a bot!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let prompt_length = 10;
|
||||
@@ -118,7 +121,7 @@ pub mod test {
|
||||
|
||||
let first = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "Hello there!".to_owned(),
|
||||
content: super::MessageContent::Text("Hello there!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let first_length = 8;
|
||||
@@ -132,7 +135,7 @@ pub mod test {
|
||||
|
||||
let second = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
message_text: "Hello!".to_owned(),
|
||||
content: super::MessageContent::Text("Hello!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let second_length = 7;
|
||||
@@ -146,8 +149,10 @@ pub mod test {
|
||||
|
||||
let third = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "This is the 3rd message in this conversation. It shall be preserved."
|
||||
.to_owned(),
|
||||
content: super::MessageContent::Text(
|
||||
"This is the 3rd message in this conversation. It shall be preserved."
|
||||
.to_owned(),
|
||||
),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let third_length = 21;
|
||||
@@ -161,7 +166,9 @@ pub mod test {
|
||||
|
||||
let forth = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
message_text: "This is yet another message that shall be preserved.".to_owned(),
|
||||
content: super::MessageContent::Text(
|
||||
"This is yet another message that shall be preserved.".to_owned(),
|
||||
),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let forth_length = 15;
|
||||
@@ -186,13 +193,13 @@ pub mod test {
|
||||
assert_eq!(2, new_conversation_messages.len());
|
||||
|
||||
assert_eq!(
|
||||
new_conversation_messages.first().unwrap().message_text,
|
||||
third.message_text
|
||||
new_conversation_messages.first().unwrap().content,
|
||||
third.content
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
new_conversation_messages.last().unwrap().message_text,
|
||||
forth.message_text
|
||||
new_conversation_messages.last().unwrap().content,
|
||||
forth.content
|
||||
);
|
||||
}
|
||||
|
||||
@@ -206,7 +213,7 @@ pub mod test {
|
||||
|
||||
let prompt = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "あなたはボットです。".to_owned(),
|
||||
content: super::MessageContent::Text("あなたはボットです。".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let prompt_length = 14;
|
||||
@@ -220,7 +227,7 @@ pub mod test {
|
||||
|
||||
let first = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "こんにちは!".to_owned(),
|
||||
content: super::MessageContent::Text("こんにちは!".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let first_length = 7;
|
||||
@@ -234,7 +241,7 @@ pub mod test {
|
||||
|
||||
let second = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
message_text: "こんにちは。今日は元気ですか。".to_owned(),
|
||||
content: super::MessageContent::Text("こんにちは。今日は元気ですか。".to_string()),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let second_length = 15;
|
||||
@@ -248,7 +255,9 @@ pub mod test {
|
||||
|
||||
let third = super::Message {
|
||||
author: super::Author::User,
|
||||
message_text: "これは第3のメッセージなので、保存されます。".to_owned(),
|
||||
content: super::MessageContent::Text(
|
||||
"これは第3のメッセージなので、保存されます。".to_string(),
|
||||
),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let third_length = 22;
|
||||
@@ -262,7 +271,9 @@ pub mod test {
|
||||
|
||||
let forth = super::Message {
|
||||
author: super::Author::Assistant,
|
||||
message_text: "これはもう一つの保存されますメッセージです。".to_owned(),
|
||||
content: super::MessageContent::Text(
|
||||
"これはもう一つの保存されますメッセージです。".to_string(),
|
||||
),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
let forth_length = 21;
|
||||
@@ -287,13 +298,13 @@ pub mod test {
|
||||
assert_eq!(2, new_conversation_messages.len());
|
||||
|
||||
assert_eq!(
|
||||
new_conversation_messages.first().unwrap().message_text,
|
||||
third.message_text
|
||||
new_conversation_messages.first().unwrap().content,
|
||||
third.content
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
new_conversation_messages.last().unwrap().message_text,
|
||||
forth.message_text
|
||||
new_conversation_messages.last().unwrap().content,
|
||||
forth.content
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use matrix_sdk::ruma::OwnedUserId;
|
||||
use mxlink::matrix_sdk::ruma::OwnedUserId;
|
||||
|
||||
use super::{Author, Message};
|
||||
use crate::conversation::matrix::{MatrixMessage, MatrixMessageType};
|
||||
use super::entity::{Author, ImageDetails, Message, MessageContent};
|
||||
use crate::conversation::matrix::{MatrixMessage, MatrixMessageContent};
|
||||
use crate::utils::text_to_speech as text_to_speech_utils;
|
||||
|
||||
pub fn convert_matrix_message_to_llm_message(
|
||||
@@ -16,12 +16,23 @@ pub fn convert_matrix_message_to_llm_message(
|
||||
}
|
||||
|
||||
fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
match matrix_message.message_type {
|
||||
MatrixMessageType::Text => {
|
||||
convert_bot_text_message(&matrix_message.message_text, &matrix_message.timestamp)
|
||||
match &matrix_message.content {
|
||||
MatrixMessageContent::Text(text) => {
|
||||
convert_bot_text_message(text, &matrix_message.timestamp)
|
||||
}
|
||||
MatrixMessageType::Notice => {
|
||||
convert_bot_notice_message(&matrix_message.message_text, &matrix_message.timestamp)
|
||||
MatrixMessageContent::Notice(text) => {
|
||||
convert_bot_notice_message(text, &matrix_message.timestamp)
|
||||
}
|
||||
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => {
|
||||
Some(Message {
|
||||
author: Author::Assistant,
|
||||
content: MessageContent::Image(ImageDetails::new(
|
||||
image_content.clone(),
|
||||
mime_type.clone(),
|
||||
media_bytes.clone()
|
||||
)),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -32,7 +43,7 @@ fn convert_bot_text_message(
|
||||
) -> Option<Message> {
|
||||
Some(Message {
|
||||
author: Author::Assistant,
|
||||
message_text: text.to_owned(),
|
||||
content: MessageContent::Text(text.to_owned()),
|
||||
timestamp: timestamp.to_owned(),
|
||||
})
|
||||
}
|
||||
@@ -52,7 +63,7 @@ fn convert_bot_notice_message(
|
||||
// This is a transcription message. We remove the prefix and consider it as a message sent by the user.
|
||||
return Some(Message {
|
||||
author: Author::User,
|
||||
message_text: text.to_owned(),
|
||||
content: MessageContent::Text(text.to_owned()),
|
||||
timestamp: timestamp.to_owned(),
|
||||
});
|
||||
}
|
||||
@@ -61,9 +72,31 @@ fn convert_bot_notice_message(
|
||||
}
|
||||
|
||||
fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
Some(Message {
|
||||
author: Author::User,
|
||||
message_text: matrix_message.message_text.clone(),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
})
|
||||
match &matrix_message.content {
|
||||
MatrixMessageContent::Text(text) => {
|
||||
Some(Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Text(text.clone()),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
})
|
||||
}
|
||||
MatrixMessageContent::Notice(text) => {
|
||||
Some(Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Text(text.clone()),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
})
|
||||
}
|
||||
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => {
|
||||
Some(Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Image(ImageDetails::new(
|
||||
image_content.clone(),
|
||||
mime_type.clone(),
|
||||
media_bytes.clone(),
|
||||
)),
|
||||
timestamp: matrix_message.timestamp.to_owned(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,20 +2,22 @@ use chrono::{DateTime, Utc};
|
||||
use regex::Regex;
|
||||
|
||||
use mxlink::matrix_sdk::ruma::OwnedUserId;
|
||||
use mxlink::matrix_sdk::ruma::events::room::message::ImageMessageEventContent;
|
||||
use mxlink::mime::Mime;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct MatrixMessage {
|
||||
pub sender_id: OwnedUserId,
|
||||
pub message_type: MatrixMessageType,
|
||||
pub message_text: String,
|
||||
pub content: MatrixMessageContent,
|
||||
pub mentioned_users: Vec<OwnedUserId>,
|
||||
pub timestamp: DateTime<Utc>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub enum MatrixMessageType {
|
||||
Text,
|
||||
Notice,
|
||||
pub enum MatrixMessageContent {
|
||||
Text(String),
|
||||
Notice(String),
|
||||
Image(ImageMessageEventContent, Mime, Vec<u8>),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
|
||||
@@ -6,6 +6,6 @@ mod utils;
|
||||
pub(crate) use room_display_name_fetcher::RoomDisplayNameFetcher;
|
||||
pub(crate) use room_event_fetcher::RoomEventFetcher;
|
||||
|
||||
pub(crate) use entity::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType};
|
||||
pub(crate) use entity::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageContent};
|
||||
|
||||
pub(crate) use utils::*;
|
||||
|
||||
@@ -18,9 +18,11 @@ use mxlink::matrix_sdk::{
|
||||
},
|
||||
};
|
||||
use mxlink::{MatrixLink, ThreadGetMessagesParams, ThreadInfo};
|
||||
use tracing::Instrument;
|
||||
|
||||
use super::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType, RoomEventFetcher};
|
||||
use super::{MatrixMessage, MatrixMessageContent, MatrixMessageProcessingParams, RoomEventFetcher};
|
||||
use crate::entity::{InteractionContext, InteractionTrigger, MessagePayload};
|
||||
use crate::utils::mime::get_mime_type_from_file_name;
|
||||
|
||||
struct DetailedMessagePayload {
|
||||
is_mentioning_bot: bool,
|
||||
@@ -28,7 +30,7 @@ struct DetailedMessagePayload {
|
||||
}
|
||||
|
||||
pub async fn get_matrix_messages_in_thread(
|
||||
matrix_link: MatrixLink,
|
||||
matrix_link: &MatrixLink,
|
||||
room: &Room,
|
||||
thread_id: OwnedEventId,
|
||||
) -> Result<Vec<MatrixMessage>, mxlink::matrix_sdk::Error> {
|
||||
@@ -40,18 +42,18 @@ pub async fn get_matrix_messages_in_thread(
|
||||
let mut messages: Vec<MatrixMessage> = Vec::new();
|
||||
|
||||
for matrix_native_message in messages_native {
|
||||
let Some(message) = convert_matrix_native_event_to_matrix_message(&matrix_native_message)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
|
||||
|
||||
messages.push(message);
|
||||
if let Some(message) = message_result {
|
||||
messages.push(message);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(messages)
|
||||
}
|
||||
|
||||
pub async fn get_matrix_messages_in_reply_chain(
|
||||
matrix_link: &MatrixLink,
|
||||
event_fetcher: &Arc<RoomEventFetcher>,
|
||||
room: &Room,
|
||||
event_id: OwnedEventId,
|
||||
@@ -62,12 +64,11 @@ pub async fn get_matrix_messages_in_reply_chain(
|
||||
let mut messages: Vec<MatrixMessage> = Vec::new();
|
||||
|
||||
for matrix_native_message in messages_native {
|
||||
let Some(message) = convert_matrix_native_event_to_matrix_message(&matrix_native_message)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
|
||||
|
||||
messages.push(message);
|
||||
if let Some(message) = message_result {
|
||||
messages.push(message);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(messages)
|
||||
@@ -150,30 +151,34 @@ pub async fn process_matrix_messages(
|
||||
let mut message = message.clone();
|
||||
|
||||
if i == 0 && !params.first_message_prefixes_to_strip.is_empty() {
|
||||
let mut message_text = message.message_text.clone();
|
||||
if let MatrixMessageContent::Text(message_text) = &message.content {
|
||||
let mut message_text = message_text.clone();
|
||||
|
||||
for prefix in ¶ms.first_message_prefixes_to_strip {
|
||||
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
|
||||
message_text = message_text_stripped.to_owned();
|
||||
for prefix in ¶ms.first_message_prefixes_to_strip {
|
||||
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
|
||||
message_text = message_text_stripped.to_owned();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
message.message_text = message_text.trim().to_owned();
|
||||
message.content = MatrixMessageContent::Text(message_text.trim().to_owned());
|
||||
}
|
||||
}
|
||||
|
||||
// We only strip `bot_user_prefixes_to_strip`-defined prefixes from messages that mention the bot user.
|
||||
if !params.bot_user_prefixes_to_strip.is_empty()
|
||||
&& message.mentioned_users.contains(¶ms.bot_user_id)
|
||||
{
|
||||
let mut message_text = message.message_text.clone();
|
||||
if let MatrixMessageContent::Text(message_text) = &message.content {
|
||||
let mut message_text = message_text.clone();
|
||||
|
||||
for prefix in ¶ms.bot_user_prefixes_to_strip {
|
||||
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
|
||||
message_text = message_text_stripped.to_owned();
|
||||
for prefix in ¶ms.bot_user_prefixes_to_strip {
|
||||
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
|
||||
message_text = message_text_stripped.to_owned();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
message.message_text = message_text.trim().to_owned();
|
||||
message.content = MatrixMessageContent::Text(message_text.trim().to_owned());
|
||||
}
|
||||
}
|
||||
|
||||
messages_filtered.push(message);
|
||||
@@ -207,23 +212,25 @@ fn is_message_from_allowed_sender(
|
||||
false
|
||||
}
|
||||
|
||||
pub fn convert_matrix_native_event_to_matrix_message(
|
||||
pub async fn convert_matrix_native_event_to_matrix_message(
|
||||
matrix_link: &MatrixLink,
|
||||
matrix_native_event: &AnySyncMessageLikeEvent,
|
||||
) -> Option<MatrixMessage> {
|
||||
) -> Result<Option<MatrixMessage>, mxlink::matrix_sdk::Error> {
|
||||
let Some(content) = matrix_native_event.original_content() else {
|
||||
// Redacted message
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let AnyMessageLikeEventContent::RoomMessage(room_message) = content else {
|
||||
// Some state event, etc.
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let (text, is_notice) = match &room_message.msgtype {
|
||||
MessageType::Text(text_content) => (text_content.body.clone(), false),
|
||||
MessageType::Notice(notice_content) => (notice_content.body.clone(), true),
|
||||
_ => return None,
|
||||
MessageType::Image(image_content) => (image_content.body.clone(), false),
|
||||
_ => return Ok(None),
|
||||
};
|
||||
|
||||
let is_reply = matches!(room_message.relates_to, Some(Relation::Reply { .. }));
|
||||
@@ -248,17 +255,45 @@ pub fn convert_matrix_native_event_to_matrix_message(
|
||||
.map(|m| m.user_ids.iter().map(|u| u.to_owned()).collect())
|
||||
.unwrap_or(vec![]);
|
||||
|
||||
Some(MatrixMessage {
|
||||
if let MessageType::Image(image_content) = &room_message.msgtype {
|
||||
let media_request = mxlink::matrix_sdk::media::MediaRequestParameters {
|
||||
source: image_content.source.to_owned(),
|
||||
format: mxlink::matrix_sdk::media::MediaFormat::File,
|
||||
};
|
||||
|
||||
let file_name = image_content.filename.clone().unwrap_or(image_content.body.clone());
|
||||
|
||||
let mime_type = get_mime_type_from_file_name(&file_name);
|
||||
|
||||
tracing::debug!("Determined mime type {} for file {}", mime_type, file_name);
|
||||
|
||||
let span = tracing::debug_span!("get_media_content", file_name = %file_name, mime_type = %mime_type);
|
||||
|
||||
let media_bytes = matrix_link
|
||||
.client()
|
||||
.media()
|
||||
.get_media_content(&media_request, true)
|
||||
.instrument(span)
|
||||
.await?;
|
||||
|
||||
return Ok(Some(MatrixMessage {
|
||||
sender_id: matrix_native_event.sender().to_owned(),
|
||||
content: MatrixMessageContent::Image(image_content.clone(), mime_type, media_bytes),
|
||||
mentioned_users,
|
||||
timestamp,
|
||||
}));
|
||||
}
|
||||
|
||||
Ok(Some(MatrixMessage {
|
||||
sender_id: matrix_native_event.sender().to_owned(),
|
||||
message_type: if is_notice {
|
||||
MatrixMessageType::Notice
|
||||
content: if is_notice {
|
||||
MatrixMessageContent::Notice(text)
|
||||
} else {
|
||||
MatrixMessageType::Text
|
||||
MatrixMessageContent::Text(text)
|
||||
},
|
||||
message_text: text,
|
||||
mentioned_users,
|
||||
timestamp,
|
||||
})
|
||||
}))
|
||||
}
|
||||
|
||||
/// Determines the interaction context for an incoming (new) room event.
|
||||
|
||||
@@ -3,7 +3,7 @@ use chrono::{TimeZone, Utc};
|
||||
use mxlink::matrix_sdk::ruma::OwnedUserId;
|
||||
|
||||
use crate::conversation::matrix::{
|
||||
MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType,
|
||||
MatrixMessage, MatrixMessageContent, MatrixMessageProcessingParams,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -17,24 +17,21 @@ fn is_message_from_allowed_sender() {
|
||||
|
||||
let bot_message = MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let allowed_user_message = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let unallowed_user_message = MatrixMessage {
|
||||
sender_id: unallowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
@@ -88,48 +85,42 @@ async fn process_matrix_messages() {
|
||||
|
||||
let allowed_user_message = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello from the user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello from the user!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let allowed_user_message_with_prefix = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "!bai Hello from the user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("!bai Hello from the user!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let allowed_user_message_with_prefix_no_space = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "!baiHello from the user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("!baiHello from the user!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let allowed_user_message_with_prefix_full_width_space = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "!bai Hello from the user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("!bai Hello from the user!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let bot_message = MatrixMessage {
|
||||
sender_id: bot_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello from the bot!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello from the bot!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let allowed_user_message_with_bot_mention = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "@baibot: Hello from the user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("@baibot: Hello from the user!".to_owned()),
|
||||
mentioned_users: vec![bot_user_id.to_owned()],
|
||||
timestamp,
|
||||
};
|
||||
@@ -137,16 +128,14 @@ async fn process_matrix_messages() {
|
||||
// The message text is the same as above - it mentions the bot, but the actually-mentioned user is another user.
|
||||
let allowed_user_message_with_another_user_mention = MatrixMessage {
|
||||
sender_id: allowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: allowed_user_message_with_bot_mention.message_text.clone(),
|
||||
content: allowed_user_message_with_bot_mention.content.clone(),
|
||||
mentioned_users: vec![allowed_user_id.to_owned()],
|
||||
timestamp,
|
||||
};
|
||||
|
||||
let unallowed_user_message = MatrixMessage {
|
||||
sender_id: unallowed_user_id.to_owned(),
|
||||
message_type: MatrixMessageType::Text,
|
||||
message_text: "Hello from an unallowed user!".to_owned(),
|
||||
content: MatrixMessageContent::Text("Hello from an unallowed user!".to_owned()),
|
||||
mentioned_users: vec![],
|
||||
timestamp,
|
||||
};
|
||||
@@ -285,7 +274,10 @@ async fn process_matrix_messages() {
|
||||
|
||||
let processed_message_texts = processed_messages
|
||||
.iter()
|
||||
.map(|message| message.message_text.clone())
|
||||
.map(|message| match &message.content {
|
||||
MatrixMessageContent::Text(text) => text.clone(),
|
||||
_ => "".to_owned(),
|
||||
})
|
||||
.collect::<Vec<String>>();
|
||||
|
||||
assert_eq!(
|
||||
|
||||
@@ -12,7 +12,7 @@ use super::matrix::{
|
||||
};
|
||||
|
||||
pub async fn create_llm_conversation_for_matrix_thread(
|
||||
matrix_link: MatrixLink,
|
||||
matrix_link: &MatrixLink,
|
||||
room: &mxlink::matrix_sdk::Room,
|
||||
thread_id: OwnedEventId,
|
||||
params: &MatrixMessageProcessingParams,
|
||||
@@ -27,12 +27,13 @@ pub async fn create_llm_conversation_for_matrix_thread(
|
||||
}
|
||||
|
||||
pub async fn create_llm_conversation_for_matrix_reply_chain(
|
||||
matrix_link: &MatrixLink,
|
||||
event_fetcher: &Arc<RoomEventFetcher>,
|
||||
room: &mxlink::matrix_sdk::Room,
|
||||
event_id: OwnedEventId,
|
||||
params: &MatrixMessageProcessingParams,
|
||||
) -> Result<Conversation, mxlink::matrix_sdk::Error> {
|
||||
let messages = get_matrix_messages_in_reply_chain(event_fetcher, room, event_id).await?;
|
||||
let messages = get_matrix_messages_in_reply_chain(matrix_link, event_fetcher, room, event_id).await?;
|
||||
|
||||
let llm_messages = filter_messages_and_convert_to_llm_messages(messages, params).await;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user